// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "math" "sourcedock.dev/petrbalvin/tensor/internal/base" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // Second-order differentiation, forward-over-reverse: the // inner derivative is the exact analytic gradient Backward produces, // and only the outer derivative runs by central differences over one // coordinate at a time. The result carries the accuracy of the exact // first derivative with the O(h²) truncation of the outer stencil, the // same trade a hand-written finite-difference Hessian makes but with // none of the first-order error. // HessianOptions tunes Hessian and HessianVectorProduct. Step is the // absolute coordinate perturbation (≤ 0 picks sqrt(eps)·max(1, |x_i|) // per coordinate, the stencil that balances truncation against // cancellation at double precision). type HessianOptions struct { Step float64 } // Hessian returns the Hessian matrix of a scalar function f at x, an // (n, n) float64 array for an n-element x. f receives a tensor that // requires grad and must return a single-element real tensor; complex // outputs are rejected like Backward does. The cost is 2n gradient // evaluations, the price of a dense second derivative by any method // that does not exploit structure; for large n prefer // HessianVectorProduct. The 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 Hessian(f func(*Tensor) (*Tensor, error), x *Tensor, opts HessianOptions) (*core.Array, error) { if x.Data().Dtype() == core.Complex { return nil, base.Errf("Hessian: complex points are not supported") } n := x.Data().Len() if n == 0 { return nil, base.Errf("Hessian: the point must not be empty") } p := flatFloats(x.Data()) out := zeros(core.Float, []int{n, n}) h := opts.Step for j := range n { step := h if step <= 0 { step = math.Sqrt(2.220446049250313e-16) * math.Max(1, math.Abs(p[j])) } gp, err := hessianColumn(f, x, p, j, step) if err != nil { return nil, err } gm, err := hessianColumn(f, x, p, j, -step) if err != nil { return nil, err } inv2h := 1 / (2 * step) for i := range n { out.SetFloatAt(i*n+j, (gp[i]-gm[i])*inv2h) } } return out, nil } // hessianColumn evaluates the analytic gradient of f at the point // perturbed by step along coordinate j, flattened. func hessianColumn(f func(*Tensor) (*Tensor, error), x *Tensor, p []float64, j int, step float64) ([]float64, error) { probe := append([]float64(nil), p...) probe[j] += step pa, err := core.FromFloats(probe, x.Data().Shape()...) if err != nil { return nil, base.Errf("Hessian: %w", err) } xt := FromArray(pa, true) y, err := f(xt) if err != nil { return nil, base.Errf("Hessian: %w", err) } if y.Data().Len() != 1 { return nil, base.Errf("Hessian: f must return a scalar, got %d elements", y.Data().Len()) } grads, err := y.reverseGrads() if err != nil { return nil, base.Errf("Hessian: %w", err) } g := grads[xt] if g == nil { return nil, base.Errf("Hessian: the objective does not depend on x, so no gradient exists") } return flatFloats(g), nil } // HessianVectorProduct returns H·v, the Hessian of the scalar f at x // contracted with the direction v, by a central difference along the // direction itself, with the step scaled so it never depends on v's // magnitude. Two gradient evaluations // answer for any n, which is what makes Newton-CG tractable where a // dense Hessian is not. As in Hessian, the evaluations leave the // accumulated gradients of every tensor f closes over untouched. func HessianVectorProduct(f func(*Tensor) (*Tensor, error), x, v *Tensor, opts HessianOptions) (*core.Array, error) { if x.Data().Dtype() == core.Complex { return nil, base.Errf("HessianVectorProduct: complex points are not supported") } if v.Data().Dtype() == core.Complex { return nil, base.Errf("HessianVectorProduct: complex directions are not supported") } n := x.Data().Len() if v.Data().Len() != n { return nil, base.Errf("HessianVectorProduct: direction has %d elements for %d variables", v.Data().Len(), n) } vn := 0.0 for _, v := range flatFloats(v.Data()) { vn += v * v } vn = math.Sqrt(vn) if vn == 0 { // H·0 = 0 in the shape of the point, the same shape the // quotient below returns: a flat vector here would change the // result's shape with the direction's norm. return zeros(core.Float, x.Data().Shape()), nil } h := opts.Step if h <= 0 { h = 1e-5 } p := flatFloats(x.Data()) vf := flatFloats(v.Data()) eval := func(sign float64) ([]float64, error) { probe := make([]float64, n) for i := range n { probe[i] = p[i] + sign*h*vf[i]/vn } pa, err := core.FromFloats(probe, x.Data().Shape()...) if err != nil { return nil, base.Errf("HessianVectorProduct: %w", err) } xt := FromArray(pa, true) y, err := f(xt) if err != nil { return nil, base.Errf("HessianVectorProduct: %w", err) } if y.Data().Len() != 1 { return nil, base.Errf("HessianVectorProduct: f must return a scalar, got %d elements", y.Data().Len()) } grads, err := y.reverseGrads() if err != nil { return nil, base.Errf("HessianVectorProduct: %w", err) } g := grads[xt] if g == nil { return nil, base.Errf("HessianVectorProduct: the objective does not depend on x, so no gradient exists") } return flatFloats(g), nil } gp, err := eval(1) if err != nil { return nil, err } gm, err := eval(-1) if err != nil { return nil, err } // The step advanced h·v/|v| along v, so the quotient is (H·v/|v|) // and carries the |v| factor back in. out := zeros(core.Float, x.Data().Shape()) inv2h := vn / (2 * h) for i := range n { out.SetFloatAt(i, (gp[i]-gm[i])*inv2h) } return out, nil }