413 lines
13 KiB
Go
413 lines
13 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package optim
|
||
|
|
|
||
|
|
import (
|
||
|
|
"fmt"
|
||
|
|
"math"
|
||
|
|
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
|
|
)
|
||
|
|
|
||
|
|
// Scalar root finding and multivariate minimisation. Two root finders
|
||
|
|
// cover the common cases: FindRoot needs only a sign-changing bracket
|
||
|
|
// and converges unconditionally; FindRootNewton needs the derivative
|
||
|
|
// and converges quadratically when a good starting guess and a smooth
|
||
|
|
// derivative are available. Minimise is the derivative-free simplex
|
||
|
|
// method, the standard choice for objectives that are noisy, opaque or
|
||
|
|
// expensive to differentiate.
|
||
|
|
|
||
|
|
// FindRoot returns a root of f in the bracket [a, b] by Brent's
|
||
|
|
// method, which combines inverse quadratic interpolation, the secant
|
||
|
|
// step and bisection. f(a) and f(b) must be finite with opposite
|
||
|
|
// signs, so a root is guaranteed inside. A tol ≤ 0 defaults to 1e-12.
|
||
|
|
func FindRoot(f func(float64) float64, a, b, tol float64) (float64, error) {
|
||
|
|
if tol <= 0 {
|
||
|
|
tol = 1e-12
|
||
|
|
}
|
||
|
|
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("FindRoot: 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 fa*fb > 0 {
|
||
|
|
return 0, base.Errf("FindRoot: the bracket [%g, %g] does not change sign (f(a)=%g, f(b)=%g)", a, b, fa, fb)
|
||
|
|
}
|
||
|
|
// Brent's iteration (fbrent): c tracks the opposite-sign end, e
|
||
|
|
// the previous step width; interpolation is tried first and
|
||
|
|
// bisection keeps the step honest.
|
||
|
|
c, fc := a, fa
|
||
|
|
d, e := b-a, b-a
|
||
|
|
for range 200 {
|
||
|
|
if fb*fc > 0 {
|
||
|
|
c, fc = a, fa
|
||
|
|
d, e = b-a, b-a
|
||
|
|
}
|
||
|
|
if math.Abs(fc) < math.Abs(fb) {
|
||
|
|
a, b, c = b, c, b
|
||
|
|
fa, fb, fc = fb, fc, fb
|
||
|
|
}
|
||
|
|
tol1 := 2*base.EpsF*math.Abs(b) + 0.5*tol
|
||
|
|
xm := 0.5 * (c - b)
|
||
|
|
if math.Abs(xm) <= tol1 || fb == 0 {
|
||
|
|
return b, nil
|
||
|
|
}
|
||
|
|
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.
|
||
|
|
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)
|
||
|
|
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 fb == 0 {
|
||
|
|
return b, nil
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return 0, base.Errf("FindRoot: no convergence in 200 iterations")
|
||
|
|
}
|
||
|
|
|
||
|
|
// FindRootNewton returns a root of f near x0 by Newton's iteration
|
||
|
|
// with the supplied derivative df. A tol ≤ 0 defaults to 1e-12 and
|
||
|
|
// maxIter ≤ 0 to 100. A vanishing derivative or an exhausted budget
|
||
|
|
// is an error, not a silent guess.
|
||
|
|
func FindRootNewton(f, df func(float64) float64, x0, tol float64, maxIter int) (float64, error) {
|
||
|
|
if tol <= 0 {
|
||
|
|
tol = 1e-12
|
||
|
|
}
|
||
|
|
if maxIter <= 0 {
|
||
|
|
maxIter = 100
|
||
|
|
}
|
||
|
|
x := x0
|
||
|
|
for range maxIter {
|
||
|
|
fx := f(x)
|
||
|
|
if math.IsNaN(fx) || math.IsInf(fx, 0) {
|
||
|
|
return 0, base.Errf("FindRootNewton: the objective left the real numbers at x=%g", x)
|
||
|
|
}
|
||
|
|
d := df(x)
|
||
|
|
if d == 0 {
|
||
|
|
return 0, base.Errf("FindRootNewton: the derivative vanishes at x=%g", x)
|
||
|
|
}
|
||
|
|
step := fx / d
|
||
|
|
x -= step
|
||
|
|
if math.Abs(step) <= tol*(1+math.Abs(x)) {
|
||
|
|
return x, nil
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return 0, base.Errf("FindRootNewton: no convergence in %d steps from x0=%g", maxIter, x0)
|
||
|
|
}
|
||
|
|
|
||
|
|
// MinimiseOptions tunes the simplex minimisation. MaxIterations ≤ 0
|
||
|
|
// means 2000, Tolerance ≤ 0 means 1e-10, InitialStep ≤ 0 means 1.
|
||
|
|
//
|
||
|
|
// Both convergence tests are absolute in the objective's own scale:
|
||
|
|
// the value spread is compared against Tolerance·max(1, |f|) and the
|
||
|
|
// simplex diameter against Tolerance·max(1, |x|). An objective whose
|
||
|
|
// values sit many orders of magnitude below one therefore counts as
|
||
|
|
// flat from the start, and its start point is returned with a nil
|
||
|
|
// error; rescale the objective (and the variables) to O(1) before
|
||
|
|
// calling Minimise when the natural units are not of that size.
|
||
|
|
type MinimiseOptions struct {
|
||
|
|
MaxIterations int
|
||
|
|
Tolerance float64
|
||
|
|
InitialStep float64
|
||
|
|
// AllowBudgetExit makes a run that exhausts MaxIterations report
|
||
|
|
// its best point instead of an error. The default is false, so a
|
||
|
|
// budget stop is never mistaken for a converged answer; the flag
|
||
|
|
// mirrors LBFGSOptions.AllowBudgetExit.
|
||
|
|
AllowBudgetExit bool
|
||
|
|
}
|
||
|
|
|
||
|
|
// Minimise returns the point and value of a local minimum of f near
|
||
|
|
// x0 by the Nelder-Mead simplex method, which needs no derivatives.
|
||
|
|
// The answer is a local minimum: multistart from several x0 when the
|
||
|
|
// objective may have several basins.
|
||
|
|
func Minimise(f func(*core.Array) (float64, error), x0 *core.Array, opts MinimiseOptions) (*core.Array, float64, error) {
|
||
|
|
if x0.Dtype() == core.Complex {
|
||
|
|
return nil, 0, base.Errf("Minimise: complex starting points are not supported")
|
||
|
|
}
|
||
|
|
n := x0.Len()
|
||
|
|
if n == 0 {
|
||
|
|
return nil, 0, base.Errf("Minimise: the starting point must have at least one element")
|
||
|
|
}
|
||
|
|
if opts.MaxIterations <= 0 {
|
||
|
|
opts.MaxIterations = 2000
|
||
|
|
}
|
||
|
|
if opts.Tolerance <= 0 {
|
||
|
|
opts.Tolerance = 1e-10
|
||
|
|
}
|
||
|
|
if opts.InitialStep <= 0 {
|
||
|
|
opts.InitialStep = 1
|
||
|
|
}
|
||
|
|
|
||
|
|
eval := func(v []float64) (float64, error) {
|
||
|
|
a, err := core.FromFloats(v, n)
|
||
|
|
if err != nil {
|
||
|
|
return 0, err
|
||
|
|
}
|
||
|
|
fv, ferr := f(a)
|
||
|
|
if ferr != nil {
|
||
|
|
return 0, ferr
|
||
|
|
}
|
||
|
|
// A non-finite objective is an error naming the point: the
|
||
|
|
// spread test compares false against NaN, so a NaN vertex
|
||
|
|
// would burn the whole budget and be reported as a budget
|
||
|
|
// problem, and with AllowBudgetExit set it can sit at index 0
|
||
|
|
// and come back as the answer. The error stays bare: every
|
||
|
|
// call site of eval wraps it once with "Minimise: %w" through
|
||
|
|
// base.Errf, which adds the entry-point name and the library
|
||
|
|
// tag, so the surfaced message carries exactly one prefix
|
||
|
|
// instead of the doubled "tensor: Minimise: tensor:
|
||
|
|
// Minimise:" a prefixed inner error produced.
|
||
|
|
if math.IsNaN(fv) || math.IsInf(fv, 0) {
|
||
|
|
return 0, fmt.Errorf("f returned the non-finite value %g at %v", fv, v)
|
||
|
|
}
|
||
|
|
return fv, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// The simplex: n+1 vertices, x0 plus one offset per coordinate.
|
||
|
|
simplex := make([][]float64, n+1)
|
||
|
|
values := make([]float64, n+1)
|
||
|
|
simplex[0] = cloneDense(x0)
|
||
|
|
for i := range n {
|
||
|
|
v := cloneDense(x0)
|
||
|
|
step := opts.InitialStep * math.Max(1, math.Abs(v[i]))
|
||
|
|
v[i] += step
|
||
|
|
simplex[i+1] = v
|
||
|
|
}
|
||
|
|
for i, v := range simplex {
|
||
|
|
fv, err := eval(v)
|
||
|
|
if err != nil {
|
||
|
|
return nil, 0, base.Errf("Minimise: %w", err)
|
||
|
|
}
|
||
|
|
values[i] = fv
|
||
|
|
}
|
||
|
|
|
||
|
|
// Scratch reused across iterations: the centroid and the reflected,
|
||
|
|
// expanded and contracted candidates. A candidate that wins is
|
||
|
|
// copied into the worst vertex's slot, whose slice then keeps its
|
||
|
|
// identity; the scratch is fully rewritten before every read.
|
||
|
|
centre := make([]float64, n)
|
||
|
|
reflected := make([]float64, n)
|
||
|
|
expanded := make([]float64, n)
|
||
|
|
contracted := make([]float64, n)
|
||
|
|
|
||
|
|
converged := false
|
||
|
|
for range opts.MaxIterations {
|
||
|
|
// Order so vertex 0 is the best and vertex n the worst.
|
||
|
|
orderSimplex(simplex, values)
|
||
|
|
spread := values[n] - values[0]
|
||
|
|
if spread <= opts.Tolerance*math.Max(1, math.Abs(values[0])) {
|
||
|
|
// A small value spread on its own is not convergence: the
|
||
|
|
// vertices can agree on the value while spanning the
|
||
|
|
// parameter space, because all of them sit on one level set
|
||
|
|
// of the objective. The spatial test catches exactly that
|
||
|
|
// simplex, whose best vertex is no minimum at all.
|
||
|
|
if simplexDiameter(simplex) <= opts.Tolerance*math.Max(1, maxAbs(simplex[0])) {
|
||
|
|
converged = true
|
||
|
|
break
|
||
|
|
}
|
||
|
|
// A stalled level set is collapsed onto its best vertex:
|
||
|
|
// each shrink halves the diameter, so the spatial test is
|
||
|
|
// reached in a bounded number of rounds and the returned
|
||
|
|
// point is the best one the simplex found, never a random
|
||
|
|
// vertex of an equally-valued set.
|
||
|
|
if err := shrinkSimplex(simplex, values, eval); err != nil {
|
||
|
|
return nil, 0, base.Errf("Minimise: %w", err)
|
||
|
|
}
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
worst := simplex[n]
|
||
|
|
clear(centre)
|
||
|
|
for i := range n {
|
||
|
|
for j := range n {
|
||
|
|
centre[j] += simplex[i][j] / float64(n)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// Reflect the worst vertex through the centroid.
|
||
|
|
for j := range n {
|
||
|
|
reflected[j] = centre[j] + (centre[j] - worst[j])
|
||
|
|
}
|
||
|
|
fr, err := eval(reflected)
|
||
|
|
if err != nil {
|
||
|
|
return nil, 0, base.Errf("Minimise: %w", err)
|
||
|
|
}
|
||
|
|
switch {
|
||
|
|
case fr < values[0]:
|
||
|
|
// Reflected better than the best: try expanding further.
|
||
|
|
for j := range n {
|
||
|
|
expanded[j] = centre[j] + 2*(centre[j]-worst[j])
|
||
|
|
}
|
||
|
|
fe, err := eval(expanded)
|
||
|
|
if err != nil {
|
||
|
|
return nil, 0, base.Errf("Minimise: %w", err)
|
||
|
|
}
|
||
|
|
if fe < fr {
|
||
|
|
copy(worst, expanded)
|
||
|
|
values[n] = fe
|
||
|
|
} else {
|
||
|
|
copy(worst, reflected)
|
||
|
|
values[n] = fr
|
||
|
|
}
|
||
|
|
case fr < values[n-1]:
|
||
|
|
copy(worst, reflected)
|
||
|
|
values[n] = fr
|
||
|
|
default:
|
||
|
|
// Reflected worse than the second worst: contract.
|
||
|
|
for j := range n {
|
||
|
|
contracted[j] = centre[j] + 0.5*(worst[j]-centre[j])
|
||
|
|
}
|
||
|
|
fc, err := eval(contracted)
|
||
|
|
if err != nil {
|
||
|
|
return nil, 0, base.Errf("Minimise: %w", err)
|
||
|
|
}
|
||
|
|
if fc < values[n] {
|
||
|
|
copy(worst, contracted)
|
||
|
|
values[n] = fc
|
||
|
|
break
|
||
|
|
}
|
||
|
|
// Shrink everything towards the best vertex.
|
||
|
|
if err := shrinkSimplex(simplex, values, eval); err != nil {
|
||
|
|
return nil, 0, base.Errf("Minimise: %w", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
// Falling out of the loop means the budget ran out, not that a
|
||
|
|
// minimum was found: the simplex still moves, and reporting its
|
||
|
|
// best vertex as the answer is the silent wrongness the level-set
|
||
|
|
// and stall exits refuse. The final move happened after the top-of-
|
||
|
|
// loop test, so the convergence pair is re-checked once first.
|
||
|
|
orderSimplex(simplex, values)
|
||
|
|
if values[n]-values[0] <= opts.Tolerance*math.Max(1, math.Abs(values[0])) &&
|
||
|
|
simplexDiameter(simplex) <= opts.Tolerance*math.Max(1, maxAbs(simplex[0])) {
|
||
|
|
converged = true
|
||
|
|
}
|
||
|
|
if !converged && !opts.AllowBudgetExit {
|
||
|
|
return nil, 0, base.Errf("Minimise: the iteration budget of %d ran out without the simplex converging", opts.MaxIterations)
|
||
|
|
}
|
||
|
|
// values[0] is f(simplex[0]) by construction: orderSimplex keeps the
|
||
|
|
// pairs together and every update writes both, so the answer is the
|
||
|
|
// held value, not a fresh evaluation of the best vertex.
|
||
|
|
fv := values[0]
|
||
|
|
out, err := core.FromFloats(simplex[0], n)
|
||
|
|
if err != nil {
|
||
|
|
return nil, 0, err
|
||
|
|
}
|
||
|
|
return out, fv, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// orderSimplex sorts the vertices together with their values by
|
||
|
|
// ascending value.
|
||
|
|
func orderSimplex(simplex [][]float64, values []float64) {
|
||
|
|
for i := 1; i < len(values); i++ {
|
||
|
|
for j := i; j > 0 && values[j] < values[j-1]; j-- {
|
||
|
|
simplex[j], simplex[j-1] = simplex[j-1], simplex[j]
|
||
|
|
values[j], values[j-1] = values[j-1], values[j]
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// simplexDiameter returns the largest L∞ distance from the simplex's
|
||
|
|
// best vertex to another one: the spatial extent of the simplex, which
|
||
|
|
// the convergence test pairs with the value spread so that a set of
|
||
|
|
// vertices lying on one level set of the objective is never read as a
|
||
|
|
// converged answer.
|
||
|
|
func simplexDiameter(simplex [][]float64) float64 {
|
||
|
|
n := len(simplex[0])
|
||
|
|
diameter := 0.0
|
||
|
|
for i := 1; i < len(simplex); i++ {
|
||
|
|
for j := range n {
|
||
|
|
if d := math.Abs(simplex[i][j] - simplex[0][j]); d > diameter {
|
||
|
|
diameter = d
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return diameter
|
||
|
|
}
|
||
|
|
|
||
|
|
// shrinkSimplex halves every vertex's distance to the best one and
|
||
|
|
// re-evaluates the moved vertices. It is both the classic Nelder-Mead
|
||
|
|
// shrink, taken when a contraction failed, and the remedy for a
|
||
|
|
// level-set stall, where the vertices agree on the value without
|
||
|
|
// spanning a small neighbourhood. Each call halves the diameter, so a
|
||
|
|
// stalled simplex collapses onto its best point in a bounded number of
|
||
|
|
// rounds.
|
||
|
|
func shrinkSimplex(simplex [][]float64, values []float64, eval func([]float64) (float64, error)) error {
|
||
|
|
n := len(simplex) - 1
|
||
|
|
for i := 1; i <= n; i++ {
|
||
|
|
for j := range n {
|
||
|
|
simplex[i][j] = simplex[0][j] + 0.5*(simplex[i][j]-simplex[0][j])
|
||
|
|
}
|
||
|
|
fv, err := eval(simplex[i])
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
values[i] = fv
|
||
|
|
}
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// signOf returns ±1 for the sign of x (0 counts as +).
|
||
|
|
func signOf(x float64) float64 {
|
||
|
|
if x < 0 {
|
||
|
|
return -1
|
||
|
|
}
|
||
|
|
return 1
|
||
|
|
}
|
||
|
|
|
||
|
|
// cloneDense copies an array's elements into a plain float64 slice.
|
||
|
|
func cloneDense(a *core.Array) []float64 {
|
||
|
|
vals := make([]float64, a.Len())
|
||
|
|
for i := range vals {
|
||
|
|
vals[i] = a.FloatAt(i)
|
||
|
|
}
|
||
|
|
return vals
|
||
|
|
}
|
||
|
|
|
||
|
|
// requireReal refuses a complex array a caller supplied as a
|
||
|
|
// constraint matrix or as a callback payload. The complex payload has
|
||
|
|
// no real part to read, so the library's FloatAt dereferences a nil
|
||
|
|
// payload and panics; every optim entry point answers complex input
|
||
|
|
// with an error instead. name identifies the entry point, what the
|
||
|
|
// payload, so the message reads like the dtype refusals the family
|
||
|
|
// already raises ("complex starting points are not supported").
|
||
|
|
func requireReal(name, what string, a *core.Array) error {
|
||
|
|
if a.Dtype() != core.Complex {
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
return base.Errf("%s: complex %s are not supported", name, what)
|
||
|
|
}
|