Files
tensor/grad/hessian.go
T

179 lines
5.8 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/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
}