195 lines
6.6 KiB
Go
195 lines
6.6 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
|
|
// SPDX-License-Identifier: MIT
|
|||
|
|
|
|||
|
|
package optim
|
|||
|
|
|
|||
|
|
import (
|
|||
|
|
"go/ast"
|
|||
|
|
"go/parser"
|
|||
|
|
"go/token"
|
|||
|
|
"math"
|
|||
|
|
"os"
|
|||
|
|
"strings"
|
|||
|
|
"testing"
|
|||
|
|
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// Fallback and diagnosis pins: the singular-Jacobian fallback of
|
|||
|
|
// FindRootSystem computing its direction from the factored matrix, a
|
|||
|
|
// NaN residual dying as a bogus damping diagnosis in
|
|||
|
|
// LevenbergMarquardt, the duplicate ArrayFromFloatsSafe export
|
|||
|
|
// shadowed by the facade's linalg copy, the doubled "Minimise:" error
|
|||
|
|
// prefix and the dead wrapVector helper.
|
|||
|
|
|
|||
|
|
// TestFindRootSystemSingularFallbackDescends runs the case:
|
|||
|
|
// r = [x−1, (x−1)−3(x−1)²] has a rank-one Jacobian with a zero second
|
|||
|
|
// column at every point, so every Newton system is singular and the
|
|||
|
|
// steepest-descent fallback carries the whole iteration. The true
|
|||
|
|
// −Jᵀr at the start (2, 0) is (−11, 0), straight downhill to the root
|
|||
|
|
// (1, 0); a direction formed from the factored matrix points the other
|
|||
|
|
// way and the run used to burn its budget with a frozen iterate.
|
|||
|
|
func TestFindRootSystemSingularFallbackDescends(t *testing.T) {
|
|||
|
|
residual := func(x *core.Array) (*core.Array, error) {
|
|||
|
|
d := x.FloatAt(0) - 1
|
|||
|
|
return mustFloats(t, []float64{d, d - 3*d*d}), nil
|
|||
|
|
}
|
|||
|
|
x, res, err := FindRootSystem(residual, mustFloats(t, []float64{2, 0}), RootSystemOptions{})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FindRootSystem: %v", err)
|
|||
|
|
}
|
|||
|
|
if res > 1e-10 {
|
|||
|
|
t.Fatalf("residual %g, want ≤ 1e-10", res)
|
|||
|
|
}
|
|||
|
|
if math.Abs(x.FloatAt(0)-1) > 1e-9 || math.Abs(x.FloatAt(1)) > 1e-9 {
|
|||
|
|
t.Fatalf("solution = (%.12g, %.12g), want (1, 0)", x.FloatAt(0), x.FloatAt(1))
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestSteepestDescentStepDescent checks the extracted fallback at unit
|
|||
|
|
// level: the direction it writes is a descent direction of ‖r‖², which
|
|||
|
|
// the linearised model makes exact, and the same formula applied to the
|
|||
|
|
// matrix base.Factor has consumed is not one, which is the defect the
|
|||
|
|
// pristine copy removes.
|
|||
|
|
func TestSteepestDescentStepDescent(t *testing.T) {
|
|||
|
|
// The Jacobian at (2, 0): rank one, columns [1, −5]ᵀ
|
|||
|
|
// and zero; r = (1, −2).
|
|||
|
|
jac := [][]float64{{1, 0}, {-5, 0}}
|
|||
|
|
r := []float64{1, -2}
|
|||
|
|
step := make([]float64, 2)
|
|||
|
|
steepestDescentStep(jac, r, step)
|
|||
|
|
|
|||
|
|
// The directional derivative of ‖r + α·J·step‖² at α = 0 is
|
|||
|
|
// 2·rᵀ(J·step): negative means the direction descends the
|
|||
|
|
// residual norm, and one small α then really lowers it.
|
|||
|
|
apply := func(dir []float64) []float64 {
|
|||
|
|
out := make([]float64, len(r))
|
|||
|
|
for k := range jac {
|
|||
|
|
s := 0.0
|
|||
|
|
for i := range dir {
|
|||
|
|
s += jac[k][i] * dir[i]
|
|||
|
|
}
|
|||
|
|
out[k] = s
|
|||
|
|
}
|
|||
|
|
return out
|
|||
|
|
}
|
|||
|
|
slope := 0.0
|
|||
|
|
for k, jv := range apply(step) {
|
|||
|
|
slope += r[k] * jv
|
|||
|
|
}
|
|||
|
|
if slope >= 0 {
|
|||
|
|
t.Fatalf("fallback direction has residual slope %g, want < 0", slope)
|
|||
|
|
}
|
|||
|
|
const alpha = 0.05
|
|||
|
|
before, after := 0.0, 0.0
|
|||
|
|
jstep := apply(step)
|
|||
|
|
for k := range r {
|
|||
|
|
before += r[k] * r[k]
|
|||
|
|
v := r[k] + alpha*jstep[k]
|
|||
|
|
after += v * v
|
|||
|
|
}
|
|||
|
|
if after >= before {
|
|||
|
|
t.Fatalf("linearised residual norm rose from %g to %g, want a fall", before, after)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// The factored matrix must not be mistaken for the Jacobian:
|
|||
|
|
// Factor pivots the rows and stores the multipliers on the
|
|||
|
|
// subdiagonal in place, and the direction from that matrix points
|
|||
|
|
// uphill here.
|
|||
|
|
factored := [][]float64{{1, 0}, {-5, 0}}
|
|||
|
|
base.Factor(factored)
|
|||
|
|
bad := make([]float64, 2)
|
|||
|
|
for i := range bad {
|
|||
|
|
s := 0.0
|
|||
|
|
for k := range r {
|
|||
|
|
s += factored[k][i] * r[k]
|
|||
|
|
}
|
|||
|
|
bad[i] = -s
|
|||
|
|
}
|
|||
|
|
badSlope := 0.0
|
|||
|
|
for k, jv := range apply(bad) {
|
|||
|
|
badSlope += r[k] * jv
|
|||
|
|
}
|
|||
|
|
if badSlope <= 0 {
|
|||
|
|
t.Fatalf("factored-matrix direction has slope %g, want > 0 (the defect the copy fixes)", badSlope)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestLevenbergMarquardtNonFiniteResidual pins the diagnosis: a
|
|||
|
|
// residual that returns NaN or Inf is reported as the non-finite value
|
|||
|
|
// it is, with its index, and never misread as a damping collapse.
|
|||
|
|
func TestLevenbergMarquardtNonFiniteResidual(t *testing.T) {
|
|||
|
|
nan := func(*core.Array) (*core.Array, error) {
|
|||
|
|
return mustFloats(t, []float64{math.NaN(), math.NaN()}), nil
|
|||
|
|
}
|
|||
|
|
_, _, err := LevenbergMarquardt(nan, mustFloats(t, []float64{0, 0}), LMOptions{})
|
|||
|
|
if err == nil {
|
|||
|
|
t.Fatal("NaN residual: want an error")
|
|||
|
|
}
|
|||
|
|
if msg := err.Error(); !strings.Contains(msg, "the residual returned the non-finite value NaN at 0") {
|
|||
|
|
t.Fatalf("NaN residual misdiagnosed: %v", err)
|
|||
|
|
} else if strings.Contains(msg, "damping") {
|
|||
|
|
t.Fatalf("NaN residual reported as a damping collapse: %v", err)
|
|||
|
|
}
|
|||
|
|
inf := func(*core.Array) (*core.Array, error) {
|
|||
|
|
return mustFloats(t, []float64{1, math.Inf(1)}), nil
|
|||
|
|
}
|
|||
|
|
_, _, err = LevenbergMarquardt(inf, mustFloats(t, []float64{0, 0}), LMOptions{})
|
|||
|
|
if err == nil {
|
|||
|
|
t.Fatal("Inf residual: want an error")
|
|||
|
|
}
|
|||
|
|
if msg := err.Error(); !strings.Contains(msg, "the residual returned the non-finite value +Inf at 1") {
|
|||
|
|
t.Fatalf("Inf residual misdiagnosed: %v", err)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestMinimiseNonFiniteSinglePrefix pins the exact wording: the
|
|||
|
|
// non-finite objective carries one "Minimise:" prefix, because the
|
|||
|
|
// inner error is bare and each call site of eval wraps once.
|
|||
|
|
func TestMinimiseNonFiniteSinglePrefix(t *testing.T) {
|
|||
|
|
f := func(*core.Array) (float64, error) { return math.NaN(), nil }
|
|||
|
|
_, _, err := Minimise(f, mustFloats(t, []float64{0, 0}), MinimiseOptions{})
|
|||
|
|
if err == nil {
|
|||
|
|
t.Fatal("NaN objective: want an error")
|
|||
|
|
}
|
|||
|
|
want := "tensor: Minimise: f returned the non-finite value NaN at [0 0]"
|
|||
|
|
if err.Error() != want {
|
|||
|
|
t.Fatalf("error = %q, want %q", err.Error(), want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestOptimSourcesDeclareNoMovedOrDeadHelpers guards the public
|
|||
|
|
// surface: ArrayFromFloatsSafe lives in linalg alone (the facade
|
|||
|
|
// forwards the name there), and the dead wrapVector wrapper is gone.
|
|||
|
|
func TestOptimSourcesDeclareNoMovedOrDeadHelpers(t *testing.T) {
|
|||
|
|
banned := map[string]string{
|
|||
|
|
"ArrayFromFloatsSafe": "a bit-identical duplicate of linalg.ArrayFromFloatsSafe, which the facade already exports",
|
|||
|
|
"wrapVector": "a dead helper with no caller in the package",
|
|||
|
|
}
|
|||
|
|
entries, err := os.ReadDir(".")
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("ReadDir: %v", err)
|
|||
|
|
}
|
|||
|
|
fset := token.NewFileSet()
|
|||
|
|
for _, e := range entries {
|
|||
|
|
name := e.Name()
|
|||
|
|
if e.IsDir() || !strings.HasSuffix(name, ".go") || strings.HasSuffix(name, "_test.go") {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
file, perr := parser.ParseFile(fset, name, nil, 0)
|
|||
|
|
if perr != nil {
|
|||
|
|
t.Fatalf("parse %s: %v", name, perr)
|
|||
|
|
}
|
|||
|
|
for _, decl := range file.Decls {
|
|||
|
|
fn, ok := decl.(*ast.FuncDecl)
|
|||
|
|
if !ok || fn.Recv != nil {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
if why, bad := banned[fn.Name.Name]; bad {
|
|||
|
|
t.Errorf("%s declares %s again: %s", name, fn.Name.Name, why)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|