204 lines
7.3 KiB
Go
204 lines
7.3 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package optim
|
||
|
||
import (
|
||
"math"
|
||
"strings"
|
||
"testing"
|
||
|
||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
)
|
||
|
||
// TestFindRootBrentTranscendentalRoots pins the documented contract on
|
||
// known transcendental roots: the answer sits inside the bracket and
|
||
// carries |f| at or below the default tolerance of 1e-10.
|
||
func TestFindRootBrentTranscendentalRoots(t *testing.T) {
|
||
cases := []struct {
|
||
name string
|
||
f func(float64) float64
|
||
a, b float64
|
||
root float64
|
||
}{
|
||
{"sin", math.Sin, 1, 4, math.Pi},
|
||
{"cos x = x", func(x float64) float64 { return math.Cos(x) - x }, 0, 1, 0.7390851332151607},
|
||
{"exp(x) = 10", func(x float64) float64 { return math.Exp(x) - 10 }, 0, 3, math.Log(10)},
|
||
{"x³ = 2x + 5", func(x float64) float64 { return x*x*x - 2*x - 5 }, 1, 3, 2.0945514815423265},
|
||
}
|
||
for _, tc := range cases {
|
||
root, err := FindRootBrent(tc.f, tc.a, tc.b, BrentOptions{})
|
||
if err != nil {
|
||
t.Fatalf("%s: FindRootBrent: %v", tc.name, err)
|
||
}
|
||
if root <= tc.a || root >= tc.b {
|
||
t.Fatalf("%s: root %g outside the open bracket (%g, %g)", tc.name, root, tc.a, tc.b)
|
||
}
|
||
if math.Abs(root-tc.root) > 1e-9 {
|
||
t.Fatalf("%s: root = %.16g, want %.16g", tc.name, root, tc.root)
|
||
}
|
||
if fx := math.Abs(tc.f(root)); fx > 1e-10 {
|
||
t.Fatalf("%s: |f(root)| = %g, want ≤ 1e-10", tc.name, fx)
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestFindRootBrentRefusesSameSignBracket pins the bracket refusal: a
|
||
// pair of endpoint values of one sign carries no guarantee, and the
|
||
// solver answers with the package's error style instead of a guess.
|
||
func TestFindRootBrentRefusesSameSignBracket(t *testing.T) {
|
||
f := func(x float64) float64 { return x*x - 1 }
|
||
_, err := FindRootBrent(f, 2, 3, BrentOptions{})
|
||
if err == nil || !strings.Contains(err.Error(), "does not change sign") {
|
||
t.Fatalf("err = %v, want the sign-change refusal", err)
|
||
}
|
||
if strings.Count(err.Error(), "FindRootBrent") != 1 {
|
||
t.Fatalf("err = %v, want exactly one entry-point prefix", err)
|
||
}
|
||
// The reversed orientation refuses for the same reason.
|
||
if _, err := FindRootBrent(f, 3, 2, BrentOptions{}); err == nil {
|
||
t.Fatal("reversed same-sign bracket: want an error")
|
||
}
|
||
// A non-finite endpoint value is refused before anything moves.
|
||
nan := func(float64) float64 { return math.NaN() }
|
||
if _, err := FindRootBrent(nan, 1, 2, BrentOptions{}); err == nil || !strings.Contains(err.Error(), "finite values") {
|
||
t.Fatalf("err = %v, want the finite-value refusal", err)
|
||
}
|
||
}
|
||
|
||
// TestFindRootBrentAnswersEndpointRoot pins that a root sitting
|
||
// exactly on a bracket end is answered as it stands, in either
|
||
// orientation, and that an exact zero hit mid-run comes straight back.
|
||
func TestFindRootBrentAnswersEndpointRoot(t *testing.T) {
|
||
f := func(x float64) float64 { return x - 2 }
|
||
for _, ends := range [][2]float64{{2, 5}, {5, 2}} {
|
||
root, err := FindRootBrent(f, ends[0], ends[1], BrentOptions{})
|
||
if err != nil {
|
||
t.Fatalf("FindRootBrent([%g, %g]): %v", ends[0], ends[1], err)
|
||
}
|
||
if root != 2 {
|
||
t.Fatalf("FindRootBrent([%g, %g]) = %g, want 2", ends[0], ends[1], root)
|
||
}
|
||
}
|
||
// A bisection whose midpoint is the root exactly.
|
||
root, err := FindRootBrent(f, 0, 4, BrentOptions{})
|
||
if err != nil || root != 2 {
|
||
t.Fatalf("FindRootBrent([0, 4]) = (%g, %v), want (2, nil)", root, err)
|
||
}
|
||
}
|
||
|
||
// TestFindRootBrentSteepSlopeRefuses pins the honest failure at the
|
||
// resolution limit: a sign change carried by a jump has no point where
|
||
// |f| can sit below the tolerance, and once the bracket closes around
|
||
// the jump the solver reports the collapse instead of a point that
|
||
// misses the documented contract.
|
||
func TestFindRootBrentSteepSlopeRefuses(t *testing.T) {
|
||
f := func(x float64) float64 {
|
||
if x < 0.5 {
|
||
return -1
|
||
}
|
||
return 1e6
|
||
}
|
||
_, err := FindRootBrent(f, 0, 1, BrentOptions{})
|
||
if err == nil || !strings.Contains(err.Error(), "collapsed to machine width") {
|
||
t.Fatalf("err = %v, want the collapsed-bracket refusal", err)
|
||
}
|
||
}
|
||
|
||
// TestFindRootBrentStaysInsideAndBeatsBisection pins two properties on
|
||
// the classic cubic x³ = 2x + 5: every evaluation lands inside the
|
||
// initial bracket, and the superlinear convergence shows in the
|
||
// evaluation count against pure bisection under the same residual
|
||
// stop.
|
||
func TestFindRootBrentStaysInsideAndBeatsBisection(t *testing.T) {
|
||
const a, b, tol = 1.0, 3.0, 1e-10
|
||
fc := func(x float64) float64 { return x*x*x - 2*x - 5 }
|
||
var evals []float64
|
||
root, err := FindRootBrent(func(x float64) float64 {
|
||
evals = append(evals, x)
|
||
return fc(x)
|
||
}, a, b, BrentOptions{Tolerance: tol})
|
||
if err != nil {
|
||
t.Fatalf("FindRootBrent: %v", err)
|
||
}
|
||
if math.Abs(fc(root)) > tol {
|
||
t.Fatalf("|f(root)| = %g, want ≤ %g", math.Abs(fc(root)), tol)
|
||
}
|
||
for _, x := range evals {
|
||
if x < a || x > b {
|
||
t.Fatalf("evaluated f at %g, outside the bracket [%g, %g]", x, a, b)
|
||
}
|
||
}
|
||
// The documented bound: two endpoint evaluations plus one per
|
||
// iteration of the default budget.
|
||
if len(evals) > 100+2 {
|
||
t.Fatalf("%d evaluations, want at most MaxIterations + 2", len(evals))
|
||
}
|
||
// Pure bisection under the same residual stop, counted the same
|
||
// way. The cubic's slope near the root makes |f| ≤ 1e-10 need a
|
||
// bracket about 1e-11 wide, which bisection reaches only after a
|
||
// long halving chain.
|
||
bisectionEvals := 0
|
||
bisect := func(x float64) float64 {
|
||
bisectionEvals++
|
||
return fc(x)
|
||
}
|
||
xl, xr := a, b
|
||
fxl := bisect(xl)
|
||
for {
|
||
mid := 0.5 * (xl + xr)
|
||
fm := bisect(mid)
|
||
if math.Abs(fm) <= tol || math.Abs(xr-xl) <= 2*base.EpsF*math.Abs(mid) {
|
||
break
|
||
}
|
||
if (fxl < 0) == (fm < 0) {
|
||
xl, fxl = mid, fm
|
||
} else {
|
||
xr = mid
|
||
}
|
||
}
|
||
if len(evals)*2 >= bisectionEvals {
|
||
t.Fatalf("Brent used %d evaluations against bisection's %d, want less than half", len(evals), bisectionEvals)
|
||
}
|
||
t.Logf("Brent %d evaluations, bisection %d for the same tolerance", len(evals), bisectionEvals)
|
||
}
|
||
|
||
// TestFindRootBrentBudgetError pins the bounded iteration count: a run
|
||
// that exhausts MaxIterations reports the budget, never a guess.
|
||
func TestFindRootBrentBudgetError(t *testing.T) {
|
||
f := func(x float64) float64 { return x*x*x - 2*x - 5 }
|
||
_, err := FindRootBrent(f, 1, 3, BrentOptions{MaxIterations: 2})
|
||
if err == nil || !strings.Contains(err.Error(), "no convergence in 2 iterations") {
|
||
t.Fatalf("err = %v, want the budget refusal", err)
|
||
}
|
||
}
|
||
|
||
// TestFindRootBrentPoleInsideBracket pins the mid-run guard: a pole
|
||
// between the endpoints mimics a sign change, and the first bisection
|
||
// of [1, 3] for 1/(x − 2) lands exactly on the pole, whose infinite
|
||
// value is an error naming the point.
|
||
func TestFindRootBrentPoleInsideBracket(t *testing.T) {
|
||
f := func(x float64) float64 { return 1 / (x - 2) }
|
||
_, err := FindRootBrent(f, 1, 3, BrentOptions{})
|
||
if err == nil || !strings.Contains(err.Error(), "left the real numbers") {
|
||
t.Fatalf("err = %v, want the non-finite refusal", err)
|
||
}
|
||
}
|
||
|
||
// TestFindRootBrentNoAllocationsInLoop pins the allocation discipline:
|
||
// the iteration is scalar bookkeeping over the callback, and a
|
||
// converged run allocates nothing.
|
||
func TestFindRootBrentNoAllocationsInLoop(t *testing.T) {
|
||
f := func(x float64) float64 { return x*x*x - 2*x - 5 }
|
||
var err error
|
||
allocs := testing.AllocsPerRun(20, func() {
|
||
_, err = FindRootBrent(f, 1, 3, BrentOptions{})
|
||
})
|
||
if err != nil {
|
||
t.Fatalf("FindRootBrent: %v", err)
|
||
}
|
||
if allocs != 0 {
|
||
t.Fatalf("%g allocations per run, want 0", allocs)
|
||
}
|
||
}
|