Files
tensor/optim/optimise_test.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

176 lines
5.9 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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)
}
}