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