feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
+308
@@ -0,0 +1,308 @@
|
||||
// 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)
|
||||
}
|
||||
Reference in New Issue
Block a user