219 lines
5.7 KiB
Go
219 lines
5.7 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package grad_test
|
||
|
||
// The godoc examples for the autograd package: the flagship workflows
|
||
// as runnable, checked snippets. Each one pins the numbers it prints,
|
||
// so a change in the adjoint of an op or in a solver's default shows
|
||
// up as a failing example rather than as stale prose.
|
||
|
||
import (
|
||
"fmt"
|
||
"log"
|
||
|
||
tensor "sourcedock.dev/petrbalvin/tensor"
|
||
"sourcedock.dev/petrbalvin/tensor/grad"
|
||
)
|
||
|
||
// A scalar loss by hand on a small graph: z = Σ x² over a two-element
|
||
// leaf, differentiated in one reverse sweep. The leaf is entered twice
|
||
// by the product, so the backward adds both contributions and the
|
||
// answer is the 2x the calculus gives.
|
||
func ExampleTensor_Backward() {
|
||
x, err := grad.FromFloat64s([]float64{2, 3}, true, 2)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
sq, err := x.Mul(x)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
loss, err := sq.Sum()
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
if err := loss.Backward(); err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
fmt.Println(loss.Data(), x.Grad())
|
||
// Output: float (1) [13] float (2) [4, 6]
|
||
}
|
||
|
||
// A matrix product and a reduction as graph nodes: the gradient of the
|
||
// sum of A·B is ones·Bᵀ, one row sum of B per row of A.
|
||
func ExampleTensor_MatMul() {
|
||
a, err := grad.FromFloat64s([]float64{1, 2, 3, 4, 5, 6}, true, 2, 3)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
b, err := grad.FromFloat64s([]float64{1, 0, 0, 1, 1, 1}, false, 3, 2)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
prod, err := a.MatMul(b)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
total, err := prod.Sum()
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
if err := total.Backward(); err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
fmt.Println(total.Data(), a.Grad())
|
||
// Output: float (1) [30] float (2, 3) [1, 1, 2, 1, 1, 2]
|
||
}
|
||
|
||
// A complex graph with a real loss. The leaf is complex, the loss is
|
||
// Σ|z|², and the gradient a complex leaf accumulates is ∂L/∂z̄, which
|
||
// for |z|² is z itself.
|
||
func ExampleTensor_Abs2() {
|
||
data, err := tensor.FromComplexes([]complex128{1 + 2i, 3 - 1i}, 2)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
z := grad.FromArray(data, true)
|
||
magnitude, err := z.Abs2()
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
loss, err := magnitude.Sum()
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
if err := loss.Backward(); err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
fmt.Println(loss.Data(), z.Grad())
|
||
// Output: float (1) [15] complex (2) [(1+2i), (3-1i)]
|
||
}
|
||
|
||
// The dense second derivative of Σ x² at (1, 2): the Hessian of a
|
||
// quadratic form is twice its matrix, here 2·I.
|
||
func ExampleHessian() {
|
||
f := func(x *grad.Tensor) (*grad.Tensor, error) {
|
||
sq, err := x.Abs2()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return sq.Sum()
|
||
}
|
||
point, err := grad.FromFloat64s([]float64{1, 2}, true, 2)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
h, err := grad.Hessian(f, point, grad.HessianOptions{})
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
fmt.Println(h)
|
||
// Output: float (2, 2) [2, 0, 0, 2]
|
||
}
|
||
|
||
// The same function and point contracted with a direction: H·v in two
|
||
// gradient evaluations instead of the dense Hessian's four.
|
||
func ExampleHessianVectorProduct() {
|
||
f := func(x *grad.Tensor) (*grad.Tensor, error) {
|
||
sq, err := x.Abs2()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return sq.Sum()
|
||
}
|
||
point, err := grad.FromFloat64s([]float64{1, 2}, true, 2)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
direction, err := grad.FromFloat64s([]float64{1, 1}, false, 2)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
hv, err := grad.HessianVectorProduct(f, point, direction, grad.HessianOptions{})
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
// A central difference along the direction, so the answer carries
|
||
// the rounding of the two gradient evaluations it is built from.
|
||
fmt.Printf("(%.4f, %.4f)\n", hv.FloatAt(0), hv.FloatAt(1))
|
||
// Output: (2.0000, 2.0000)
|
||
}
|
||
|
||
// Newton-CG on the quadratic Σ (x − c)², whose minimiser is c and
|
||
// whose value there is zero. The curvature comes from the
|
||
// Hessian-vector product, so no dense Hessian is ever formed.
|
||
func ExampleMinimiseNewtonCG() {
|
||
centre, err := grad.FromFloat64s([]float64{1.5, -2.5}, false, 2)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
f := func(x *grad.Tensor) (*grad.Tensor, error) {
|
||
diff, err := x.Sub(centre)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
sq, err := diff.Abs2()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return sq.Sum()
|
||
}
|
||
x0, err := tensor.FromFloats([]float64{0, 0}, 2)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
x, value, err := grad.MinimiseNewtonCG(f, x0, grad.NewtonCGOptions{})
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
fmt.Printf("x = (%.4f, %.4f), f = %.4f\n", x.FloatAt(0), x.FloatAt(1), value)
|
||
// Output: x = (1.5000, -2.5000), f = 0.0000
|
||
}
|
||
|
||
// Hamiltonian Monte Carlo on the two-dimensional standard normal,
|
||
// whose log density is −‖q‖²/2. The seed makes the chain reproducible,
|
||
// so the moments of the first component are fixed numbers and not a
|
||
// range: the target has mean zero and variance one.
|
||
func ExampleSampleHMC() {
|
||
logDensity := func(q *grad.Tensor) (*grad.Tensor, error) {
|
||
sq, err := q.Abs2()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
total, err := sq.Sum()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return total.Scale(-0.5)
|
||
}
|
||
q0, err := tensor.FromFloats([]float64{2, -2}, 2)
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
samples, err := grad.SampleHMC(logDensity, q0, grad.HMCOptions{
|
||
Step: 0.25,
|
||
Steps: 16,
|
||
BurnIn: 500,
|
||
Thin: 1,
|
||
Samples: 2000,
|
||
Seed: 7,
|
||
})
|
||
if err != nil {
|
||
log.Fatal(err)
|
||
}
|
||
rows := samples.Shape()[0]
|
||
mean := 0.0
|
||
for row := range rows {
|
||
mean += samples.FloatAt(row * 2)
|
||
}
|
||
mean /= float64(rows)
|
||
variance := 0.0
|
||
for row := range rows {
|
||
d := samples.FloatAt(row*2) - mean
|
||
variance += d * d / float64(rows)
|
||
}
|
||
fmt.Printf("%v: mean %.3f, variance %.3f\n", samples.Shape(), mean, variance)
|
||
// Output: [2000 2]: mean -0.012, variance 1.013
|
||
}
|