Files
tensor/optim/nonlinearconstr.go
T

314 lines
10 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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)
}