// Copyright (c) 2026 Petr Balvín (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) }