176 lines
5.9 KiB
Go
176 lines
5.9 KiB
Go
// 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"
|
||
"testing"
|
||
)
|
||
|
||
// TestFindRoot checks Brent's method on classic brackets, including
|
||
// the polynomial Whittaker and Robinson used to sell the method.
|
||
func TestFindRoot(t *testing.T) {
|
||
// x³ − 2x − 5 has its root near 2.0945514815423265 in [2, 3].
|
||
root, err := FindRoot(func(x float64) float64 { return x*x*x - 2*x - 5 }, 2, 3, 0)
|
||
if err != nil {
|
||
t.Fatalf("FindRoot: %v", err)
|
||
}
|
||
want := 2.0945514815423265
|
||
if math.Abs(root-want) > 1e-12 {
|
||
t.Fatalf("root = %.16g, want %.16g", root, want)
|
||
}
|
||
// sin(x) on [3, 4] brackets π.
|
||
pi, err := FindRoot(math.Sin, 3, 4, 0)
|
||
if err != nil {
|
||
t.Fatalf("FindRoot sin: %v", err)
|
||
}
|
||
if math.Abs(pi-math.Pi) > 1e-12 {
|
||
t.Fatalf("sin root = %.16g, want π", pi)
|
||
}
|
||
// A root exactly at an endpoint is returned.
|
||
zero, err := FindRoot(func(x float64) float64 { return x - 1 }, 1, 4, 0)
|
||
if err != nil {
|
||
t.Fatalf("FindRoot endpoint: %v", err)
|
||
}
|
||
if zero != 1 {
|
||
t.Fatalf("endpoint root = %v, want 1", zero)
|
||
}
|
||
// A bracket that does not change sign is refused.
|
||
if _, err := FindRoot(func(x float64) float64 { return x*x - 1 }, -2, 2, 0); err == nil {
|
||
t.Fatal("expected an error for a bracket without a sign change")
|
||
}
|
||
// NaN at the bracket is refused.
|
||
if _, err := FindRoot(func(x float64) float64 { return math.NaN() }, -1, 1, 0); err == nil {
|
||
t.Fatal("expected an error for a NaN bracket")
|
||
}
|
||
}
|
||
|
||
// TestFindRootNewton checks the Newton iteration on the square root
|
||
// of two and pins the two failure contracts.
|
||
func TestFindRootNewton(t *testing.T) {
|
||
root, err := FindRootNewton(func(x float64) float64 { return x*x - 2 },
|
||
func(x float64) float64 { return 2 * x }, 1, 1e-14, 50)
|
||
if err != nil {
|
||
t.Fatalf("FindRootNewton: %v", err)
|
||
}
|
||
if math.Abs(root-math.Sqrt2) > 1e-12 {
|
||
t.Fatalf("root = %.16g, want √2", root)
|
||
}
|
||
// A vanishing derivative is refused rather than dividing by zero.
|
||
if _, err := FindRootNewton(func(x float64) float64 { return x*x - 2 },
|
||
func(x float64) float64 { return 0 }, 1, 0, 0); err == nil {
|
||
t.Fatal("expected an error for a vanishing derivative")
|
||
}
|
||
// An exhausted budget is an error.
|
||
if _, err := FindRootNewton(func(x float64) float64 { return x*x*x - 2 },
|
||
func(x float64) float64 { return 3 * x * x }, 1000, 1e-14, 2); err == nil {
|
||
t.Fatal("expected an error for an exhausted iteration budget")
|
||
}
|
||
}
|
||
|
||
// TestMinimise checks the simplex method on a convex sphere and on
|
||
// the Rosenbrock valley, the standard trap for naive descent.
|
||
func TestMinimise(t *testing.T) {
|
||
// Sphere: the minimum is the origin.
|
||
sphere := func(p *core.Array) (float64, error) {
|
||
sum := 0.0
|
||
for i := range p.Len() {
|
||
sum += p.FloatAt(i) * p.FloatAt(i)
|
||
}
|
||
return sum, nil
|
||
}
|
||
start, err := core.FromFloats([]float64{2, -3, 1.5}, 3)
|
||
if err != nil {
|
||
t.Fatalf("FromFloats: %v", err)
|
||
}
|
||
point, value, err := Minimise(sphere, start, MinimiseOptions{})
|
||
if err != nil {
|
||
t.Fatalf("Minimise sphere: %v", err)
|
||
}
|
||
// The default stopping rule bounds the function spread at 1e-10,
|
||
// which on a quadratic puts the point at roughly 1e-5 per
|
||
// coordinate; the contract is the spread, not exact zero.
|
||
if value > 1e-8 {
|
||
t.Fatalf("sphere minimum value = %v, want <= 1e-8", value)
|
||
}
|
||
for i := range point.Len() {
|
||
if math.Abs(point.FloatAt(i)) > 1e-3 {
|
||
t.Fatalf("sphere minimiser[%d] = %v, want ≈ 0", i, point.FloatAt(i))
|
||
}
|
||
}
|
||
// A tighter tolerance buys a tighter minimum on the same objective.
|
||
point, value, err = Minimise(sphere, start, MinimiseOptions{Tolerance: 1e-22})
|
||
if err != nil {
|
||
t.Fatalf("Minimise sphere tight: %v", err)
|
||
}
|
||
if value > 1e-20 {
|
||
t.Fatalf("tight sphere minimum value = %v, want <= 1e-20", value)
|
||
}
|
||
|
||
// Rosenbrock in two dimensions: minimum at (1, 1).
|
||
rosenbrock := func(p *core.Array) (float64, error) {
|
||
x, y := p.FloatAt(0), p.FloatAt(1)
|
||
return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil
|
||
}
|
||
rs, _ := core.FromFloats([]float64{-1.2, 1}, 2)
|
||
point, value, err = Minimise(rosenbrock, rs, MinimiseOptions{})
|
||
if err != nil {
|
||
t.Fatalf("Minimise Rosenbrock: %v", err)
|
||
}
|
||
if math.Abs(point.FloatAt(0)-1) > 1e-4 || math.Abs(point.FloatAt(1)-1) > 1e-4 {
|
||
t.Fatalf("Rosenbrock minimiser = (%.8g, %.8g), want (1, 1)",
|
||
point.FloatAt(0), point.FloatAt(1))
|
||
}
|
||
if value > 1e-8 {
|
||
t.Fatalf("Rosenbrock minimum value = %v, want ≈ 0", value)
|
||
}
|
||
|
||
// An objective that errors must surface the error.
|
||
failing := func(p *core.Array) (float64, error) {
|
||
if p.FloatAt(0) > 0.5 {
|
||
return 0, base.Errf("objective exploded")
|
||
}
|
||
return p.FloatAt(0) * p.FloatAt(0), nil
|
||
}
|
||
fs, _ := core.FromFloats([]float64{0}, 1)
|
||
if _, _, err := Minimise(failing, fs, MinimiseOptions{}); err == nil {
|
||
t.Fatal("expected the objective's error to propagate")
|
||
}
|
||
|
||
// An empty starting point is refused.
|
||
empty, _ := core.FromFloats(nil, 0)
|
||
if _, _, err := Minimise(sphere, empty, MinimiseOptions{}); err == nil {
|
||
t.Fatal("expected an error for an empty starting point")
|
||
}
|
||
if _, _, err := Minimise(sphere, mustComplexPoint(t), MinimiseOptions{}); err == nil {
|
||
t.Fatal("expected an error for a complex starting point")
|
||
}
|
||
}
|
||
|
||
func mustComplexPoint(t *testing.T) *core.Array {
|
||
t.Helper()
|
||
a, err := core.FromComplexes([]complex128{1}, 1)
|
||
if err != nil {
|
||
t.Fatalf("FromComplexes: %v", err)
|
||
}
|
||
return a
|
||
}
|
||
|
||
// TestFindRootNewtonLargeScale pins the relative step tolerance: for a
|
||
// root at 1e8 the floating-point step granularity (eps·|x| ≈ 2e-8)
|
||
// exceeds an absolute 1e-12 tolerance, which used to report
|
||
// non-convergence from an accurate answer.
|
||
func TestFindRootNewtonLargeScale(t *testing.T) {
|
||
root, err := FindRootNewton(func(x float64) float64 { return x*x - 1e16 },
|
||
func(x float64) float64 { return 2 * x }, 3e8, 1e-12, 100)
|
||
if err != nil {
|
||
t.Fatalf("FindRootNewton: %v", err)
|
||
}
|
||
if math.Abs(root-1e8) > 1e-4 {
|
||
t.Fatalf("root = %.12g, want 1e8", root)
|
||
}
|
||
}
|