feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,200 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user