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")
|
||
}
|
||
}
|