feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,193 @@
|
||||
// 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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user