// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "math" "testing" ode "sourcedock.dev/petrbalvin/tensor/integrate" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // decayWith builds f for y' = −θy with θ a scalar parameter leaf. func decayWith(theta *Tensor) func(t float64, y *Tensor) (*Tensor, error) { return func(t float64, y *Tensor) (*Tensor, error) { rate, err := theta.Neg() if err != nil { return nil, err } return y.Mul(rate) } } // scalarVector returns a length-1 real array. func scalarVector(t *testing.T, v float64) *core.Array { t.Helper() a, err := core.FromFloats([]float64{v}, 1) if err != nil { t.Fatalf("FromFloats: %v", err) } return a } // TestAdjointODEDecay differentiates y' = −θy with L = y(1): the exact // sensitivities are dL/dy0 = e^{−θ} and dL/dθ = −e^{−θ}, and a run // with no parameters at all must still answer the initial-state one. func TestAdjointODEDecay(t *testing.T) { const want = 0.4965853037914095 // e^{−0.7} theta, err := FromFloat64s([]float64{0.7}, true, 1) if err != nil { t.Fatalf("FromFloat64s: %v", err) } gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0, 1, scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12}) if err != nil { t.Fatalf("AdjointODE: %v", err) } if math.Abs(gradY0.FloatAt(0)-want) > 1e-6 { t.Fatalf("dL/dy0 = %.14g, want %.14g", gradY0.FloatAt(0), want) } if len(blocks) != 1 || math.Abs(blocks[0].FloatAt(0)+want) > 1e-6 { t.Fatalf("dL/dθ = %v, want −%.14g", blocks[0].FloatAt(0), want) } // The parameter's own accumulated gradient must be untouched. if theta.Grad() != nil { t.Fatal("AdjointODE must leave the parameters' gradients untouched") } solo, soloBlocks, err := AdjointODE(decayWith(theta), nil, 0, 1, scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12}) if err != nil { t.Fatalf("AdjointODE without parameters: %v", err) } if len(soloBlocks) != 0 || math.Abs(solo.FloatAt(0)-want) > 1e-6 { t.Fatalf("parameter-free run: dL/dy0 = %.14g with %d blocks", solo.FloatAt(0), len(soloBlocks)) } } // TestAdjointODEOscillator differentiates the harmonic oscillator // y” = −ω²y with L = y(1) and y0 = (1, 0): dL/dω = −sin(ω), // dL/dy0 = (cos ω, sin ω/ω). func TestAdjointODEOscillator(t *testing.T) { omega, err := FromFloat64s([]float64{1.3}, true, 1) if err != nil { t.Fatalf("FromFloat64s: %v", err) } f := func(t float64, y *Tensor) (*Tensor, error) { u, err := y.Slice(0, 0, 1) if err != nil { return nil, err } v, err := y.Slice(0, 1, 2) if err != nil { return nil, err } sq, err := omega.Mul(omega) if err != nil { return nil, err } acc, err := u.Mul(sq) if err != nil { return nil, err } drag, err := acc.Scale(-1) if err != nil { return nil, err } return v.Concat(drag, 0) } y0, err := core.FromFloats([]float64{1, 0}, 2) if err != nil { t.Fatalf("FromFloats: %v", err) } seed, err := core.FromFloats([]float64{1, 0}, 2) if err != nil { t.Fatalf("FromFloats: %v", err) } gradY0, blocks, err := AdjointODE(f, []*Tensor{omega}, 0, 1, y0, seed, ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12}) if err != nil { t.Fatalf("AdjointODE: %v", err) } if math.Abs(gradY0.FloatAt(0)-math.Cos(1.3)) > 1e-6 { t.Fatalf("dL/du0 = %.14g, want cos(1.3) = %.14g", gradY0.FloatAt(0), math.Cos(1.3)) } if math.Abs(gradY0.FloatAt(1)-math.Sin(1.3)/1.3) > 1e-6 { t.Fatalf("dL/dv0 = %.14g, want sin(1.3)/1.3 = %.14g", gradY0.FloatAt(1), math.Sin(1.3)/1.3) } if math.Abs(blocks[0].FloatAt(0)+math.Sin(1.3)) > 1e-6 { t.Fatalf("dL/dω = %.14g, want −sin(1.3) = %.14g", blocks[0].FloatAt(0), -math.Sin(1.3)) } } // TestAdjointODEBackwardTime runs the forward pass itself backwards // (t1 < t0): y(t) = e^{−θ(t−1)} from y(1) = 1 has y(0) = e^θ and both // sensitivities equal e^θ. func TestAdjointODEBackwardTime(t *testing.T) { const want = 1.6487212707001282 // e^{0.5} theta, _ := FromFloat64s([]float64{0.5}, true, 1) gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 1, 0, scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12}) if err != nil { t.Fatalf("AdjointODE: %v", err) } if math.Abs(gradY0.FloatAt(0)-want) > 1e-6 { t.Fatalf("dL/dy0 = %.14g, want %.14g", gradY0.FloatAt(0), want) } if math.Abs(blocks[0].FloatAt(0)-want) > 1e-6 { t.Fatalf("dL/dθ = %.14g, want %.14g", blocks[0].FloatAt(0), want) } } // TestAdjointODEFiniteDifference checks a nonlinear two-parameter // system against central differences of the forward solve itself, // perturbing the parameter leaves around the adjoint run. func TestAdjointODEFiniteDifference(t *testing.T) { th1, _ := FromFloat64s([]float64{0.8}, true, 1) th2, _ := FromFloat64s([]float64{1.1}, true, 1) f := func(t float64, y *Tensor) (*Tensor, error) { y1, err := y.Slice(0, 0, 1) if err != nil { return nil, err } y2, err := y.Slice(0, 1, 2) if err != nil { return nil, err } drag, err := y1.Mul(th1) if err != nil { return nil, err } drive, err := y2.Mul(th2) if err != nil { return nil, err } r1, err := drive.Sub(drag) if err != nil { return nil, err } pool, err := y1.Mul(y2) if err != nil { return nil, err } r2, err := pool.Scale(-1) if err != nil { return nil, err } return r1.Concat(r2, 0) } y0, _ := core.FromFloats([]float64{1, 0.5}, 2) seed, _ := core.FromFloats([]float64{1, 2}, 2) gradY0, blocks, err := AdjointODE(f, []*Tensor{th1, th2}, 0, 1, y0, seed, ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-13}) if err != nil { t.Fatalf("AdjointODE: %v", err) } // Central differences on the loss y1(1) + 2·y2(1), with the leaf // data swapped out for the perturbed values. forwardLoss := func() float64 { end, err := ode.IntegrateODE(func(t float64, ya *core.Array) (*core.Array, error) { out, err := f(t, FromArray(ya, false)) if err != nil { return nil, err } return out.Data(), nil }, 0, 1, y0, ode.ODEOptions{RelTol: 1e-11, AbsTol: 1e-15}) if err != nil { t.Fatalf("forward solve: %v", err) } return end.FloatAt(0) + 2*end.FloatAt(1) } perturb := func(p *Tensor, eps float64) { swapped, _ := core.FromFloats([]float64{p.Data().FloatAt(0) + eps}, 1) p.ReplaceWith(swapped) } restore := func(p *Tensor, v float64) { orig, _ := core.FromFloats([]float64{v}, 1) p.ReplaceWith(orig) } const eps = 1e-5 for k, p := range []*Tensor{th1, th2} { orig := p.Data().FloatAt(0) perturb(p, eps) up := forwardLoss() perturb(p, -2*eps) down := forwardLoss() restore(p, orig) fd := (up - down) / (2 * eps) got := blocks[k].FloatAt(0) if math.Abs(got-fd) > 1e-3*math.Max(1, math.Abs(fd)) { t.Fatalf("dL/dθ%d: adjoint %.8g, finite difference %.8g", k+1, got, fd) } } if gradY0.Len() != 2 { t.Fatalf("dL/dy0 has length %d, want 2", gradY0.Len()) } } // TestAdjointODEDegenerateSpan returns the loss seed unchanged and // zero parameter gradients when the span carries no dynamics. func TestAdjointODEDegenerateSpan(t *testing.T) { theta, _ := FromFloat64s([]float64{0.7}, true, 1) gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0.5, 0.5, scalarVector(t, 1), scalarVector(t, 2), ode.ODEOptions{}) if err != nil { t.Fatalf("AdjointODE: %v", err) } if gradY0.FloatAt(0) != 2 { t.Fatalf("dL/dy0 = %.14g, want the seed 2", gradY0.FloatAt(0)) } if blocks[0].FloatAt(0) != 0 { t.Fatalf("dL/dθ = %.14g, want 0", blocks[0].FloatAt(0)) } } // TestAdjointODEErrors pins the validation contract. func TestAdjointODEErrors(t *testing.T) { theta, _ := FromFloat64s([]float64{0.7}, true, 1) y0 := scalarVector(t, 1) seed := scalarVector(t, 1) rank2, _ := core.FromFloats([]float64{1, 1}, 1, 2) if _, _, err := AdjointODE(nil, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil { t.Fatal("expected an error for a nil function") } if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, nil, seed, ode.ODEOptions{}); err == nil { t.Fatal("expected an error for a nil state") } if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, rank2, seed, ode.ODEOptions{}); err == nil { t.Fatal("expected an error for a rank-2 state") } if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, y0, nil, ode.ODEOptions{}); err == nil { t.Fatal("expected an error for a nil loss seed") } badSeed, _ := core.FromFloats([]float64{1, 1}, 2) if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, y0, badSeed, ode.ODEOptions{}); err == nil { t.Fatal("expected an error for a loss seed of the wrong length") } frozen := FromArray(scalarVector(t, 0.7), false) if _, _, err := AdjointODE(decayWith(frozen), []*Tensor{frozen}, 0, 1, y0, seed, ode.ODEOptions{}); err == nil { t.Fatal("expected an error for a parameter that does not require grad") } wrongShape := func(t float64, y *Tensor) (*Tensor, error) { three, _ := core.FromFloats([]float64{1, 2, 3}, 3) return FromArray(three, false), nil } if _, _, err := AdjointODE(wrongShape, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil { t.Fatal("expected an error for a wrong-shaped derivative") } ignoresState := func(t float64, y *Tensor) (*Tensor, error) { one, _ := core.FromFloats([]float64{1}, 1) return FromArray(one, false), nil } if _, _, err := AdjointODE(ignoresState, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil { t.Fatal("expected an error when f ignores its state") } } // TestAdjointODERestoresGradientsOnError pins the error-path contract: // when f consumes the parameter but ignores the state, no state // gradient can flow and the run errors, and the caller's own // accumulated parameter gradient must come back untouched instead of // being polluted by the aborted pass: the reverse sweep is not run at // all, so nothing writes the leaf gradients on the way out. func TestAdjointODERestoresGradientsOnError(t *testing.T) { theta, _ := FromFloat64s([]float64{0.7}, true, 1) preset, _ := core.FromFloats([]float64{3}, 1) theta.SetGrad(preset) f := func(t float64, y *Tensor) (*Tensor, error) { return theta.Scale(2) } if _, _, err := AdjointODE(f, []*Tensor{theta}, 0, 1, scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{}); err == nil { t.Fatal("expected an error when f ignores its state") } if theta.Grad() == nil || theta.Grad().FloatAt(0) != 3 { t.Fatalf("the parameter's gradient was not restored: %v, want 3", theta.Grad()) } } // TestAdjointODEDisconnectedParameter pins that a parameter the // dynamics never touch is an error, not a silent zero gradient block. func TestAdjointODEDisconnectedParameter(t *testing.T) { theta, _ := FromFloat64s([]float64{0.7}, true, 1) f := func(t float64, y *Tensor) (*Tensor, error) { return y.Scale(2) } _, blocks, err := AdjointODE(f, []*Tensor{theta}, 0, 1, scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{}) if err == nil { t.Fatal("expected an error for a parameter disconnected from the dynamics") } if blocks != nil { t.Fatalf("an errored run returned blocks: %v", blocks) } } // TestAdjointODESuccessKeepsPresetGradient pins the same restore on // the success path: a preset accumulated gradient survives the run. func TestAdjointODESuccessKeepsPresetGradient(t *testing.T) { theta, _ := FromFloat64s([]float64{0.5}, true, 1) preset, _ := core.FromFloats([]float64{7}, 1) theta.SetGrad(preset) gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0, 1, scalarVector(t, 1), scalarVector(t, 2), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12}) if err != nil { t.Fatalf("AdjointODE: %v", err) } const want = 0.6065306597126334 // e^{−0.5} if math.Abs(gradY0.FloatAt(0)-2*want) > 1e-6 { t.Fatalf("dL/dy0 = %.14g, want 2e^{−0.5}", gradY0.FloatAt(0)) } if theta.Grad() == nil || theta.Grad().FloatAt(0) != 7 { t.Fatalf("the preset gradient did not survive a successful run: %v", theta.Grad()) } if blocks[0].FloatAt(0) >= 0 { t.Fatalf("dL/dθ = %v, want negative", blocks[0].FloatAt(0)) } }