Files
tensor/grad/newtoncg.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

249 lines
7.3 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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
}