// Copyright (c) 2026 Petr Balvín (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 }