// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "math" "runtime" "testing" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // The gradient pool's contract, measured: a repeated backward sweep on // one process must keep the heap flat rather than growing with the // iteration count, the sweep's arithmetic must be bit-for-bit // reproducible across runs that share the pool, and the per-sweep cost // itself is pinned by benchmarks that separate graph construction from // the reverse pass. // tapeChain builds a chain of n element-wise nodes over x and w and // reduces it to a scalar, the fixture the sweep benchmarks repeat. func tapeChain(t testing.TB, x, w *Tensor, n int) *Tensor { t.Helper() h := x for i := range n { var err error switch i % 4 { case 0: h, err = h.Add(w) case 1: h, err = h.Mul(w) case 2: h, err = h.Tanh() default: h, err = h.Scale(0.25) } if err != nil { t.Fatal(err) } } s, err := h.Sum() if err != nil { t.Fatal(err) } return s } func tapeLeaf(t testing.TB, seed, n int) *Tensor { t.Helper() v := make([]float64, n) for i := range v { v[i] = 0.25 + float64((i*7+seed)%13)*0.125 } a, err := core.FromFloats(v, n) if err != nil { t.Fatal(err) } return FromArray(a, true) } // TestGradPoolFlatHeap runs one repeated backward workload and checks // the live heap stops growing: the pool's retention caps must bound // what one process holds, however many sweeps it serves. func TestGradPoolFlatHeap(t *testing.T) { x, w := tapeLeaf(t, 1, 8), tapeLeaf(t, 2, 8) s := tapeChain(t, x, w, 128) run := func() { for range 500 { x.ZeroGrad() w.ZeroGrad() if err := s.Backward(); err != nil { t.Fatal(err) } } } var early, late runtime.MemStats runtime.GC() run() runtime.GC() runtime.ReadMemStats(&early) run() run() runtime.GC() runtime.ReadMemStats(&late) // Two more batches of a thousand sweeps may add pool slack but not // a growth trend: the second reading stays within a small factor of // the first, which a leaking pool would break. if late.HeapInuse > early.HeapInuse*2+1<<20 { t.Fatalf("heap grew across repeated sweeps: %d then %d bytes in use", early.HeapInuse, late.HeapInuse) } } // TestGradBackwardDeterminismBits runs the same program twice through // the pooled sweep and demands identical gradient bits: recycling a // buffer must never leak a previous sweep's values into a result. func TestGradBackwardDeterminismBits(t *testing.T) { gradOf := func() []float64 { x, w := tapeLeaf(t, 3, 8), tapeLeaf(t, 4, 8) s := tapeChain(t, x, w, 64) if err := s.Backward(); err != nil { t.Fatal(err) } gx, gw := x.Grad(), w.Grad() if gx == nil || gw == nil { t.Fatal("missing leaf gradient") } out := make([]float64, 0, gx.Len()+gw.Len()) out = append(out, gx.RawFloats()[:gx.Len()]...) out = append(out, gw.RawFloats()[:gw.Len()]...) return out } // Warm the pool with unrelated sweeps, so the measured runs borrow // recycled buffers carrying other work's values. for range 64 { a, b := tapeLeaf(t, 9, 8), tapeLeaf(t, 10, 8) s := tapeChain(t, a, b, 32) if err := s.Backward(); err != nil { t.Fatal(err) } } first, second := gradOf(), gradOf() if len(first) != len(second) { t.Fatalf("gradient lengths differ: %d and %d", len(first), len(second)) } for i := range first { if math.Float64bits(first[i]) != math.Float64bits(second[i]) { t.Fatalf("gradient bit %d differs: %x and %x", i, math.Float64bits(first[i]), math.Float64bits(second[i])) } } } // BenchmarkTapeChainSweepBackward measures the reverse sweep alone // on a 129-node chain that is built once: every allocation here is the // sweep's own, not the graph's. func BenchmarkTapeChainSweepBackward(b *testing.B) { x, w := tapeLeaf(b, 1, 8), tapeLeaf(b, 2, 8) s := tapeChain(b, x, w, 128) b.ReportAllocs() for b.Loop() { x.ZeroGrad() w.ZeroGrad() if err := s.Backward(); err != nil { b.Fatal(err) } } } // BenchmarkWaveTapeChainForwardRebuild measures building the same chain // afresh with no backward pass, the per-node graph-construction cost // the sweep benchmarks otherwise carry inside their loop. func BenchmarkWaveTapeChainForwardRebuild(b *testing.B) { x, w := tapeLeaf(b, 1, 8), tapeLeaf(b, 2, 8) b.ReportAllocs() for b.Loop() { s := tapeChain(b, x, w, 128) if s.Data().Len() != 1 { b.Fatal("unexpected shape") } } } // BenchmarkWaveWideFanSweepBackward measures the reverse sweep of a // 64-way fan over one shared leaf: the fold-heavy edge pattern, built // once. func BenchmarkWaveWideFanSweepBackward(b *testing.B) { x := tapeLeaf(b, 3, 16) leaves := make([]*Tensor, 64) for i := range leaves { leaves[i] = tapeLeaf(b, 10+i, 16) } acc, err := x.Mul(leaves[0]) if err != nil { b.Fatal(err) } for _, l := range leaves[1:] { p, err := x.Mul(l) if err != nil { b.Fatal(err) } if acc, err = acc.Add(p); err != nil { b.Fatal(err) } } s, err := acc.Sum() if err != nil { b.Fatal(err) } b.ReportAllocs() for b.Loop() { x.ZeroGrad() for _, l := range leaves { l.ZeroGrad() } if err := s.Backward(); err != nil { b.Fatal(err) } } }