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