Files
tensor/optim/brent.go
T
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

145 lines
4.8 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"
)
// Scalar root bracketing by Brent's method with the iteration's
// budget and tolerance under the caller's control. FindRoot answers
// the common case with fixed settings; FindRootBrent exists for the
// caller who must pin the residual it accepts and the evaluations it
// pays, and who wants a bracket float64 can no longer refine reported
// as a failure instead of receiving a point that misses the
// tolerance.
// BrentOptions tunes FindRootBrent, in the vocabulary of
// RootSystemOptions: Tolerance ≤ 0 means 1e-10, MaxIterations ≤ 0
// means 100. The tolerance is a threshold on |f| at the returned
// point, the scalar counterpart of the residual infinity norm, and a
// run costs at most MaxIterations + 2 evaluations of f.
type BrentOptions struct {
Tolerance float64
MaxIterations int
}
// FindRootBrent returns a root of f inside the bracket [a, b] by
// Brent's method: inverse quadratic interpolation, with the secant
// step as its fallback and bisection as the guarantee. f(a) and f(b)
// must be finite with opposite signs, so a root is guaranteed inside;
// a bracket that does not change sign is an error. The method never
// evaluates f outside the bracket it was given, the returned point
// always lies inside the initial bracket with |f| at or below the
// tolerance, and an exhausted budget is an error, never a silent
// guess.
func FindRootBrent(f func(float64) float64, a, b float64, opts BrentOptions) (float64, error) {
if opts.Tolerance <= 0 {
opts.Tolerance = 1e-10
}
if opts.MaxIterations <= 0 {
opts.MaxIterations = 100
}
fa, fb := f(a), f(b)
if math.IsNaN(fa) || math.IsNaN(fb) || math.IsInf(fa, 0) || math.IsInf(fb, 0) {
return 0, base.Errf("FindRootBrent: the bracket must evaluate to finite values, got f(%g)=%g, f(%g)=%g", a, fa, b, fb)
}
if fa == 0 {
return a, nil
}
if fb == 0 {
return b, nil
}
if sameSign(fa, fb) {
return 0, base.Errf("FindRootBrent: the bracket [%g, %g] does not change sign (f(a)=%g, f(b)=%g)", a, b, fa, fb)
}
// Brent's iteration: b is the best estimate, c the opposite-sign
// end of the bracket, e the width of the step before last. An
// interpolation step is accepted only when it lands safely inside
// the bracket; otherwise the midpoint is taken, so every
// evaluation falls in [b, c] and the sign change survives every
// round.
c, fc := a, fa
d, e := b-a, b-a
for range opts.MaxIterations {
if sameSign(fb, fc) {
// The last step swallowed the opposite sign: drop back to
// the previous iterate, which still brackets the root.
c, fc = a, fa
d, e = b-a, b-a
}
if math.Abs(fc) < math.Abs(fb) {
// Keep b the better end: the smaller residual.
a, b, c = b, c, b
fa, fb, fc = fb, fc, fb
}
// The width below which float64 can no longer separate two
// points of the bracket: the tolerance itself stays out of the
// floor, so the tolerance means |f| and nothing else.
tol1 := 2 * base.EpsF * math.Abs(b)
xm := 0.5 * (c - b)
if fb == 0 || math.Abs(fb) <= opts.Tolerance {
return b, nil
}
if math.Abs(xm) <= tol1 {
return 0, base.Errf("FindRootBrent: the bracket collapsed to machine width at x=%g with |f|=%g, above the tolerance %g",
b, math.Abs(fb), opts.Tolerance)
}
if math.Abs(e) >= tol1 && math.Abs(fa) > math.Abs(fb) {
s := fb / fa
var p, q float64
if a == c {
// Secant.
p = 2 * xm * s
q = 1 - s
} else {
// Inverse quadratic interpolation through (a, fa),
// (b, fb) and (c, fc).
q = fa / fc
r := fb / fc
p = s * (2*xm*q*(q-r) - (b-a)*(r-1))
q = (q - 1) * (r - 1) * (s - 1)
}
if p > 0 {
q = -q
}
p = math.Abs(p)
// Accept the interpolated step only when it stays well
// short of the bracket's far end and at least half the
// previous step; bisection otherwise.
if 2*p < min(3*xm*q-math.Abs(tol1*q), math.Abs(e*q)) {
e, d = d, p/q
} else {
d, e = xm, xm
}
} else {
d, e = xm, xm
}
a, fa = b, fb
if math.Abs(d) > tol1 {
b += d
} else {
b += tol1 * signOf(xm)
}
fb = f(b)
if math.IsNaN(fb) || math.IsInf(fb, 0) {
return 0, base.Errf("FindRootBrent: the objective left the real numbers at x=%g", b)
}
if fb == 0 {
return b, nil
}
}
return 0, base.Errf("FindRootBrent: no convergence in %d iterations", opts.MaxIterations)
}
// sameSign reports whether two finite values carry the same sign. The
// comparison is by sign and not by product, so a pair of tiny values
// whose product underflows to zero is still recognised as
// same-signed, which the product test FindRoot uses would miss.
func sameSign(x, y float64) bool {
return (x > 0 && y > 0) || (x < 0 && y < 0)
}