// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package optim import ( "math" "sourcedock.dev/petrbalvin/tensor/internal/base" "sourcedock.dev/petrbalvin/tensor/internal/core" "sourcedock.dev/petrbalvin/tensor/linalg" "testing" ) // TestFindRootSystemLinear checks the Newton step on a linear system: // one iteration must land on the dense solve's answer. func TestFindRootSystemLinear(t *testing.T) { vals := []float64{ 3, 1, 1, 1, 4, 2, 0, 1, 5, } b := []float64{1, 2, 3} residual := func(x *core.Array) (*core.Array, error) { out := core.New(core.Float, 3) for i := range 3 { s := 0.0 for j := range 3 { s += vals[i*3+j] * x.FloatAt(j) } out.RawFloats()[i] = s - b[i] } return out, nil } x, res, err := FindRootSystem(residual, mustFloats(t, []float64{0, 0, 0}), RootSystemOptions{}) if err != nil { t.Fatalf("FindRootSystem: %v", err) } if res > 1e-12 { t.Fatalf("residual %g, want ≤ 1e-12", res) } ref, err := linalg.Solve(mustFloats(t, vals, 3, 3), mustFloats(t, b)) if err != nil { t.Fatalf("Solve: %v", err) } for i := range 3 { if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-12 { t.Fatalf("x[%d] = %.16g, want %.16g", i, x.FloatAt(i), ref.FloatAt(i)) } } } // TestFindRootSystemCircleLine solves x² + y² = 4 with x = y from the // interior: the root is x = y = √2. func TestFindRootSystemCircleLine(t *testing.T) { residual := func(x *core.Array) (*core.Array, error) { cx, cy := x.FloatAt(0), x.FloatAt(1) return mustFloats(t, []float64{cx*cx + cy*cy - 4, cx - cy}), nil } x, res, err := FindRootSystem(residual, mustFloats(t, []float64{0.5, 0.5}), RootSystemOptions{}) if err != nil { t.Fatalf("FindRootSystem: %v", err) } if res > 1e-10 { t.Fatalf("residual %g, want ≤ 1e-10", res) } root := math.Sqrt2 if math.Abs(x.FloatAt(0)-root) > 1e-10 || math.Abs(x.FloatAt(1)-root) > 1e-10 { t.Fatalf("solution = (%.12g, %.12g), want (√2, √2)", x.FloatAt(0), x.FloatAt(1)) } } // TestFindRootSystemCircleHyperbola solves x² + y² = 5 with xy = 2, // whose two roots are (1, 2) and (2, 1); from an interior start the // iteration must land on one of them with a vanishing residual. The // damped step is exercised by the curved crossing of the two curves. func TestFindRootSystemCircleHyperbola(t *testing.T) { residual := func(x *core.Array) (*core.Array, error) { cx, cy := x.FloatAt(0), x.FloatAt(1) return mustFloats(t, []float64{cx*cx + cy*cy - 5, cx*cy - 2}), nil } x, res, err := FindRootSystem(residual, mustFloats(t, []float64{0.4, 2.2}), RootSystemOptions{}) if err != nil { t.Fatalf("FindRootSystem: %v", err) } if res > 1e-10 { t.Fatalf("residual %g, want ≤ 1e-10", res) } root1 := []float64{1, 2} root2 := []float64{2, 1} hit := func(root []float64) bool { return math.Abs(x.FloatAt(0)-root[0]) < 1e-9 && math.Abs(x.FloatAt(1)-root[1]) < 1e-9 } if !hit(root1) && !hit(root2) { t.Fatalf("solution = (%.12g, %.12g), want (1, 2) or (2, 1)", x.FloatAt(0), x.FloatAt(1)) } } // TestFindRootSystemTrig mixes transcendental equations: cos x = y and // sin x = y meet where x = π/4. func TestFindRootSystemTrig(t *testing.T) { residual := func(x *core.Array) (*core.Array, error) { cx, cy := x.FloatAt(0), x.FloatAt(1) return mustFloats(t, []float64{math.Cos(cx) - cy, math.Sin(cx) - cy}), nil } x, res, err := FindRootSystem(residual, mustFloats(t, []float64{0.3, 0.1}), 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)-math.Pi/4) > 1e-9 || math.Abs(x.FloatAt(1)-math.Sqrt2/2) > 1e-9 { t.Fatalf("solution = (%.12g, %.12g), want (π/4, 1/√2)", x.FloatAt(0), x.FloatAt(1)) } } func TestFindRootSystemErrors(t *testing.T) { // The origin is a genuine root of (x², y²) and must be answered. x, res, err := FindRootSystem(func(x *core.Array) (*core.Array, error) { return mustFloats(t, []float64{x.FloatAt(0) * x.FloatAt(0), x.FloatAt(1) * x.FloatAt(1)}), nil }, mustFloats(t, []float64{0, 0}), RootSystemOptions{}) if err != nil || res > 0 || x.FloatAt(0) != 0 { t.Fatalf("root at the origin: x = %v, res = %v, %v", x.RawFloats(), res, err) } // Duplicated equations have a rank-one Jacobian everywhere: the // Newton system is singular while the residual is nonzero. rankOne := func(x *core.Array) (*core.Array, error) { cx := x.FloatAt(0) return mustFloats(t, []float64{cx*cx - 1, cx*cx - 1}), nil } if _, _, err := FindRootSystem(rankOne, mustFloats(t, []float64{2}), RootSystemOptions{}); err == nil { t.Fatal("hopeless rank-one system: want an error") } wrongShape := func(x *core.Array) (*core.Array, error) { return mustFloats(t, []float64{1, 1, 1}), nil } if _, _, err := FindRootSystem(wrongShape, mustFloats(t, []float64{1, 1}), RootSystemOptions{}); err == nil { t.Fatal("wrong residual shape: want an error") } fails := func(*core.Array) (*core.Array, error) { return nil, base.Errf("residual failed") } if _, _, err := FindRootSystem(fails, mustFloats(t, []float64{1, 1}), RootSystemOptions{}); err == nil { t.Fatal("residual error: want an error") } // An exhausted budget must be reported, not approximated away. circleResidual := func(x *core.Array) (*core.Array, error) { cx, cy := x.FloatAt(0), x.FloatAt(1) return mustFloats(t, []float64{cx*cx + cy*cy - 4, cx - cy}), nil } if _, _, err := FindRootSystem(circleResidual, mustFloats(t, []float64{1, 1}), RootSystemOptions{MaxIterations: 1, Tolerance: 1e-20}); err == nil { t.Fatal("exhausted budget: want an error") } if _, _, err := FindRootSystem(circleResidual, mustFloats(t, []float64{}), RootSystemOptions{}); err == nil { t.Fatal("empty guess: want an error") } } // TestFindRootSystemDtypes pins the starting-vector promotion: Int and // Float32 starts must behave exactly like Float64 ones. The old // RawFloats() fast path handed the iteration a nil or zero slice for // those dtypes. func TestFindRootSystemDtypes(t *testing.T) { residual := func(x *core.Array) (*core.Array, error) { cx, cy := x.FloatAt(0), x.FloatAt(1) return mustFloats(t, []float64{cx*cx + cy*cy - 4, cx - cy}), nil } ints, err := core.FromInts([]int64{1, 1}, 2) if err != nil { t.Fatalf("FromInts: %v", err) } thirtyTwo, err := core.FromFloat32s([]float32{0.5, 0.5}, 2) if err != nil { t.Fatalf("FromFloat32s: %v", err) } for name, x0 := range map[string]*core.Array{"int": ints, "float32": thirtyTwo} { x, res, err := FindRootSystem(residual, x0, RootSystemOptions{}) if err != nil { t.Fatalf("FindRootSystem(%s start): %v", name, err) } if res > 1e-10 { t.Fatalf("FindRootSystem(%s start): residual %g, want ≤ 1e-10", name, res) } if math.Abs(x.FloatAt(0)-math.Sqrt2) > 1e-8 || math.Abs(x.FloatAt(1)-math.Sqrt2) > 1e-8 { t.Fatalf("FindRootSystem(%s start): solution = (%.12g, %.12g), want (√2, √2)", name, x.FloatAt(0), x.FloatAt(1)) } } }