feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -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))
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user