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