Files
tensor/grad/example_test.go
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

219 lines
5.7 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_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
}