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