Files
tensor/grad/hessian_test.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

194 lines
4.5 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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")
}
}