Files
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

83 lines
3.8 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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