83 lines
3.8 KiB
Go
83 lines
3.8 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
|
|
// SPDX-License-Identifier: MIT
|
|||
|
|
|
|||
|
|
// Package grad is reverse-mode automatic differentiation over the array
|
|||
|
|
// surface: a computation written as ordinary Go calls over [Tensor]
|
|||
|
|
// values records a graph, and one call to [Tensor.Backward] propagates
|
|||
|
|
// the gradient from the output back to every leaf that requires it.
|
|||
|
|
//
|
|||
|
|
// # The graph
|
|||
|
|
//
|
|||
|
|
// Leaves come from [FromFloat64s] or [FromArray]. Every differentiable
|
|||
|
|
// method records the operation it performs together with its inputs, and
|
|||
|
|
// Backward sweeps the recorded nodes in reverse, applying each node's
|
|||
|
|
// adjoint. The method set covers the arithmetic, the matrix products
|
|||
|
|
// (single and batched), the element-wise transcendentals, the
|
|||
|
|
// reductions, slicing, concatenation, axis permutation and the Fourier
|
|||
|
|
// transforms, with [Tensor.Conj], [Tensor.Real], [Tensor.Imag] and
|
|||
|
|
// [Tensor.Abs2] carrying complex values into a real loss.
|
|||
|
|
//
|
|||
|
|
// loss, err := total.Mean() // the forward pass records the nodes
|
|||
|
|
// if err != nil {
|
|||
|
|
// return err
|
|||
|
|
// }
|
|||
|
|
// if err := loss.Backward(); err != nil {
|
|||
|
|
// return err
|
|||
|
|
// }
|
|||
|
|
// x.Grad() // the accumulated gradient
|
|||
|
|
//
|
|||
|
|
// # The contract
|
|||
|
|
//
|
|||
|
|
// - float, float32 and complex128 tensors differentiate; an int tensor
|
|||
|
|
// is refused by every differentiable method.
|
|||
|
|
// - The loss must be real: Backward seeds the output with ones, and a
|
|||
|
|
// complex output is rejected with an error naming Real, Imag, Abs
|
|||
|
|
// and Abs2 as the reducers that turn it into a real scalar.
|
|||
|
|
// - Gradients carry the leaf's dtype. A mixed-dtype graph narrows each
|
|||
|
|
// gradient to the dtype of the tensor it accumulates into before the
|
|||
|
|
// leaf is written.
|
|||
|
|
// - Backward accumulates into the gradients already present, so
|
|||
|
|
// [Tensor.ZeroGrad] precedes a fresh pass unless accumulation is
|
|||
|
|
// wanted.
|
|||
|
|
// - The graph is rebuilt on every forward pass. Each operation's
|
|||
|
|
// backward closure captures the operands as they were when the
|
|||
|
|
// operation ran, so a [Tensor.ReplaceWith] afterwards changes the
|
|||
|
|
// next pass and not the recorded one.
|
|||
|
|
//
|
|||
|
|
// # Complex graphs
|
|||
|
|
//
|
|||
|
|
// Complex tensors differentiate under the Wirtinger convention: the
|
|||
|
|
// gradient a complex leaf accumulates is ∂L/∂z̄, the coefficient g of
|
|||
|
|
// dL = 2·Re(g·dz), which is the direction gradient descent steps along.
|
|||
|
|
// The adjoint of a holomorphic y = f(z) is therefore dz = g·conj(f′(z)),
|
|||
|
|
// and every complex adjoint in this package conjugates exactly where the
|
|||
|
|
// calculus puts it. A real tensor inside a complex graph narrows the
|
|||
|
|
// incoming gradient by 2·Re, the factor that also cancels the ½ the Real
|
|||
|
|
// and Imag backward paths contribute, so a graph mixing the two dtypes
|
|||
|
|
// composes exactly.
|
|||
|
|
//
|
|||
|
|
// # Second-order and solver tools
|
|||
|
|
//
|
|||
|
|
// On top of the graph sit the second derivative and the methods that
|
|||
|
|
// need one: [Hessian] (dense, 2n gradient evaluations),
|
|||
|
|
// [HessianVectorProduct] (H·v in two gradient evaluations),
|
|||
|
|
// [MinimiseNewtonCG] (truncated conjugate gradients on the Hessian
|
|||
|
|
// system with an Armijo line search), [SampleHMC] (Hamiltonian Monte
|
|||
|
|
// Carlo on any differentiable unnormalised log density) and [AdjointODE]
|
|||
|
|
// (adjoint sensitivities of an ODE solution at the cost of one extra
|
|||
|
|
// solve).
|
|||
|
|
//
|
|||
|
|
// Each of them differentiates through a reverse pass that commits
|
|||
|
|
// nothing, so the accumulated gradients of the tensors the caller's
|
|||
|
|
// closure holds are left exactly as they were, whether the call succeeds
|
|||
|
|
// or fails.
|
|||
|
|
//
|
|||
|
|
// # What it does not do
|
|||
|
|
//
|
|||
|
|
// There is no forward-mode differentiation, nothing beyond the second
|
|||
|
|
// derivative, no graph serialisation and no parameter registry: the
|
|||
|
|
// caller owns the leaves and the graph is a transient record of one
|
|||
|
|
// forward pass. The operation set is closed, so a new primitive is a new
|
|||
|
|
// method here and never a user-registered op.
|
|||
|
|
package grad
|