Files
tensor/grad/newtoncg_test.go
T
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

198 lines
5.2 KiB
Go
Raw 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"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// TestNewtonCGQuadratic pins the exactly-Newtonian case: a quadratic
// with SPD Hessian converges to the analytic minimiser in a couple of
// steps.
func TestNewtonCGQuadratic(t *testing.T) {
a := []float64{4, 1, 1, 3}
b := []float64{-1, 2}
x0, err := core.FromFloats([]float64{0.5, -1.25}, 2)
if err != nil {
t.Fatalf("FromFloats: %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
}
halfA, err := az.Scale(0.5)
if err != nil {
return nil, err
}
azx, err := halfA.MatMul(z)
if err != nil {
return nil, err
}
lin, err := azx.Add(bz)
if err != nil {
return nil, err
}
prod, err := lin.Mul(z)
if err != nil {
return nil, err
}
return prod.Sum()
}
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{Tolerance: 1e-12})
if err != nil {
t.Fatalf("MinimiseNewtonCG: %v", err)
}
// x* = −A⁻¹b: solve 4x+y = 1, x+3y = −2 so x = 5/11, y = −9/11.
if math.Abs(x.FloatAt(0)-5.0/11) > 1e-9 || math.Abs(x.FloatAt(1)+9.0/11) > 1e-9 {
t.Fatalf("minimiser = (%g, %g), want (5/11, -9/11)", x.FloatAt(0), x.FloatAt(1))
}
// f* = ½x*ᵀAx* + bᵀx* = 253/242 − 23/11 = −253/242.
const want = -253.0 / 242.0
if math.Abs(fv-want) > 1e-10 {
t.Fatalf("value = %.12g, want %.12g", fv, want)
}
}
// TestNewtonCGRosenbrock pins a nonquadratic valley: the classic
// Rosenbrock minimum at (1, 1) from the far side.
func TestNewtonCGRosenbrock(t *testing.T) {
x0, err := core.FromFloats([]float64{-1.5, 2}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
f := func(z *Tensor) (*Tensor, error) {
x0t, err := z.Slice(0, 0, 1)
if err != nil {
return nil, err
}
x1t, err := z.Slice(0, 1, 2)
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
}
term1, 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
}
term2, err := x0m1.Pow(2)
if err != nil {
return nil, err
}
term2s, err := term2.Scale(100)
if err != nil {
return nil, err
}
total, err := term1.Add(term2s)
if err != nil {
return nil, err
}
return total.Sum()
}
x, _, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{Tolerance: 1e-7, MaxIterations: 200})
if err != nil {
t.Fatalf("MinimiseNewtonCG: %v", err)
}
if math.Abs(x.FloatAt(0)-1) > 1e-4 || math.Abs(x.FloatAt(1)-1) > 1e-4 {
t.Fatalf("minimiser = (%.6f, %.6f), want (1, 1)", x.FloatAt(0), x.FloatAt(1))
}
}
// TestNewtonCGNegativeCurvature pins the fallback: a double well
// whose start sits in the concave region between the minima. The CG
// must take its steepest-descent fallback there and still land in a
// well.
func TestNewtonCGNegativeCurvature(t *testing.T) {
x0, err := core.FromFloats([]float64{0.1, 0.2}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
// f = Σ(x⁴ − x²): Hessian 12x² − 2 is negative for |x| < 1/√6,
// so the start is concave; the wells sit at ±1/√2 per coordinate.
f := func(z *Tensor) (*Tensor, error) {
q, err := z.Pow(4)
if err != nil {
return nil, err
}
sq, err := z.Abs2()
if err != nil {
return nil, err
}
d, err := q.Sub(sq)
if err != nil {
return nil, err
}
return d.Sum()
}
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{Tolerance: 1e-9})
if err != nil {
t.Fatalf("MinimiseNewtonCG: %v", err)
}
const well = 1.0 / math.Sqrt2
for i := range 2 {
if math.Abs(math.Abs(x.FloatAt(i))-well) > 1e-6 {
t.Fatalf("coordinate %d = %g, want magnitude %g", i, x.FloatAt(i), well)
}
}
// f at a well: Σ(1/4 − 1/2) = −1/2.
if math.Abs(fv+0.5) > 1e-9 {
t.Fatalf("value = %.12g, want -0.5", fv)
}
}
// TestNewtonCGScalarInput pins the n = 1 path and the error contract.
func TestNewtonCGScalarInput(t *testing.T) {
x0, err := core.FromFloats([]float64{3}, 1)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
f := func(z *Tensor) (*Tensor, error) {
sq, err := z.Pow(2)
if err != nil {
return nil, err
}
four, err := sq.Scale(4)
if err != nil {
return nil, err
}
return four.Sum()
}
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{})
if err != nil {
t.Fatalf("MinimiseNewtonCG: %v", err)
}
if math.Abs(x.FloatAt(0)) > 1e-7 || math.Abs(fv) > 1e-12 {
t.Fatalf("minimiser = %g, value = %g", x.FloatAt(0), fv)
}
// A slice of the point itself: the objective is disconnected from
// the minimised point, so the run exhausts its iterations on a
// constant value and errors loudly.
c, _ := core.FromFloats([]float64{1}, 1)
if _, _, err := MinimiseNewtonCG(func(z *Tensor) (*Tensor, error) { return z.Slice(0, 0, 1) }, c, NewtonCGOptions{}); err == nil {
t.Fatal("a slice of the point itself minimised without error")
}
}