Files

176 lines
5.9 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// 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)
}
}