198 lines
5.2 KiB
Go
198 lines
5.2 KiB
Go
// 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")
|
||
}
|
||
}
|