Files
tensor/grad/newtoncg.go
T

249 lines
7.3 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 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
}