249 lines
7.3 KiB
Go
249 lines
7.3 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
|
|
// SPDX-License-Identifier: MIT
|
|||
|
|
|
|||
|
|
package grad
|
|||
|
|
|
|||
|
|
import (
|
|||
|
|
"math"
|
|||
|
|
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// Newton-CG minimisation: truncated conjugate gradients on
|
|||
|
|
// the Hessian system, driven by autograd. It lives in the grad package
|
|||
|
|
// because it is meaningless without the graph: the gradients come from
|
|||
|
|
// Backward and the Hessian never forms, each CG iteration buying one
|
|||
|
|
// Hessian-vector product for two backward passes. That is the
|
|||
|
|
// optimiser large problems want, where a dense second derivative does
|
|||
|
|
// not fit memory and the numerical-difference optimisers of the optim
|
|||
|
|
// package lose their accuracy.
|
|||
|
|
|
|||
|
|
// NewtonCGOptions tunes MinimiseNewtonCG. MaxIterations bounds the
|
|||
|
|
// outer Newton steps (default 100); Tolerance stops when the gradient
|
|||
|
|
// norm falls under it (default 1e-8); MaxCGIterations bounds the inner
|
|||
|
|
// CG solve per outer step (default n, the problem dimension).
|
|||
|
|
type NewtonCGOptions struct {
|
|||
|
|
MaxIterations int
|
|||
|
|
Tolerance float64
|
|||
|
|
MaxCGIterations int
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// MinimiseNewtonCG returns the point and value of a local minimum of
|
|||
|
|
// the scalar objective f near x0 by the Newton-CG method: each step
|
|||
|
|
// solves H·s = −∇f with truncated conjugate gradients (negative
|
|||
|
|
// curvature stops the solve and falls back to the first direction),
|
|||
|
|
// then an Armijo backtracking line search secures descent. f receives
|
|||
|
|
// a leaf tensor and must return a single-element real tensor. A
|
|||
|
|
// non-finite objective, an unreachable Armijo condition or an
|
|||
|
|
// exhausted iteration budget is an error naming the state it stopped
|
|||
|
|
// in; the converged answer is a fresh array the caller owns. The
|
|||
|
|
// gradient evaluations differentiate the graph without committing
|
|||
|
|
// anything, so the accumulated gradients of the tensors f closes over
|
|||
|
|
// are left exactly as they were, on the success and the error path
|
|||
|
|
// alike.
|
|||
|
|
func MinimiseNewtonCG(f func(*Tensor) (*Tensor, error), x0 *core.Array, opts NewtonCGOptions) (*core.Array, float64, error) {
|
|||
|
|
const name = "MinimiseNewtonCG"
|
|||
|
|
if f == nil {
|
|||
|
|
return nil, 0, errf("%s: f must not be nil", name)
|
|||
|
|
}
|
|||
|
|
if x0 == nil {
|
|||
|
|
return nil, 0, errf("%s: the starting point must not be nil", name)
|
|||
|
|
}
|
|||
|
|
n := x0.Len()
|
|||
|
|
if n == 0 {
|
|||
|
|
return nil, 0, errf("%s: the starting point must have at least one element", name)
|
|||
|
|
}
|
|||
|
|
if x0.Dtype() == core.Complex {
|
|||
|
|
return nil, 0, errf("%s: complex starting points are not supported", name)
|
|||
|
|
}
|
|||
|
|
maxIter := opts.MaxIterations
|
|||
|
|
if maxIter <= 0 {
|
|||
|
|
maxIter = 100
|
|||
|
|
}
|
|||
|
|
tol := opts.Tolerance
|
|||
|
|
if tol <= 0 {
|
|||
|
|
tol = 1e-8
|
|||
|
|
}
|
|||
|
|
maxCG := opts.MaxCGIterations
|
|||
|
|
if maxCG <= 0 {
|
|||
|
|
maxCG = n
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// eval runs the objective and its reverse pass at point p, returning
|
|||
|
|
// the loss and the flattened gradient. The pass commits nothing, so
|
|||
|
|
// the caller's own gradients survive the evaluation untouched.
|
|||
|
|
eval := func(p *core.Array) (float64, *core.Array, error) {
|
|||
|
|
xt := FromArray(p, true)
|
|||
|
|
y, err := f(xt)
|
|||
|
|
if err != nil {
|
|||
|
|
return 0, nil, errf("%s: %w", name, err)
|
|||
|
|
}
|
|||
|
|
if y.Data().Len() != 1 {
|
|||
|
|
return 0, nil, errf("%s: the objective must return a scalar, got %d elements", name, y.Data().Len())
|
|||
|
|
}
|
|||
|
|
v := y.Data().FloatAt(0)
|
|||
|
|
if math.IsNaN(v) || math.IsInf(v, 0) {
|
|||
|
|
return 0, nil, errf("%s: the objective is non-finite (%g)", name, v)
|
|||
|
|
}
|
|||
|
|
grads, err := y.reverseGrads()
|
|||
|
|
if err != nil {
|
|||
|
|
return 0, nil, errf("%s: %w", name, err)
|
|||
|
|
}
|
|||
|
|
g := grads[xt]
|
|||
|
|
if g == nil {
|
|||
|
|
return 0, nil, errf("%s: the objective does not depend on the starting point", name)
|
|||
|
|
}
|
|||
|
|
return v, g, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
x := x0
|
|||
|
|
f0, g, err := eval(x)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, 0, err
|
|||
|
|
}
|
|||
|
|
for iter := 1; iter <= maxIter; iter++ {
|
|||
|
|
gnorm := flatNorm(g)
|
|||
|
|
if gnorm <= tol {
|
|||
|
|
return clonePoint(x), f0, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// Truncated CG on H·s = −g. The Hessian acts through the
|
|||
|
|
// Hessian-vector product, two backward passes per iteration.
|
|||
|
|
xt := FromArray(x, true)
|
|||
|
|
s := make([]float64, n)
|
|||
|
|
r := make([]float64, n)
|
|||
|
|
p := make([]float64, n)
|
|||
|
|
gs := make([]float64, n)
|
|||
|
|
gFloats := flatFloats(g)
|
|||
|
|
for i := range n {
|
|||
|
|
r[i] = -gFloats[i]
|
|||
|
|
p[i] = r[i]
|
|||
|
|
gs[i] = gFloats[i]
|
|||
|
|
}
|
|||
|
|
rr := 0.0
|
|||
|
|
for i := range n {
|
|||
|
|
rr += r[i] * r[i]
|
|||
|
|
}
|
|||
|
|
for cg := 0; cg < maxCG; cg++ {
|
|||
|
|
pArr, herr := core.FromFloats(p, n)
|
|||
|
|
if herr != nil {
|
|||
|
|
return nil, 0, errf("%s: %w", name, herr)
|
|||
|
|
}
|
|||
|
|
hp, herr2 := HessianVectorProduct(f, xt, FromArray(pArr, false), HessianOptions{})
|
|||
|
|
if herr2 != nil {
|
|||
|
|
return nil, 0, errf("%s: %w", name, herr2)
|
|||
|
|
}
|
|||
|
|
hpF := flatFloats(hp)
|
|||
|
|
pHp := 0.0
|
|||
|
|
for i := range n {
|
|||
|
|
pHp += p[i] * hpF[i]
|
|||
|
|
}
|
|||
|
|
if pHp <= 0 {
|
|||
|
|
// Negative or vanishing curvature: the quadratic model
|
|||
|
|
// is not convex here. The first iteration falls back to
|
|||
|
|
// the steepest descent direction; later ones keep what
|
|||
|
|
// the solve has accumulated.
|
|||
|
|
if cg == 0 {
|
|||
|
|
copy(s, p)
|
|||
|
|
}
|
|||
|
|
break
|
|||
|
|
}
|
|||
|
|
alpha := rr / pHp
|
|||
|
|
for i := range n {
|
|||
|
|
s[i] += alpha * p[i]
|
|||
|
|
r[i] -= alpha * hpF[i]
|
|||
|
|
}
|
|||
|
|
rrNew := 0.0
|
|||
|
|
for i := range n {
|
|||
|
|
rrNew += r[i] * r[i]
|
|||
|
|
}
|
|||
|
|
if math.Sqrt(rrNew) <= 0.1*gnorm {
|
|||
|
|
break
|
|||
|
|
}
|
|||
|
|
beta := rrNew / rr
|
|||
|
|
for i := range n {
|
|||
|
|
p[i] = r[i] + beta*p[i]
|
|||
|
|
}
|
|||
|
|
rr = rrNew
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// Armijo backtracking along s; gᵀs is negative by construction.
|
|||
|
|
gsDot := 0.0
|
|||
|
|
for i := range n {
|
|||
|
|
gsDot += gs[i] * s[i]
|
|||
|
|
}
|
|||
|
|
if gsDot >= 0 {
|
|||
|
|
return nil, 0, errf("%s: the CG direction does not descend at step %d", name, iter)
|
|||
|
|
}
|
|||
|
|
step := 1.0
|
|||
|
|
xFloats := flatFloats(x)
|
|||
|
|
var xNew *core.Array
|
|||
|
|
var fNew float64
|
|||
|
|
accepted := false
|
|||
|
|
for range 40 {
|
|||
|
|
vals := make([]float64, n)
|
|||
|
|
for i := range n {
|
|||
|
|
vals[i] = xFloats[i] + step*s[i]
|
|||
|
|
}
|
|||
|
|
cand, cerr := core.FromFloats(vals, x.Shape()...)
|
|||
|
|
if cerr != nil {
|
|||
|
|
return nil, 0, errf("%s: %w", name, cerr)
|
|||
|
|
}
|
|||
|
|
// cand assigns to the outer xNew; a := here would shadow
|
|||
|
|
// it and hand the post-loop update a nil.
|
|||
|
|
fNew, g, err = eval(cand)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, 0, err
|
|||
|
|
}
|
|||
|
|
if fNew <= f0+1e-4*step*gsDot {
|
|||
|
|
accepted = true
|
|||
|
|
xNew = cand
|
|||
|
|
break
|
|||
|
|
}
|
|||
|
|
step /= 2
|
|||
|
|
}
|
|||
|
|
if !accepted {
|
|||
|
|
return nil, 0, errf("%s: the line search found no descent at step %d (f = %.6g)", name, iter, f0)
|
|||
|
|
}
|
|||
|
|
x = xNew
|
|||
|
|
f0 = fNew
|
|||
|
|
}
|
|||
|
|
// The last accepted step updated g after the loop-top test, so a
|
|||
|
|
// run whose tolerance was met exactly on the final iteration must
|
|||
|
|
// re-test before the budget refusal reports it; the message below
|
|||
|
|
// would otherwise print a gradient already under the tolerance.
|
|||
|
|
if flatNorm(g) <= tol {
|
|||
|
|
return clonePoint(x), f0, nil
|
|||
|
|
}
|
|||
|
|
return nil, 0, errf("%s: no convergence in %d steps (gradient norm %.3g)", name, maxIter, flatNorm(g))
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// flatNorm returns the Euclidean norm of a flattened gradient. The sum
|
|||
|
|
// runs in ascending element order on the raw payload when it can; the
|
|||
|
|
// walk is bounded by the element count, not the payload, because a
|
|||
|
|
// rebased view's storage may run longer than its own elements.
|
|||
|
|
func flatNorm(g *core.Array) float64 {
|
|||
|
|
s := 0.0
|
|||
|
|
if !g.Strided() && g.Dtype() == core.Float {
|
|||
|
|
gs := g.RawFloats()
|
|||
|
|
for i := range g.Len() {
|
|||
|
|
s += gs[i] * gs[i]
|
|||
|
|
}
|
|||
|
|
return math.Sqrt(s)
|
|||
|
|
}
|
|||
|
|
for i := range g.Len() {
|
|||
|
|
s += g.FloatAt(i) * g.FloatAt(i)
|
|||
|
|
}
|
|||
|
|
return math.Sqrt(s)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// clonePoint copies the converged point so the caller owns it.
|
|||
|
|
func clonePoint(x *core.Array) *core.Array {
|
|||
|
|
vals := make([]float64, x.Len())
|
|||
|
|
for i := range vals {
|
|||
|
|
vals[i] = x.FloatAt(i)
|
|||
|
|
}
|
|||
|
|
out, _ := core.FromFloats(vals, x.Shape()...)
|
|||
|
|
return out
|
|||
|
|
}
|