Files
tensor/grad/doc.go
T

83 lines
3.8 KiB
Go
Raw Permalink 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 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