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