314 lines
10 KiB
Go
314 lines
10 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"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
"sourcedock.dev/petrbalvin/tensor/linalg"
|
||
)
|
||
|
||
// Nonlinear equality and inequality constraints through the augmented
|
||
// Lagrangian machinery the linear rows of linear.go already run on.
|
||
//
|
||
// How the nonlinear rows enter: each row carries the same rowPenalty
|
||
// state the linear rows carry (the multipliers, the violation and its
|
||
// side, and the equality flag), so the penalty term and the slope that
|
||
// chains onto the gradient are the shared functions
|
||
// rowPenalty.term and rowPenalty.slope unchanged. What differs is only
|
||
// where the row's value and gradient come from: a linear row reads
|
||
// A·x and a fixed coefficient vector, a nonlinear row calls the
|
||
// caller's function at the measured point and takes central
|
||
// differences of it for the chain rule. The outer loop itself is the
|
||
// same schedule as MinimiseConstrained's: the multipliers ascend by
|
||
// the violation, the quadratic penalty grows by the same factors, and
|
||
// feasibility is judged on the true rows against the same fixed
|
||
// 1e-10 threshold, independent of opts.Tolerance, which stays the
|
||
// inner solver's projected-gradient tolerance. The inner solves stay
|
||
// inexact by design (MinimiseLBFGS with AllowBudgetExit), judged by
|
||
// that outer feasibility check, exactly as the linear entry's document
|
||
// says. Mixed problems need nothing new: a linear row a·x ≤ u is a
|
||
// legal nonlinear row h(x) = a·x − u.
|
||
|
||
// NonlinearConstraints carries functional rows: each equality is
|
||
// g(x) = 0 and each inequality h(x) ≤ 0, both as callbacks receiving
|
||
// the candidate point as a rank-1 array. The functions must return
|
||
// finite values; an error they return is fatal for the run.
|
||
type NonlinearConstraints struct {
|
||
Equalities []func(*core.Array) (float64, error)
|
||
Inequalities []func(*core.Array) (float64, error)
|
||
}
|
||
|
||
// nonlinearRow is one functional row: the shared rowPenalty state (the
|
||
// equality rows sit at lower = upper = 0, the inequality rows at
|
||
// lower = −∞, upper = 0), the caller's function, the scratch for the
|
||
// function's gradient at the last measured point, and the function's
|
||
// own signed value there, which the inequality ascent rides so a slack
|
||
// row's multiplier decays instead of freezing.
|
||
type nonlinearRow struct {
|
||
rowPenalty
|
||
fn func(*core.Array) (float64, error)
|
||
grad []float64
|
||
value float64
|
||
}
|
||
|
||
// MinimiseNonlinearConstrained returns the point, the value of f and
|
||
// the row multipliers of a local minimum of f subject to the nonlinear
|
||
// rows of cons. The multipliers hold the equality rows' signed
|
||
// estimates first, then the inequality rows' non-negative ones, in
|
||
// declaration order; they are the outer loop's final estimates, which
|
||
// converge to the KKT multipliers of the rows active at the answer,
|
||
// while an inactive inequality row's estimate decays to its KKT zero.
|
||
// A row with both slices empty is the unconstrained problem and
|
||
// delegates to MinimiseLBFGS with nil multipliers.
|
||
//
|
||
// The constants of the schedule (the initial penalty of 10, its growth
|
||
// of 10 per round, 40 outer rounds) and the fixed feasibility
|
||
// threshold of 1e-10 are the linear entry's; a run that spends its 40
|
||
// rounds without reaching feasibility is refused with the worst
|
||
// remaining violation, and a badly scaled row can put the inner step
|
||
// beyond the line search's reach, which is refused with the stall
|
||
// diagnostic, never returned as a converged answer.
|
||
func MinimiseNonlinearConstrained(f func(*core.Array) (float64, error),
|
||
grad func(*core.Array) (*core.Array, error),
|
||
x0 *core.Array, cons NonlinearConstraints, opts LBFGSOptions) (*core.Array, float64, []float64, error) {
|
||
const name = "MinimiseNonlinearConstrained"
|
||
n := x0.Len()
|
||
if n == 0 {
|
||
return nil, 0, nil, base.Errf("%s: the starting point must have at least one element", name)
|
||
}
|
||
if x0.Dtype() == core.Complex {
|
||
return nil, 0, nil, base.Errf("%s: complex starting points are not supported", name)
|
||
}
|
||
for i, fn := range cons.Equalities {
|
||
if fn == nil {
|
||
return nil, 0, nil, base.Errf("%s: equality row %d is a nil function", name, i+1)
|
||
}
|
||
}
|
||
for i, fn := range cons.Inequalities {
|
||
if fn == nil {
|
||
return nil, 0, nil, base.Errf("%s: inequality row %d is a nil function", name, i+1)
|
||
}
|
||
}
|
||
if len(cons.Equalities) == 0 && len(cons.Inequalities) == 0 {
|
||
pt, fv, err := MinimiseLBFGS(f, grad, x0, opts)
|
||
return pt, fv, nil, err
|
||
}
|
||
rows := make([]nonlinearRow, 0, len(cons.Equalities)+len(cons.Inequalities))
|
||
for _, fn := range cons.Equalities {
|
||
rows = append(rows, nonlinearRow{
|
||
equal: true,
|
||
fn: fn,
|
||
grad: make([]float64, n),
|
||
})
|
||
}
|
||
for _, fn := range cons.Inequalities {
|
||
rows = append(rows, nonlinearRow{
|
||
lower: math.Inf(-1),
|
||
fn: fn,
|
||
grad: make([]float64, n),
|
||
})
|
||
}
|
||
|
||
x := make([]float64, n)
|
||
for i := range n {
|
||
x[i] = x0.FloatAt(i)
|
||
}
|
||
if opts.Tolerance <= 0 {
|
||
opts.Tolerance = 1e-8
|
||
}
|
||
mu := almMuStart
|
||
worst := 0.0
|
||
// The chain-rule stencils, the point the caller's rows are measured
|
||
// at, and the working copy of a point handed to them: allocated once
|
||
// for the whole run. The stencil carries the offset on one
|
||
// coordinate at a time, restored as soon as that coordinate is
|
||
// differenced, so nothing is copied per coordinate and no slice per
|
||
// variable per row reaches the heap.
|
||
point := make([]float64, n)
|
||
xp := make([]float64, n)
|
||
xm := make([]float64, n)
|
||
// The inner solve's start point and the augmented gradient buffer:
|
||
// one array each for the whole run. MinimiseLBFGS reads the start
|
||
// once at entry and copies every gradient out at once, so both are
|
||
// fully rewritten before the next reader sees them.
|
||
start := core.New(core.Float, n)
|
||
gradBuf := core.New(core.Float, n)
|
||
measure := func(p []float64, wantGrad bool) (float64, error) {
|
||
worst := 0.0
|
||
arr := linalg.ArrayFromFloatsSafe(p, n)
|
||
if wantGrad {
|
||
copy(xp, p)
|
||
copy(xm, p)
|
||
}
|
||
for i := range rows {
|
||
r := &rows[i]
|
||
v, err := r.fn(arr)
|
||
if err != nil {
|
||
return 0, base.Errf("%s: %w", name, err)
|
||
}
|
||
if math.IsNaN(v) || math.IsInf(v, 0) {
|
||
return 0, base.Errf("%s: a row function returned the non-finite value %g", name, v)
|
||
}
|
||
r.value = v
|
||
vio, side := 0.0, 0.0
|
||
switch {
|
||
case v > r.upper:
|
||
vio, side = v-r.upper, 1
|
||
case v < r.lower:
|
||
vio, side = r.lower-v, -1
|
||
}
|
||
r.violate, r.side = vio, side
|
||
worst = math.Max(worst, vio)
|
||
if wantGrad {
|
||
// Central differences of the row function: the chain
|
||
// rule needs the row's gradient at the measured point,
|
||
// the functional counterpart of a linear row's fixed
|
||
// coefficients.
|
||
for j := range n {
|
||
eps := math.Sqrt(base.EpsF) * math.Max(1, math.Abs(p[j]))
|
||
xp[j] += eps
|
||
xm[j] -= eps
|
||
fp, err := r.fn(linalg.ArrayFromFloatsSafe(xp, n))
|
||
if err != nil {
|
||
return 0, base.Errf("%s: %w", name, err)
|
||
}
|
||
fm, err := r.fn(linalg.ArrayFromFloatsSafe(xm, n))
|
||
if err != nil {
|
||
return 0, base.Errf("%s: %w", name, err)
|
||
}
|
||
r.grad[j] = (fp - fm) / (2 * eps)
|
||
xp[j], xm[j] = p[j], p[j]
|
||
}
|
||
}
|
||
}
|
||
return worst, nil
|
||
}
|
||
|
||
// readPoint copies a point into the shared working buffer: the rows
|
||
// are measured through it and it is fully overwritten every call.
|
||
readPoint := func(p *core.Array) []float64 {
|
||
for i := range n {
|
||
point[i] = p.FloatAt(i)
|
||
}
|
||
return point
|
||
}
|
||
|
||
for range almOuterRound {
|
||
objective := func(p *core.Array) (float64, error) {
|
||
fv, err := f(p)
|
||
if err != nil {
|
||
return 0, err
|
||
}
|
||
if _, err := measure(readPoint(p), false); err != nil {
|
||
return 0, err
|
||
}
|
||
total := fv
|
||
for i := range rows {
|
||
total += rows[i].term(mu)
|
||
}
|
||
return total, nil
|
||
}
|
||
objectiveGrad := func(p *core.Array) (*core.Array, error) {
|
||
g, err := grad(p)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if err := requireReal(name, "gradients", g); err != nil {
|
||
return nil, err
|
||
}
|
||
if g.Len() != n {
|
||
return nil, base.Errf("%s: the gradient callback returned %d elements for %d variables",
|
||
name, g.Len(), n)
|
||
}
|
||
if _, err := measure(readPoint(p), true); err != nil {
|
||
return nil, err
|
||
}
|
||
out := gradBuf
|
||
for j := range n {
|
||
out.RawFloats()[j] = g.FloatAt(j)
|
||
}
|
||
for i := range rows {
|
||
w := rows[i].slope(mu)
|
||
if w == 0 {
|
||
continue
|
||
}
|
||
for j := range n {
|
||
out.RawFloats()[j] += w * rows[i].grad[j]
|
||
}
|
||
}
|
||
return out, nil
|
||
}
|
||
innerGrad := objectiveGrad
|
||
if grad == nil {
|
||
innerGrad = nil // the inner solver differences the objective
|
||
}
|
||
|
||
copy(start.RawFloats(), x)
|
||
// The inner solve is inexact by design, as in the linear
|
||
// entry: the outer feasibility check is what judges it.
|
||
inner := opts
|
||
inner.AllowBudgetExit = true
|
||
pt, _, err := MinimiseLBFGS(objective, innerGrad, start, inner)
|
||
if err != nil {
|
||
return nil, 0, nil, base.Errf("%s: %w", name, err)
|
||
}
|
||
for i := range n {
|
||
x[i] = pt.FloatAt(i)
|
||
}
|
||
|
||
// Assigned, not declared: the function-level worst carries the
|
||
// figure the budget refusal reports.
|
||
var merr error
|
||
worst, merr = measure(x, false)
|
||
if merr != nil {
|
||
return nil, 0, nil, base.Errf("%s: %w", name, merr)
|
||
}
|
||
if worst <= almFeasibleAt {
|
||
// The answer carries f's own value, not the augmented
|
||
// Lagrangian's. The multipliers take the estimates the
|
||
// rounds converged to, with complementary slackness pinned
|
||
// at the answer itself: an inequality row that sits strictly
|
||
// slack has KKT multiplier exactly zero, while an active or
|
||
// equality row keeps its estimate.
|
||
for i := range rows {
|
||
r := &rows[i]
|
||
if !r.equal && r.value < -almFeasibleAt {
|
||
r.high = 0
|
||
}
|
||
}
|
||
fv, ferr := f(linalg.ArrayFromFloatsSafe(x, n))
|
||
if ferr != nil {
|
||
return nil, 0, nil, base.Errf("%s: %w", name, ferr)
|
||
}
|
||
multipliers := make([]float64, len(rows))
|
||
for i := range rows {
|
||
multipliers[i] = rows[i].high
|
||
}
|
||
out, fv := packResult(x, fv)
|
||
return out, fv, multipliers, nil
|
||
}
|
||
for i := range rows {
|
||
r := &rows[i]
|
||
// The ascent mirrors the linear entry: an equality row's
|
||
// signed multiplier rides the side, an inequality row's
|
||
// non-negative one rides the row's own signed value, so a
|
||
// violated row lifts it exactly as before while a row that
|
||
// turned slack pulls it back toward its KKT zero instead
|
||
// of freezing at what the violated rounds left.
|
||
if r.equal {
|
||
r.high += mu * r.side * r.violate
|
||
} else {
|
||
r.high = max(0, r.high+mu*r.value)
|
||
}
|
||
}
|
||
if mu < almMuCeiling {
|
||
mu *= almMuGrowth
|
||
}
|
||
}
|
||
return nil, 0, nil, base.Errf("%s: %d rounds left the worst row violation at %g", name, almOuterRound, worst)
|
||
}
|