353 lines
9.8 KiB
Go
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)
|
|
}
|