231 lines
7.5 KiB
Go
231 lines
7.5 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
|
|
// SPDX-License-Identifier: MIT
|
|||
|
|
|
|||
|
|
package grad
|
|||
|
|
|
|||
|
|
import (
|
|||
|
|
"math"
|
|||
|
|
"strings"
|
|||
|
|
"testing"
|
|||
|
|
|
|||
|
|
ode "sourcedock.dev/petrbalvin/tensor/integrate"
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// Regression tests for gradient hygiene: the Add backward
|
|||
|
|
// handed both inputs the same gradient instance, AdjointODE polluted
|
|||
|
|
// trainable leaves closed over by f but outside params, and
|
|||
|
|
// MinimiseNewtonCG panicked on nil inputs where the rest of the
|
|||
|
|
// package returns errors.
|
|||
|
|
|
|||
|
|
// TestAddBackwardIndependentGradients pins the aliasing fix: the two
|
|||
|
|
// inputs of an Add receive two independent gradient buffers with
|
|||
|
|
// identical values, so a write through one leaf's gradient cannot
|
|||
|
|
// corrupt the other's.
|
|||
|
|
func TestAddBackwardIndependentGradients(t *testing.T) {
|
|||
|
|
x, err := FromFloat64s([]float64{2, 3}, true, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat64s: %v", err)
|
|||
|
|
}
|
|||
|
|
y, err := FromFloat64s([]float64{4, 5}, true, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat64s: %v", err)
|
|||
|
|
}
|
|||
|
|
z, err := x.Add(y)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Add: %v", err)
|
|||
|
|
}
|
|||
|
|
loss, err := z.Sum()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Sum: %v", err)
|
|||
|
|
}
|
|||
|
|
if err := loss.Backward(); err != nil {
|
|||
|
|
t.Fatalf("Backward: %v", err)
|
|||
|
|
}
|
|||
|
|
if x.Grad() == y.Grad() {
|
|||
|
|
t.Fatal("Add handed both leaves the same gradient instance")
|
|||
|
|
}
|
|||
|
|
// d(x+y)/dx = 1 and d(x+y)/dy = 1, element for element.
|
|||
|
|
for i := range 2 {
|
|||
|
|
if x.Grad().FloatAt(i) != 1 {
|
|||
|
|
t.Fatalf("x.grad[%d] = %v, want 1", i, x.Grad().FloatAt(i))
|
|||
|
|
}
|
|||
|
|
if y.Grad().FloatAt(i) != 1 {
|
|||
|
|
t.Fatalf("y.grad[%d] = %v, want 1", i, y.Grad().FloatAt(i))
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
// A write through one leaf's buffer must leave the other's intact.
|
|||
|
|
x.Grad().SetFloatAt(0, 999)
|
|||
|
|
if y.Grad().FloatAt(0) != 1 {
|
|||
|
|
t.Fatalf("y.grad[0] = %v after a write through x's gradient, want 1",
|
|||
|
|
y.Grad().FloatAt(0))
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestAddSameLeafAccumulatesTwice pins Add(x, x): the same leaf as both
|
|||
|
|
// inputs accumulates both contributions into one gradient of 2.
|
|||
|
|
func TestAddSameLeafAccumulatesTwice(t *testing.T) {
|
|||
|
|
x, err := FromFloat64s([]float64{0.5, -1.25}, true, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat64s: %v", err)
|
|||
|
|
}
|
|||
|
|
z, err := x.Add(x)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Add: %v", err)
|
|||
|
|
}
|
|||
|
|
loss, err := z.Sum()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Sum: %v", err)
|
|||
|
|
}
|
|||
|
|
if err := loss.Backward(); err != nil {
|
|||
|
|
t.Fatalf("Backward: %v", err)
|
|||
|
|
}
|
|||
|
|
// d(x+x)/dx = 2, and the accumulated buffer is one array.
|
|||
|
|
for i := range 2 {
|
|||
|
|
if x.Grad().FloatAt(i) != 2 {
|
|||
|
|
t.Fatalf("x.grad[%d] = %v, want 2", i, x.Grad().FloatAt(i))
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestAdjointODELeavesHiddenLeavesClean pins the vjp fix: a trainable
|
|||
|
|
// leaf closed over by f but not listed in params receives no gradient,
|
|||
|
|
// because the Jacobian-vector products come from a pass that commits
|
|||
|
|
// nothing.
|
|||
|
|
func TestAdjointODELeavesHiddenLeavesClean(t *testing.T) {
|
|||
|
|
theta, err := FromFloat64s([]float64{0.7}, true, 1)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat64s: %v", err)
|
|||
|
|
}
|
|||
|
|
hidden, err := FromFloat64s([]float64{1.5}, true, 1)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat64s: %v", err)
|
|||
|
|
}
|
|||
|
|
f := func(t float64, y *Tensor) (*Tensor, error) {
|
|||
|
|
rate, err := theta.Neg()
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
hm, err := y.Mul(hidden)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return hm.Mul(rate)
|
|||
|
|
}
|
|||
|
|
y0, _ := core.FromFloats([]float64{1}, 1)
|
|||
|
|
seed, _ := core.FromFloats([]float64{1}, 1)
|
|||
|
|
_, blocks, err := AdjointODE(f, []*Tensor{theta}, 0, 1, y0, seed,
|
|||
|
|
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("AdjointODE: %v", err)
|
|||
|
|
}
|
|||
|
|
// dL/dθ = −1.5·e^{−1.05}, the central-difference answer the exact
|
|||
|
|
// dynamics give.
|
|||
|
|
want := -1.5 * math.Exp(-1.05)
|
|||
|
|
if math.Abs(blocks[0].FloatAt(0)-want) > 1e-6 {
|
|||
|
|
t.Fatalf("dL/dθ = %.14g, want %.14g", blocks[0].FloatAt(0), want)
|
|||
|
|
}
|
|||
|
|
if hidden.Grad() != nil {
|
|||
|
|
t.Fatalf("the hidden leaf's gradient = %v, want nil", hidden.Grad())
|
|||
|
|
}
|
|||
|
|
if theta.Grad() != nil {
|
|||
|
|
t.Fatalf("the parameter's gradient = %v, want nil", theta.Grad())
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestAdjointODEThetaAgainstCentralDifferences checks the returned
|
|||
|
|
// parameter sensitivity of the same hidden-leaf system against central
|
|||
|
|
// differences of the forward solve, so the cleanup left the θ gradient
|
|||
|
|
// every bit as accurate as it was.
|
|||
|
|
func TestAdjointODEThetaAgainstCentralDifferences(t *testing.T) {
|
|||
|
|
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
|
|||
|
|
hidden, _ := FromFloat64s([]float64{1.5}, true, 1)
|
|||
|
|
f := func(t float64, y *Tensor) (*Tensor, error) {
|
|||
|
|
rate, err := theta.Neg()
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
hm, err := y.Mul(hidden)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return hm.Mul(rate)
|
|||
|
|
}
|
|||
|
|
y0, _ := core.FromFloats([]float64{1}, 1)
|
|||
|
|
seed, _ := core.FromFloats([]float64{1}, 1)
|
|||
|
|
_, blocks, err := AdjointODE(f, []*Tensor{theta}, 0, 1, y0, seed,
|
|||
|
|
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("AdjointODE: %v", err)
|
|||
|
|
}
|
|||
|
|
// Central differences on the loss y(1), with the parameter leaf's
|
|||
|
|
// data swapped out for the perturbed values.
|
|||
|
|
forwardLoss := func() float64 {
|
|||
|
|
end, ferr := ode.IntegrateODE(func(t float64, ya *core.Array) (*core.Array, error) {
|
|||
|
|
out, oerr := f(t, FromArray(ya, false))
|
|||
|
|
if oerr != nil {
|
|||
|
|
return nil, oerr
|
|||
|
|
}
|
|||
|
|
return out.Data(), nil
|
|||
|
|
}, 0, 1, y0, ode.ODEOptions{RelTol: 1e-11, AbsTol: 1e-15})
|
|||
|
|
if ferr != nil {
|
|||
|
|
t.Fatalf("forward solve: %v", ferr)
|
|||
|
|
}
|
|||
|
|
return end.FloatAt(0)
|
|||
|
|
}
|
|||
|
|
const eps = 1e-6
|
|||
|
|
orig := theta.Data().FloatAt(0)
|
|||
|
|
up, _ := core.FromFloats([]float64{orig + eps}, 1)
|
|||
|
|
theta.ReplaceWith(up)
|
|||
|
|
hi := forwardLoss()
|
|||
|
|
dn, _ := core.FromFloats([]float64{orig - eps}, 1)
|
|||
|
|
theta.ReplaceWith(dn)
|
|||
|
|
lo := forwardLoss()
|
|||
|
|
back, _ := core.FromFloats([]float64{orig}, 1)
|
|||
|
|
theta.ReplaceWith(back)
|
|||
|
|
fd := (hi - lo) / (2 * eps)
|
|||
|
|
if math.Abs(blocks[0].FloatAt(0)-fd) > 1e-5*math.Max(1, math.Abs(fd)) {
|
|||
|
|
t.Fatalf("dL/dθ: adjoint %.10g, central difference %.10g",
|
|||
|
|
blocks[0].FloatAt(0), fd)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestMinimiseNewtonCGNilInputs pins the validation contract: a nil
|
|||
|
|
// objective and a nil starting point are errors naming the argument,
|
|||
|
|
// in step with SampleHMC, never panics.
|
|||
|
|
func TestMinimiseNewtonCGNilInputs(t *testing.T) {
|
|||
|
|
f := func(z *Tensor) (*Tensor, error) { return z.Sum() }
|
|||
|
|
if _, _, err := MinimiseNewtonCG(nil, nil, NewtonCGOptions{}); err == nil {
|
|||
|
|
t.Fatal("expected an error for a nil objective")
|
|||
|
|
} else if !strings.Contains(err.Error(), "must not be nil") {
|
|||
|
|
t.Fatalf("error = %v, want a must-not-be-nil refusal", err)
|
|||
|
|
}
|
|||
|
|
x0, err := core.FromFloats([]float64{1}, 1)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
if _, _, err := MinimiseNewtonCG(nil, x0, NewtonCGOptions{}); err == nil {
|
|||
|
|
t.Fatal("expected an error for a nil objective")
|
|||
|
|
} else if !strings.Contains(err.Error(), "f must not be nil") {
|
|||
|
|
t.Fatalf("error = %v, want a refusal naming f", err)
|
|||
|
|
}
|
|||
|
|
if _, _, err := MinimiseNewtonCG(f, nil, NewtonCGOptions{}); err == nil {
|
|||
|
|
t.Fatal("expected an error for a nil starting point")
|
|||
|
|
} else if !strings.Contains(err.Error(), "starting point must not be nil") {
|
|||
|
|
t.Fatalf("error = %v, want a refusal naming the starting point", err)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestODETraceSingleNodeRefused pins the interpolation guard: a trace
|
|||
|
|
// with fewer than two recorded nodes has no interval to interpolate
|
|||
|
|
// over, so the accessor refuses instead of indexing out of range.
|
|||
|
|
func TestODETraceSingleNodeRefused(t *testing.T) {
|
|||
|
|
tr := &odeTrace{times: []float64{1}, states: [][]float64{{2, 3}},
|
|||
|
|
slopes: [][]float64{{0, 0}}, dim: 2}
|
|||
|
|
if _, err := tr.at(1); err == nil {
|
|||
|
|
t.Fatal("expected an error for a one-node trace")
|
|||
|
|
} else if !strings.Contains(err.Error(), "at least two") {
|
|||
|
|
t.Fatalf("error = %v, want a refusal naming the node count", err)
|
|||
|
|
}
|
|||
|
|
}
|