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