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