309 lines
10 KiB
Go
309 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"
|
||
)
|
||
|
||
// Linear constraints handled by the augmented Lagrangian. Each row of
|
||
// A carries a two-sided bound l ≤ A·x ≤ u: a row with l = u is an
|
||
// equality constraint, a one-sided row opens the other side with an
|
||
// infinity. The outer loop escalates the quadratic penalty and ascends
|
||
// the row multipliers until the true violations vanish; the inner
|
||
// problem is the box-capable L-BFGS, so the options bounds compose
|
||
// with the linear rows.
|
||
|
||
// LinearConstraints holds the rows of l ≤ A·x ≤ u. A is an r×n matrix
|
||
// over the n variables, Lower and Upper hold one entry per row, and an
|
||
// infinite entry opens that side. A row with Lower = Upper is an
|
||
// equality constraint.
|
||
type LinearConstraints struct {
|
||
A *core.Array
|
||
Lower []float64
|
||
Upper []float64
|
||
}
|
||
|
||
// rowPenalty is one row's augmented-Lagrangian state: the coefficient
|
||
// vector, the multipliers (High carries the equality row's signed
|
||
// multiplier; the two one-sided multipliers of an inequality row are
|
||
// both non-negative) and the violation measured at the point of the
|
||
// last objective or gradient call.
|
||
type rowPenalty struct {
|
||
coeffs []float64
|
||
lower float64
|
||
upper float64
|
||
equal bool
|
||
high float64
|
||
low float64
|
||
violate float64
|
||
side float64 // +1 above the upper wall, -1 below the lower one
|
||
}
|
||
|
||
// The augmented-Lagrangian schedule both constraint entries walk: the
|
||
// linear rows of MinimiseConstrained and the nonlinear rows of
|
||
// MinimiseNonlinearConstrained share one penalty path, one ceiling and
|
||
// one feasibility gate, so the two cannot silently diverge.
|
||
const (
|
||
almMuStart = 10.0
|
||
almMuGrowth = 10.0
|
||
almMuCeiling = 1e10
|
||
almOuterRound = 40
|
||
// almFeasibleAt is the absolute row violation the outer loop
|
||
// accepts as feasible. Like the other tolerances in the package it
|
||
// is absolute, deliberately independent of opts.Tolerance (which
|
||
// stays the inner solver's projected-gradient tolerance): tying
|
||
// the two made a caller who asked for a tighter inner solve get a
|
||
// hard failure instead of a tighter answer.
|
||
almFeasibleAt = 1e-10
|
||
)
|
||
|
||
// measure evaluates every row at p, records each violation and its
|
||
// side for the penalty terms, and returns the worst true violation.
|
||
func measure(p *core.Array, rows []rowPenalty) float64 {
|
||
worst := 0.0
|
||
n := p.Len()
|
||
for i := range rows {
|
||
r := &rows[i]
|
||
ax := 0.0
|
||
for j := range n {
|
||
ax += r.coeffs[j] * p.FloatAt(j)
|
||
}
|
||
v, side := 0.0, 0.0
|
||
switch {
|
||
case ax > r.upper:
|
||
v, side = ax-r.upper, 1
|
||
case ax < r.lower:
|
||
v, side = r.lower-ax, -1
|
||
}
|
||
r.violate, r.side = v, side
|
||
if v > worst {
|
||
worst = v
|
||
}
|
||
}
|
||
return worst
|
||
}
|
||
|
||
// slope is d(penalty row)/d(ax) at the measured point: the multiplier
|
||
// of the active side plus the penalty's derivative. Zero inside the
|
||
// row's bounds.
|
||
func (r *rowPenalty) slope(mu float64) float64 {
|
||
switch r.side {
|
||
case 0:
|
||
return 0
|
||
case 1:
|
||
// An equality row violates to one side only, so its slope is
|
||
// the same whichever sign the violation carries.
|
||
return r.high + mu*r.violate
|
||
default:
|
||
if r.equal {
|
||
return r.high - mu*r.violate
|
||
}
|
||
return -(r.low + mu*r.violate)
|
||
}
|
||
}
|
||
|
||
// term is the row's contribution to the augmented Lagrangian at the
|
||
// measured point: the multiplier times the signed violation plus the
|
||
// quadratic penalty. An equality row's multiplier is signed and rides
|
||
// the side; an inequality row's one-sided multipliers are
|
||
// non-negative.
|
||
func (r *rowPenalty) term(mu float64) float64 {
|
||
if r.side == 0 {
|
||
return 0
|
||
}
|
||
mult := r.high
|
||
if !r.equal && r.side < 0 {
|
||
mult = r.low
|
||
}
|
||
if r.equal {
|
||
mult *= r.side
|
||
}
|
||
return mult*r.violate + 0.5*mu*r.violate*r.violate
|
||
}
|
||
|
||
// MinimiseConstrained returns the point and value of a local minimum
|
||
// of f subject to the box walls carried by opts and the linear rows
|
||
// l ≤ A·x ≤ u. It is the only entry point whose matrix is read
|
||
// directly, so only here is a complex A refused beside the usual dtype
|
||
// check on the starting point. Each row is used in the caller's own
|
||
// units: the violation, the multipliers and the penalty inherit the
|
||
// row's scale, so a row whose coefficients sit many orders above
|
||
// another's is enforced to a correspondingly tighter absolute
|
||
// precision. Each outer round minimises f plus the rows' augmented
|
||
// Lagrangian terms, then ascends the row multipliers and scales the
|
||
// quadratic penalty. The inner solves are MinimiseLBFGS warm-started
|
||
// from the previous round's point; the supplied gradient, when given,
|
||
// is the gradient of f and the rows' contributions are chained onto it
|
||
// analytically, and the callback must return n real values. A start
|
||
// outside the feasible set is fine: the outer rounds pull it back, and
|
||
// failure to reach feasibility is an error naming the worst remaining
|
||
// violation. Feasibility is judged on the caller's rows against a
|
||
// fixed absolute threshold of 1e-10, independent of opts.Tolerance:
|
||
// the option stays the inner solver's projected-gradient tolerance,
|
||
// and asking for a tighter inner solve never turns into a hard failure
|
||
// of the outer loop. A badly scaled row (a·x = 1 with a of 1e9 or
|
||
// more) can put the step that reduces the augmented Lagrangian beyond
|
||
// the inner line search's reach: such a run is refused with the stall
|
||
// diagnostic, never returned as a converged answer.
|
||
func MinimiseConstrained(f func(*core.Array) (float64, error), grad func(*core.Array) (*core.Array, error),
|
||
x0 *core.Array, cons LinearConstraints, opts LBFGSOptions) (*core.Array, float64, error) {
|
||
const name = "MinimiseConstrained"
|
||
n := x0.Len()
|
||
if n == 0 {
|
||
return nil, 0, base.Errf("%s: the starting point must have at least one element", name)
|
||
}
|
||
if x0.Dtype() == core.Complex {
|
||
return nil, 0, base.Errf("%s: complex starting points are not supported", name)
|
||
}
|
||
if cons.A == nil {
|
||
return nil, 0, base.Errf("%s: the constraint matrix is nil", name)
|
||
}
|
||
if err := requireReal(name, "constraint matrices", cons.A); err != nil {
|
||
return nil, 0, err
|
||
}
|
||
if cons.A.NDim() != 2 || cons.A.Shape()[1] != n {
|
||
return nil, 0, base.Errf("%s: the constraint matrix is %s, want r×%d", name, base.ShapeText(cons.A.Shape()), n)
|
||
}
|
||
rowCount := cons.A.Shape()[0]
|
||
if rowCount == 0 {
|
||
return nil, 0, base.Errf("%s: the constraint matrix has no rows", name)
|
||
}
|
||
if len(cons.Lower) != rowCount || len(cons.Upper) != rowCount {
|
||
return nil, 0, base.Errf("%s: the bounds hold %d and %d entries for %d rows",
|
||
name, len(cons.Lower), len(cons.Upper), rowCount)
|
||
}
|
||
rows := make([]rowPenalty, rowCount)
|
||
for i := range rows {
|
||
r := &rows[i]
|
||
r.coeffs = make([]float64, n)
|
||
for j := range n {
|
||
a := cons.A.FloatAt(i*n + j)
|
||
if math.IsNaN(a) || math.IsInf(a, 0) {
|
||
return nil, 0, base.Errf("%s: row %d carries a non-finite coefficient", name, i+1)
|
||
}
|
||
r.coeffs[j] = a
|
||
}
|
||
if math.IsNaN(cons.Lower[i]) || math.IsNaN(cons.Upper[i]) || cons.Lower[i] > cons.Upper[i] {
|
||
return nil, 0, base.Errf("%s: row %d has bounds [%g, %g]", name, i+1, cons.Lower[i], cons.Upper[i])
|
||
}
|
||
r.lower, r.upper = cons.Lower[i], cons.Upper[i]
|
||
r.equal = r.lower == r.upper
|
||
}
|
||
|
||
x := make([]float64, n)
|
||
for i := range n {
|
||
x[i] = x0.FloatAt(i)
|
||
}
|
||
if opts.Tolerance <= 0 {
|
||
opts.Tolerance = 1e-8
|
||
}
|
||
mu := almMuStart
|
||
// 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)
|
||
|
||
for range almOuterRound {
|
||
objective := func(p *core.Array) (float64, error) {
|
||
fv, err := f(p)
|
||
if err != nil {
|
||
return 0, err
|
||
}
|
||
measure(p, rows)
|
||
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
|
||
}
|
||
// The payload is validated here, before anything reads it:
|
||
// RawFloats() is nil for a complex or non-float payload and
|
||
// as short as the callback's array, so slicing it first
|
||
// panics instead of reporting the contract breach.
|
||
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)
|
||
}
|
||
measure(p, rows)
|
||
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].coeffs[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: the outer loop judges
|
||
// the point by the rows' feasibility, so a run that spends its
|
||
// iteration budget without reaching the projected-gradient
|
||
// tolerance still carries the outer loop forward.
|
||
inner := opts
|
||
inner.AllowBudgetExit = true
|
||
pt, _, err := MinimiseLBFGS(objective, innerGrad, start, inner)
|
||
if err != nil {
|
||
return nil, 0, base.Errf("%s: %w", name, err)
|
||
}
|
||
for i := range n {
|
||
x[i] = pt.FloatAt(i)
|
||
}
|
||
|
||
if measure(linalg.ArrayFromFloatsSafe(x, n), rows) <= almFeasibleAt {
|
||
// The returned value is f(x) itself, not the augmented
|
||
// Lagrangian the inner solver minimised: the penalty terms
|
||
// may be small but their multipliers are not, and the
|
||
// caller asked for the objective.
|
||
fv, ferr := f(linalg.ArrayFromFloatsSafe(x, n))
|
||
if ferr != nil {
|
||
return nil, 0, base.Errf("%s: %w", name, ferr)
|
||
}
|
||
out, _ := packResult(x, fv)
|
||
return out, fv, nil
|
||
}
|
||
for i := range rows {
|
||
r := &rows[i]
|
||
switch {
|
||
case r.equal:
|
||
r.high += mu * r.side * r.violate
|
||
case r.side > 0:
|
||
r.high = max(0, r.high+mu*r.violate)
|
||
case r.side < 0:
|
||
r.low = max(0, r.low+mu*r.violate)
|
||
}
|
||
}
|
||
if mu < almMuCeiling {
|
||
mu *= almMuGrowth
|
||
}
|
||
}
|
||
worst := measure(linalg.ArrayFromFloatsSafe(x, n), rows)
|
||
return nil, 0, base.Errf("%s: %d rounds left the worst row violation at %g", name, almOuterRound, worst)
|
||
}
|