// Copyright (c) 2026 Petr Balvín (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) } }