Files
tensor/grad/pool_bench_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

201 lines
5.2 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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)
}
}
}