Files
tensor/grad/pool.go
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

353 lines
9.8 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"sync"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Gradient buffer recycling for the backward sweep. A sweep allocates
// one gradient array per node output and per folded contribution; the
// arrays die within the sweep that made them, except the ones committed
// to leaves, which escape to the caller. The pool reclaims the
// intermediates: a borrowed array arrives with a fully zeroed payload,
// the pool retains a bounded number of elements, and nothing is
// recycled while a live reference to it exists. That last rule rests on
// an invariant the closures maintain: a backward closure never returns
// the incoming gradient buffer itself, never returns one buffer in two
// slots and never returns a view of another live array, so an entry the
// sweep releases is unreachable from the graph, from the returned map
// and from every other entry.
//
// An array the pool declines (an exotic dtype, a rank above six, a
// payload above the cap, a full bucket) falls back to ordinary
// allocation; correctness never depends on a hit.
// gradSlot is one gradient array together with the shape it was built
// for, which is the pool's reuse key. Slots travel instead of bare
// arrays so the sweep can release a buffer without re-deriving its
// shape, which would cost an allocation of its own.
type gradSlot struct {
arr *core.Array
sh []int
}
// poolKey identifies a gradient buffer exactly: dtype, rank, element
// count and dimensions. The fixed dimension array keeps the key
// comparable, so buckets need no stored shape for matching.
type poolKey struct {
dt core.Dtype
nd int8
n int
d [6]int
}
// poolKeyOf builds the key for dt and shape, reporting false for
// anything the pool does not accept.
func poolKeyOf(dt core.Dtype, shape []int) (poolKey, bool) {
var k poolKey
if len(shape) == 0 || len(shape) > len(k.d) {
return k, false
}
switch dt {
case core.Float, core.Float32, core.Complex:
default:
return k, false
}
n := 1
for i, d := range shape {
n *= d
k.d[i] = d
}
if n > gradPoolMaxArrayElems {
return k, false
}
k.dt, k.nd, k.n = dt, int8(len(shape)), n
return k, true
}
// The retention caps: no more than gradPoolMaxArrayElems elements in
// one array, gradPoolMaxElems retained across all buckets and
// gradPoolPerBucket arrays of one exact shape. A buffer outside the
// caps is dropped to the garbage collector instead of retained, so the
// pool cannot pin memory beyond these bounds however hard one workload
// pushes it.
const (
gradPoolMaxArrayElems = 1 << 20
gradPoolMaxElems = 1 << 20
gradPoolPerBucket = 32
)
var gradPool = struct {
sync.Mutex
buckets map[poolKey][]*core.Array
elems int
}{buckets: make(map[poolKey][]*core.Array)}
// freshGrad allocates a zeroed array of k's dtype and shape, taking
// ownership of the payload the way the constructor documents: grad
// writes that payload only through the array's own raw accessor.
func freshGrad(k poolKey, shape []int) *core.Array {
var a *core.Array
switch k.dt {
case core.Float:
a, _ = core.FloatsFromArray(make([]float64, k.n), shape...)
case core.Float32:
a, _ = core.FromFloat32Slice(make([]float32, k.n), shape...)
default:
a, _ = core.ComplexFromArray(make([]complex128, k.n), shape...)
}
if a == nil {
a, _ = core.Zeros(k.dt, shape...)
}
return a
}
// clearGradPayload zeroes every slot of a's payload, the borrow-side
// rule: a recycled buffer must never carry the previous sweep's values
// into a reader.
func clearGradPayload(a *core.Array) {
switch a.Dtype() {
case core.Float:
clear(a.RawFloats())
case core.Float32:
clear(a.RawFloat32s())
case core.Complex:
clear(a.RawComplexes())
}
}
// releaseGrad offers a dead gradient array to the pool. The caller must
// have proven the array unreachable: releasing one a tape, a graph or a
// caller still holds would let the next borrower corrupt it. The shape
// must be the shape the array was built with; a mismatch is caught by
// the element-count check and drops the array instead of pooling it.
func releaseGrad(a *core.Array, shape []int) {
if a == nil || a.Strided() {
return
}
k, ok := poolKeyOf(a.Dtype(), shape)
if !ok || k.n != a.Len() {
return
}
gradPool.Lock()
b := gradPool.buckets[k]
if len(b) >= gradPoolPerBucket || gradPool.elems+k.n > gradPoolMaxElems {
gradPool.Unlock()
return
}
gradPool.buckets[k] = append(b, a)
gradPool.elems += k.n
gradPool.Unlock()
}
// gradArena is one sweep's private free list. A sweep borrows and
// releases in near-LIFO order, so most round trips stay on the calling
// goroutine under no lock; the global pool absorbs overflow and
// supplies misses, and a sweep-end flush returns the leftovers under a
// single lock. Every sweep owns its arena, so concurrent sweeps on
// different graphs never share one.
type gradArena struct {
free []gradFree
elems int
pooled bool
}
type gradFree struct {
arr *core.Array
sh []int
k poolKey
}
// gradArenaMaxFree bounds what one arena carries between flushes; a
// sweep that releases beyond it hands the surplus to the global pool.
const gradArenaMaxFree = 256
var gradArenaPool = sync.Pool{New: func() any { return &gradArena{} }}
func borrowArena() *gradArena {
ar := gradArenaPool.Get().(*gradArena)
ar.pooled = true
// A recycled arena comes back with its previous free list: the
// slice is emptied here, or stale entries would both starve the
// scan and pin the buffers they still name.
ar.free = ar.free[:0]
ar.elems = 0
return ar
}
// borrowGrad returns a zeroed array of dt and shape: the arena's own
// free list first, then the global pool, then fresh allocation. A nil
// arena means the legacy sweep path, which allocates exactly what it
// allocated before and takes no part in the pool.
func (ar *gradArena) borrowGrad(dt core.Dtype, shape []int) *core.Array {
if ar == nil {
a, _ := core.Zeros(dt, shape...)
return a
}
k, ok := poolKeyOf(dt, shape)
if !ok {
a, _ := core.Zeros(dt, shape...)
return a
}
for i := len(ar.free) - 1; i >= 0; i-- {
e := ar.free[i]
if e.k != k {
continue
}
ar.free[i] = ar.free[len(ar.free)-1]
ar.free = ar.free[:len(ar.free)-1]
ar.elems -= k.n
clearGradPayload(e.arr)
return e.arr
}
gradPool.Lock()
b := gradPool.buckets[k]
if len(b) > 0 {
a := b[len(b)-1]
gradPool.buckets[k] = b[:len(b)-1]
gradPool.elems -= k.n
gradPool.Unlock()
clearGradPayload(a)
return a
}
gradPool.Unlock()
return freshGrad(k, shape)
}
// releaseGrad returns a dead gradient array to the arena, falling
// through to the global pool when the arena is full. A nil arena means
// the caller owns the lifetime, so the array is left to the collector.
func (ar *gradArena) releaseGrad(a *core.Array, shape []int) {
if ar == nil || a == nil || a.Strided() {
return
}
k, ok := poolKeyOf(a.Dtype(), shape)
if !ok || k.n != a.Len() {
return
}
if len(ar.free) >= gradArenaMaxFree || ar.elems+k.n > gradPoolMaxElems {
releaseGrad(a, shape)
return
}
ar.free = append(ar.free, gradFree{arr: a, sh: shape, k: k})
ar.elems += k.n
}
// flush returns everything the arena still holds to the global pool.
// Buffers the pool declines are dropped to the collector; the arena
// itself returns to the sync.Pool for the next sweep.
func (ar *gradArena) flush() {
if ar == nil {
return
}
gradPool.Lock()
for _, e := range ar.free {
b := gradPool.buckets[e.k]
if len(b) >= gradPoolPerBucket || gradPool.elems+e.k.n > gradPoolMaxElems {
continue
}
gradPool.buckets[e.k] = append(b, e.arr)
gradPool.elems += e.k.n
}
gradPool.Unlock()
clear(ar.free)
ar.free = ar.free[:0]
ar.elems = 0
if ar.pooled {
gradArenaPool.Put(ar)
}
}
// fillGradSlotC writes z into every complex element of s.
func fillGradSlotC(s gradSlot, z complex128) {
if s.arr == nil {
return
}
cs := s.arr.RawComplexes()[:s.arr.Len()]
for i := range cs {
cs[i] = z
}
}
// fillGradSlot writes v into every element of s, the seed and fill
// helper. Each dtype takes the same spelling fillConst writes: a
// float32 destination narrows the constant once and stores it, a
// complex one stores complex(v, 0).
func fillGradSlot(s gradSlot, v float64) {
if s.arr == nil {
return
}
n := s.arr.Len()
switch s.arr.Dtype() {
case core.Float:
fs := s.arr.RawFloats()[:n]
for i := range fs {
fs[i] = v
}
case core.Float32:
fs := s.arr.RawFloat32s()[:n]
fv := float32(v)
for i := range fs {
fs[i] = fv
}
case core.Complex:
cs := s.arr.RawComplexes()[:n]
z := complex(v, 0)
for i := range cs {
cs[i] = z
}
}
}
// tapeFrame is one node's position in the reverse sweep's explicit
// walk: the node being expanded and the next operand index to visit.
type tapeFrame struct {
node *gradNode
next int
}
// tapeWork is the sweep's traversal scratch: the topological order, the
// walk stack and the seen set. It is borrowed per sweep and returned
// with its references cleared, so a pooled copy never pins a dead
// graph; a workload whose graph exceeds the retention cap drops the
// buffers to the collector instead of pinning them.
type tapeWork struct {
order []*gradNode
stack []tapeFrame
seen map[*Tensor]bool
}
const gradPoolMaxTapeNodes = 1 << 16
var tapeWorkPool = sync.Pool{New: func() any {
return &tapeWork{seen: make(map[*Tensor]bool)}
}}
func borrowTapeWork() *tapeWork {
w := tapeWorkPool.Get().(*tapeWork)
w.order = w.order[:0]
w.stack = w.stack[:0]
clear(w.seen)
if cap(w.order) > gradPoolMaxTapeNodes || cap(w.stack) > gradPoolMaxTapeNodes {
return &tapeWork{seen: make(map[*Tensor]bool)}
}
return w
}
func releaseTapeWork(w *tapeWork) {
if w == nil {
return
}
clear(w.order[:cap(w.order)])
clear(w.stack[:cap(w.stack)])
clear(w.seen)
if cap(w.order) > gradPoolMaxTapeNodes || cap(w.stack) > gradPoolMaxTapeNodes {
return
}
tapeWorkPool.Put(w)
}