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
|
||
}
|