194 lines
4.5 KiB
Go
194 lines
4.5 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
|
|
// SPDX-License-Identifier: MIT
|
|||
|
|
|
|||
|
|
package grad
|
|||
|
|
|
|||
|
|
import (
|
|||
|
|
"math"
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|||
|
|
"testing"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// TestHessianQuadratic pins the exact case: for f(x) = ½xᵀAx + bᵀx the
|
|||
|
|
// Hessian is A whatever the point.
|
|||
|
|
func TestHessianQuadratic(t *testing.T) {
|
|||
|
|
a := []float64{4, 1, 1, 3}
|
|||
|
|
b := []float64{-1, 2}
|
|||
|
|
x0 := []float64{0.5, -1.25}
|
|||
|
|
xt, err := FromFloat64s(x0, false, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat64s: %v", err)
|
|||
|
|
}
|
|||
|
|
f := func(z *Tensor) (*Tensor, error) {
|
|||
|
|
az, err := FromFloat64s(a, false, 2, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
bz, err := FromFloat64s(b, false, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
halfAz, err := az.Scale(0.5)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
azx, err := halfAz.MatMul(z)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
sum1, err := azx.Add(bz)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
// ½xᵀAx + bᵀx = ((½A)x + b)·x
|
|||
|
|
return sum1.Mul(z)
|
|||
|
|
}
|
|||
|
|
// ((½A)x + b)·x is elementwise; the scalar loss needs the sum.
|
|||
|
|
fScalar := func(z *Tensor) (*Tensor, error) {
|
|||
|
|
p, err := f(z)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return p.Sum()
|
|||
|
|
}
|
|||
|
|
h, err := Hessian(fScalar, xt, HessianOptions{})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Hessian: %v", err)
|
|||
|
|
}
|
|||
|
|
for i := range 2 {
|
|||
|
|
for j := range 2 {
|
|||
|
|
if math.Abs(h.FloatAt(i*2+j)-a[i*2+j]) > 1e-6 {
|
|||
|
|
t.Fatalf("H[%d][%d] = %g, want %g", i, j, h.FloatAt(i*2+j), a[i*2+j])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestHessianRosenbrock pins a nonquadratic landscape against the
|
|||
|
|
// analytic Hessian of the 2-D Rosenbrock function.
|
|||
|
|
func TestHessianRosenbrock(t *testing.T) {
|
|||
|
|
x0 := []float64{-0.5, 1.25}
|
|||
|
|
xt, err := FromFloat64s(x0, false, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat64s: %v", err)
|
|||
|
|
}
|
|||
|
|
f := func(z *Tensor) (*Tensor, error) {
|
|||
|
|
els := []int{0, 1}
|
|||
|
|
x0t, err := z.Slice(0, els[0], els[0]+1)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
x1t, err := z.Slice(0, els[1], els[1]+1)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
x0sq, err := x0t.Pow(2)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
diff, err := x1t.Sub(x0sq)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
term1v, err := diff.Pow(2)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
one, err := FromFloat64s([]float64{1}, false, 1)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
x0m1, err := x0t.Sub(one)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
term2v, err := x0m1.Pow(2)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
term2s, err := term2v.Scale(100)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
total, err := term1v.Add(term2s)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return total.Sum()
|
|||
|
|
}
|
|||
|
|
h, err := Hessian(f, xt, HessianOptions{})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Hessian: %v", err)
|
|||
|
|
}
|
|||
|
|
x, y := x0[0], x0[1]
|
|||
|
|
// Analytic Hessian of f = (y − x²)² + 100(x − 1)².
|
|||
|
|
h00 := 12*x*x - 4*y + 200
|
|||
|
|
h01 := -4 * x
|
|||
|
|
h11 := 2.0
|
|||
|
|
want := [][]float64{{h00, h01}, {h01, h11}}
|
|||
|
|
for i := range 2 {
|
|||
|
|
for j := range 2 {
|
|||
|
|
if math.Abs(h.FloatAt(i*2+j)-want[i][j]) > 1e-4*math.Max(1, math.Abs(want[i][j])) {
|
|||
|
|
t.Fatalf("H[%d][%d] = %g, want %g", i, j, h.FloatAt(i*2+j), want[i][j])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestHessianVectorProduct pins H·v against the dense Hessian.
|
|||
|
|
func TestHessianVectorProduct(t *testing.T) {
|
|||
|
|
a := []float64{4, 1, 1, 3}
|
|||
|
|
x0 := []float64{0.5, -1.25}
|
|||
|
|
xt, err := FromFloat64s(x0, false, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat64s: %v", err)
|
|||
|
|
}
|
|||
|
|
f := func(z *Tensor) (*Tensor, error) {
|
|||
|
|
az, err := FromFloat64s(a, false, 2, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
halfAz, err := az.Scale(0.5)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
azx, err := halfAz.MatMul(z)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
p, err := azx.Mul(z)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return p.Sum()
|
|||
|
|
}
|
|||
|
|
vArr, err := core.FromFloats([]float64{2, -1}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
v := FromArray(vArr, false)
|
|||
|
|
hv, err := HessianVectorProduct(f, xt, v, HessianOptions{})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("HessianVectorProduct: %v", err)
|
|||
|
|
}
|
|||
|
|
// A·v exactly.
|
|||
|
|
want := []float64{4*2 + 1*(-1), 1*2 + 3*(-1)}
|
|||
|
|
for i := range 2 {
|
|||
|
|
if math.Abs(hv.FloatAt(i)-want[i]) > 1e-5 {
|
|||
|
|
t.Fatalf("Hv[%d] = %g, want %g", i, hv.FloatAt(i), want[i])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestHessianRejectsVectorOutput pins the scalar contract.
|
|||
|
|
func TestHessianRejectsVectorOutput(t *testing.T) {
|
|||
|
|
xt, err := FromFloat64s([]float64{1, 2}, false, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat64s: %v", err)
|
|||
|
|
}
|
|||
|
|
f := func(z *Tensor) (*Tensor, error) { return z, nil }
|
|||
|
|
if _, err := Hessian(f, xt, HessianOptions{}); err == nil {
|
|||
|
|
t.Fatal("Hessian accepted a vector output")
|
|||
|
|
}
|
|||
|
|
}
|