145 lines
4.8 KiB
Go
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)
|
||
|
|
}
|