// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "math" "sync" "testing" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // Pins for the pooled reverse sweep: the pooled Backward must answer // bit-identically to the legacy map-returning sweep on the same graph, // concurrent sweeps on separate graphs must answer the serial reference // bits, and the pool's retention caps must refuse releases past them. // pinLit builds a leaf of n elements from fixed literals, the fixture // shape the tape benchmarks use. func pinLit(seed, n int) *Tensor { 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 { panic(err) } return FromArray(a, true) } // chainLeafBits builds the same deep chain the tape benchmarks build, // runs one pooled Backward and returns the two leaf gradients' raw // bits. The chain fans both leaves into every node, so every fold // accumulates multiple contributions. func chainLeafBits(seedA, seedB, nodes int) ([]float64, error) { x, w := pinLit(seedA, 8), pinLit(seedB, 8) s, err := deepChain(x, w, nodes) if err != nil { return nil, err } x.ZeroGrad() w.ZeroGrad() if err := s.Backward(); err != nil { return nil, err } xg, wg := x.Grad(), w.Grad() if xg == nil || wg == nil { return nil, errf("pinned chain: missing leaf gradient") } out := make([]float64, 0, 16) out = append(out, xg.RawFloats()[:8]...) out = append(out, wg.RawFloats()[:8]...) return out, nil } func TestPooledBackwardBitsMatchLegacySweep(t *testing.T) { pooled, err := chainLeafBits(1, 2, 64) if err != nil { t.Fatalf("pooled sweep: %v", err) } // The same graph through the legacy sweep: reverseGrads commits // nothing, so its map carries the leaves' gradients from this pass // alone, which is what the pooled sweep commits on a fresh leaf. x, w := pinLit(1, 8), pinLit(2, 8) s, err := deepChain(x, w, 64) if err != nil { t.Fatalf("legacy chain: %v", err) } grads, err := s.reverseGrads() if err != nil { t.Fatalf("legacy sweep: %v", err) } gx, gw := grads[x], grads[w] if gx == nil || gw == nil { t.Fatal("legacy sweep returned no leaf gradient") } legacy := append(append([]float64{}, gx.RawFloats()[:8]...), gw.RawFloats()[:8]...) if len(legacy) != len(pooled) { t.Fatalf("length %d, want %d", len(pooled), len(legacy)) } for i := range pooled { if math.Float64bits(pooled[i]) != math.Float64bits(legacy[i]) { t.Fatalf("leaf gradient %d: pooled %#x, legacy %#x", i, math.Float64bits(pooled[i]), math.Float64bits(legacy[i])) } } } func TestConcurrentBackwardDeterminism(t *testing.T) { ref, err := chainLeafBits(3, 4, 48) if err != nil { t.Fatalf("serial reference: %v", err) } const sweeps = 40 outs := make([][]float64, 2) errs := make([]error, 2) var wg sync.WaitGroup for g := range 2 { wg.Go(func() { for range sweeps { bits, err := chainLeafBits(3, 4, 48) if err != nil { errs[g] = err return } outs[g] = bits } }) } wg.Wait() for g := range 2 { if errs[g] != nil { t.Fatalf("goroutine %d: %v", g, errs[g]) } for i := range ref { if math.Float64bits(outs[g][i]) != math.Float64bits(ref[i]) { t.Fatalf("goroutine %d element %d: %#x, want %#x", g, i, math.Float64bits(outs[g][i]), math.Float64bits(ref[i])) } } } } func TestGradPoolCapsAreEnforced(t *testing.T) { k, ok := poolKeyOf(core.Float, []int{8}) if !ok { t.Fatal("poolKeyOf refused a float shape of 8") } keep, err := core.Zeros(core.Float, 8) if err != nil { t.Fatal(err) } // Bucket cap: a full bucket refuses the next release even when the // pool's element budget has room. gradPool.Lock() savedB, savedE := gradPool.buckets[k], gradPool.elems gradPool.buckets[k] = make([]*core.Array, gradPoolPerBucket) gradPool.elems = gradPoolMaxElems - 16 gradPool.Unlock() releaseGrad(keep, []int{8}) gradPool.Lock() gotLen := len(gradPool.buckets[k]) gradPool.buckets[k], gradPool.elems = savedB, savedE gradPool.Unlock() if gotLen != gradPoolPerBucket { t.Fatalf("bucket accepted a release past its cap: %d entries, cap %d", gotLen, gradPoolPerBucket) } // Element cap: a full pool refuses the next release however empty // the bucket is. gradPool.Lock() gradPool.buckets[k] = nil gradPool.elems = gradPoolMaxElems gradPool.Unlock() releaseGrad(keep, []int{8}) gradPool.Lock() gotElems := gradPool.elems gradPool.buckets[k], gradPool.elems = savedB, savedE gradPool.Unlock() if gotElems != gradPoolMaxElems { t.Fatalf("pool accepted a release past its element cap: %d, cap %d", gotElems, gradPoolMaxElems) } }