// Copyright (c) 2026 Petr Balvín (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") } }