Files

194 lines
4.5 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// 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")
}
}