// Copyright (c) 2026 Petr BalvĂ­n (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) }