feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+188
View File
@@ -0,0 +1,188 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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))
}
}
}