Files
tensor/grad/gradient_hygiene_pins_test.go
T

231 lines
7.5 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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)
}
}