// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "fmt" "math" "slices" "strconv" "strings" "sourcedock.dev/petrbalvin/tensor/internal/core" "sourcedock.dev/petrbalvin/tensor/internal/engine" ) // Automatic differentiation. Tensor wraps an immutable Array // with reverse-mode autograd: every differentiable operation records a // node in the computation graph, and Backward propagates gradients from // the output to every leaf that requires them. This is the training // substrate on top of the Array surface. // // Contract: // - float, float32 and complex tensors are differentiable; int // tensors error loudly. Backward needs a real-valued loss: a // complex loss is an error, and complex leaves receive the // conjugate-Wirtinger gradient dL/dz̄, the coefficient g of // dL = 2·Re(g·dz). // - Gradients carry the same dtype as the data they update; a // mixed-dtype graph narrows each leaf gradient to the leaf's dtype. // - Backward accumulates into existing leaf gradients; call ZeroGrad // before the next backward. // - The graph is rebuilt on every forward pass; tensors are cheap // values and never alias the arrays they wrap. // Tensor is a differentiable n-dimensional array. type Tensor struct { data *core.Array grad *core.Array reqGrad bool node *gradNode } // gradNode is one recorded operation in the reverse-mode graph. The // operands are held inline, so the tape pays no second allocation for a // one- or two-element input list, and the backward closures read the // operand arrays they were handed at forward time: Tensor.data is // mutable through ReplaceWith, so nothing may be read off an operand at // backward time. A node carries no sweep state: everything a reverse // pass needs lives in locals, so two passes over one graph never share // a cell and never see each other's stamps. type gradNode struct { op string grad func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error in [2]*Tensor out *Tensor arity int8 } // opResult is a differentiable result and its graph node in one // allocation: a node lives exactly as long as the tensor carrying it, // so the tape pays one allocation per operation instead of two. type opResult struct { t Tensor n gradNode } // FromArray wraps a as a tensor; requiresGrad marks it as a leaf whose // gradient Backward fills. func FromArray(a *core.Array, requiresGrad bool) *Tensor { return &Tensor{data: a, reqGrad: requiresGrad} } // FromFloat64s builds a float tensor from vals and a shape, mirroring // FromFloats. func FromFloat64s(vals []float64, requiresGrad bool, shape ...int) (*Tensor, error) { a, err := core.FromFloats(vals, shape...) if err != nil { return nil, err } return FromArray(a, requiresGrad), nil } // Data returns the underlying array. func (t *Tensor) Data() *core.Array { return t.data } // Grad returns the accumulated gradient, or nil before Backward. func (t *Tensor) Grad() *core.Array { return t.grad } // RequiresGrad reports whether the tensor is a trainable leaf. func (t *Tensor) RequiresGrad() bool { return t.reqGrad } // ZeroGrad discards the accumulated gradient. func (t *Tensor) ZeroGrad() { t.grad = nil } // Backward propagates gradients from t back to every leaf that requires // them, accumulating into the leaves' existing gradients. It seeds the // output with ones, so a scalar loss is the usual caller. The seed is // real: a complex output means the "loss" is not a scalar objective, // and Backward rejects it with an error pointing at Real, Abs or Abs2. // The reverse sweep recycles its intermediate gradient buffers; the // arrays committed here escape that recycling, because from Grad on // they belong to the caller. func (t *Tensor) Backward() error { grads, err := t.reverseGradsPooled() if err != nil { return err } // Commit the accumulated gradients into the leaves, merging with any // gradient left over from an earlier Backward. The sweep returned // leaf entries only. for leaf, s := range grads { if s.arr == nil || !leaf.reqGrad { continue } if leaf.grad == nil { leaf.grad = s.arr continue } acc, err := core.Add(leaf.grad, s.arr) if err != nil { return err } releaseGrad(s.arr, s.sh) leaf.grad = acc } return nil } // reverseGrads runs the reverse pass over t's graph and returns the // gradient every reached tensor receives from this pass alone. Nothing // is committed: the pass neither reads nor writes any tensor's // accumulated Grad, so a caller that differentiates the graph for its // own purposes (the second-order helpers, whose objective closure may // touch the caller's trainable tensors) leaves the graph's gradients // exactly as it found them. // // The sweep walks the graph once to fix the topological order and // again in reverse, folding each node's contributions into its inputs. // Both orders are the ones a recursive children-first walk produces: // the fold order decides the rounding of an accumulation, so it is // part of the result and not an implementation detail. Everything the // pass needs is a local: two passes over one graph, concurrent // included, cannot see each other's state. func (t *Tensor) reverseGrads() (map[*Tensor]*core.Array, error) { if !t.reqGrad { return nil, errf("Backward: tensor does not require grad") } if isComplexArr(t.data) { return nil, errf("Backward: the loss must be real-valued; reduce the complex result with Real, Imag, Abs or Abs2 first") } ones, err := core.Ones(t.data.Dtype(), t.data.Shape()...) if err != nil { return nil, err } if t.node == nil { // A leaf tensor's gradient w.r.t. itself is ones, accumulated // like any other backward pass. return map[*Tensor]*core.Array{t: ones}, nil } type frame struct { node *gradNode next int } // Topological order of the graph, children first: an explicit // post-order walk that marks a node when it is first reached and // appends it once its operands are done, the order the recursive // walk produced. seen := make(map[*Tensor]bool) order := make([]*gradNode, 0, 64) stack := make([]frame, 0, 64) seen[t] = true stack = append(stack, frame{node: t.node}) for len(stack) > 0 { f := &stack[len(stack)-1] if f.next < int(f.node.arity) { in := f.node.in[f.next] f.next++ if in.node != nil && !seen[in] { seen[in] = true stack = append(stack, frame{node: in.node}) } continue } order = append(order, f.node) stack = stack[:len(stack)-1] } grads := map[*Tensor]*core.Array{t: ones} var slots [2]gradSlot for _, node := range slices.Backward(order) { gOut, ok := grads[node.out] if !ok { continue } for i := range slots { slots[i] = gradSlot{} } // A nil arena: this path hands the map to callers who read it // after the sweep, so its buffers are plain allocations the // collector reclaims, never pooled ones. if err := node.grad(gradSlot{arr: gOut}, &slots, nil); err != nil { return nil, err } for j := range int(node.arity) { in := node.in[j] gj := slots[j].arr if !in.reqGrad || gj == nil { continue } if gj.Dtype() != in.data.Dtype() { sl, nerr := narrowGradient(gradSlot{arr: gj}, in.data.Dtype()) if nerr != nil { return nil, nerr } gj = sl.arr if gj == nil { continue } } // Accumulate into the incoming-gradient map: a tensor used // twice (like z in z·z) receives both contributions. if prev, ok := grads[in]; ok { var aerr error gj, aerr = core.Add(prev, gj) if aerr != nil { return nil, aerr } } grads[in] = gj } } return grads, nil } // reverseGradsPooled is Backward's own reverse pass: the same walk, the // same fold order and the same per-element arithmetic as reverseGrads, // with the gradient buffers borrowed from the sweep's arena. The // returned map holds leaf entries only; every intermediate gradient is // released the moment its producing node has consumed it. That release // is safe because the closures maintain one invariant: a closure // returns freshly built buffers in its slots, never the incoming // gradient itself, never one buffer in two slots, and never a view of // another live array, so at the producing node's fold nothing but the // map entry itself can still reach the buffer. A leaf entry is never // released here: Backward commits it to the leaf, where it escapes the // pool into caller ownership. An error path commits nothing to any // tensor: the caller's gradients are exactly as they were. Buffers // already released to the arena return through its flush, where the // pool zeroes every one of them at the next borrow; buffers still // referenced when the error surfaces go to the collector unreleased. func (t *Tensor) reverseGradsPooled() (map[*Tensor]gradSlot, error) { if !t.reqGrad { return nil, errf("Backward: tensor does not require grad") } if isComplexArr(t.data) { return nil, errf("Backward: the loss must be real-valued; reduce the complex result with Real, Imag, Abs or Abs2 first") } ar := borrowArena() seed := gradSlot{sh: t.data.Shape()} switch dt := t.data.Dtype(); dt { case core.Float, core.Float32, core.Complex: seed.arr = ar.borrowGrad(dt, seed.sh) if dt == core.Complex { fillGradSlotC(seed, complex(1, 0)) } else { fillGradSlot(seed, 1) } default: ones, err := core.Ones(t.data.Dtype(), seed.sh...) if err != nil { ar.flush() return nil, err } seed.arr = ones } if t.node == nil { ar.flush() return map[*Tensor]gradSlot{t: seed}, nil } w := borrowTapeWork() fail := func(err error) (map[*Tensor]gradSlot, error) { releaseTapeWork(w) ar.flush() return nil, err } // The same explicit children-first post-order walk reverseGrads // runs, on the borrowed traversal scratch. w.seen[t] = true w.stack = append(w.stack, tapeFrame{node: t.node}) for len(w.stack) > 0 { f := &w.stack[len(w.stack)-1] if f.next < int(f.node.arity) { in := f.node.in[f.next] f.next++ if in.node != nil && !w.seen[in] { w.seen[in] = true w.stack = append(w.stack, tapeFrame{node: in.node}) } continue } w.order = append(w.order, f.node) w.stack = w.stack[:len(w.stack)-1] } grads := map[*Tensor]gradSlot{t: seed} var slots [2]gradSlot for _, node := range slices.Backward(w.order) { gOut, ok := grads[node.out] if !ok { continue } for i := range slots { slots[i] = gradSlot{} } if err := node.grad(gOut, &slots, ar); err != nil { return fail(err) } // The producing node has consumed its output gradient, and no // closure returned the buffer itself, so this entry was the // last live reference. delete(grads, node.out) ar.releaseGrad(gOut.arr, gOut.sh) for j := range int(node.arity) { in := node.in[j] s := slots[j] if s.arr == nil { continue } if !in.reqGrad { ar.releaseGrad(s.arr, s.sh) continue } gj := s if gj.arr.Dtype() != in.data.Dtype() { nj, nerr := narrowGradient(gj, in.data.Dtype()) if nerr != nil { return fail(nerr) } if nj.arr == gj.arr { gj = nj } else { ar.releaseGrad(gj.arr, gj.sh) gj = nj } } // Accumulate into the incoming-gradient map: a tensor used // twice (like z in z·z) receives both contributions, folded // in the same order the walk has always produced. if prev, ok := grads[in]; ok { acc, aerr := addGradSlots(ar, prev, gj) if aerr != nil { return fail(aerr) } ar.releaseGrad(prev.arr, prev.sh) ar.releaseGrad(gj.arr, gj.sh) grads[in] = acc continue } grads[in] = gj } } releaseTapeWork(w) // Intermediates were released at their producing nodes; the sweep // below keeps the leaf-only return invariant explicit rather than // trusted. for k, s := range grads { if k.node != nil { ar.releaseGrad(s.arr, s.sh) delete(grads, k) } } ar.flush() return grads, nil } // gradShape returns the shape the incoming gradient carries, or, on // the legacy sweep path whose slots travel without a shape, a fresh // copy of a's own. For a shape-preserving op the two are equal, which // is what Add's and Mul's use of g.sh already rests on. func gradShape(g gradSlot, a *core.Array) []int { if g.sh != nil { return g.sh } return a.Shape() } // elemSweepMin is the element count below which the new gradient // helpers stay on the calling goroutine: below the element-wise spawn // floor a parallel split cannot pay for its own closure, and the // helpers' buffers at that size are the tape's small tensors. const elemSweepMin = 1024 // copyGradSlot returns an independent buffer holding g's values, // shaped sh. The copy is what keeps every gradient entry exclusively // owned, which in turn is what lets the sweep recycle an entry the // moment its producing node is done: a passthrough of the incoming // buffer would leave a second live reference and break that proof. The // same-dtype dense path writes the copy element for element, the same // values Astype's clone produced; anything else keeps Astype itself. func copyGradSlot(ar *gradArena, g gradSlot, sh []int) (gradSlot, error) { if g.arr == nil { return gradSlot{}, nil } dt := g.arr.Dtype() if g.arr.Strided() || len(sh) == 0 || len(sh) > 6 { c, err := core.Astype(g.arr, dt) if err != nil { return gradSlot{}, err } return gradSlot{arr: c, sh: sh}, nil } n := g.arr.Len() var out *core.Array switch dt { case core.Float, core.Float32, core.Complex: out = ar.borrowGrad(dt, sh) default: c, err := core.Astype(g.arr, dt) if err != nil { return gradSlot{}, err } return gradSlot{arr: c, sh: sh}, nil } switch dt { case core.Float: copy(out.RawFloats()[:n], g.arr.RawFloats()[:n]) case core.Float32: copy(out.RawFloat32s()[:n], g.arr.RawFloat32s()[:n]) default: copy(out.RawComplexes()[:n], g.arr.RawComplexes()[:n]) } return gradSlot{arr: out, sh: sh}, nil } // negGradSlot returns an independent buffer holding g negated. Each // dtype multiplies by the negated one in the expression MulI writes, // so the bits are MulI(g, -1)'s own. func negGradSlot(ar *gradArena, g gradSlot, sh []int) (gradSlot, error) { if g.arr == nil { return gradSlot{}, nil } dt := g.arr.Dtype() if g.arr.Strided() || len(sh) == 0 || len(sh) > 6 { c := core.MulI(g.arr, -1) return gradSlot{arr: c, sh: sh}, nil } n := g.arr.Len() var out *core.Array switch dt { case core.Float, core.Float32, core.Complex: out = ar.borrowGrad(dt, sh) default: c := core.MulI(g.arr, -1) return gradSlot{arr: c, sh: sh}, nil } switch dt { case core.Float: gs, os := g.arr.RawFloats()[:n], out.RawFloats()[:n] for i := range os { os[i] = gs[i] * -1 } case core.Float32: gs, os := g.arr.RawFloat32s()[:n], out.RawFloat32s()[:n] for i := range os { os[i] = float32(float64(gs[i]) * -1) } default: gs, os := g.arr.RawComplexes()[:n], out.RawComplexes()[:n] for i := range os { os[i] = gs[i] * complex(-1, 0) } } return gradSlot{arr: out, sh: sh}, nil } // addGradSlots folds two contributions to one tensor's gradient into a // fresh buffer. Both carry the tensor's gradient dtype and shape, so // the dense path writes core.Add's own per-element sum into a borrowed // buffer; anything else keeps core.Add unchanged. func addGradSlots(ar *gradArena, p, q gradSlot) (gradSlot, error) { if p.arr == nil { return q, nil } if q.arr == nil { return p, nil } dt := p.arr.Dtype() if dt != q.arr.Dtype() || p.arr.Strided() || q.arr.Strided() || len(p.sh) == 0 || len(p.sh) > 6 { out, err := core.Add(p.arr, q.arr) if err != nil { return gradSlot{}, err } return gradSlot{arr: out, sh: p.sh}, nil } n := p.arr.Len() out := ar.borrowGrad(dt, p.sh) switch dt { case core.Float: os, as, bs := out.RawFloats()[:n], p.arr.RawFloats()[:n], q.arr.RawFloats()[:n] if n < elemSweepMin { for i := range os { os[i] = as[i] + bs[i] } break } elemPass(n, func(s, e int) { for i := s; i < e; i++ { os[i] = as[i] + bs[i] } }) case core.Float32: os, as, bs := out.RawFloat32s()[:n], p.arr.RawFloat32s()[:n], q.arr.RawFloat32s()[:n] if n < elemSweepMin { for i := range os { os[i] = float32(float64(as[i]) + float64(bs[i])) } break } elemPass(n, func(s, e int) { for i := s; i < e; i++ { os[i] = float32(float64(as[i]) + float64(bs[i])) } }) default: os, as, bs := out.RawComplexes()[:n], p.arr.RawComplexes()[:n], q.arr.RawComplexes()[:n] if n < elemSweepMin { for i := range os { os[i] = as[i] + bs[i] } break } elemPass(n, func(s, e int) { for i := s; i < e; i++ { os[i] = as[i] + bs[i] } }) } return gradSlot{arr: out, sh: p.sh}, nil } // mulGradSlot writes g⊙b into a borrowed buffer shaped sh: the same // per-element product core.Mul computes, dtype for dtype. The second // return is false when the operands leave the dense same-dtype paths, // and the caller keeps core.Mul. func mulGradSlot(ar *gradArena, g gradSlot, b *core.Array, sh []int) (gradSlot, bool) { if g.arr == nil || b == nil || g.arr.Strided() || b.Strided() || len(sh) == 0 || len(sh) > 6 || g.arr.Len() != b.Len() { return gradSlot{}, false } dt := g.arr.Dtype() if dt != b.Dtype() { return gradSlot{}, false } n := g.arr.Len() switch dt { case core.Float: out := ar.borrowGrad(dt, sh) gs, bs, os := g.arr.RawFloats()[:n], b.RawFloats()[:n], out.RawFloats()[:n] if n < elemSweepMin { for i := range os { os[i] = gs[i] * bs[i] } } else { elemPass(n, func(s, e int) { for i := s; i < e; i++ { os[i] = gs[i] * bs[i] } }) } return gradSlot{arr: out, sh: sh}, true case core.Float32: out := ar.borrowGrad(dt, sh) gs, bs, os := g.arr.RawFloat32s()[:n], b.RawFloat32s()[:n], out.RawFloat32s()[:n] if n < elemSweepMin { for i := range os { os[i] = float32(float64(gs[i]) * float64(bs[i])) } } else { elemPass(n, func(s, e int) { for i := s; i < e; i++ { os[i] = float32(float64(gs[i]) * float64(bs[i])) } }) } return gradSlot{arr: out, sh: sh}, true case core.Complex: out := ar.borrowGrad(dt, sh) gs, bs, os := g.arr.RawComplexes()[:n], b.RawComplexes()[:n], out.RawComplexes()[:n] for i := range os { os[i] = gs[i] * bs[i] } return gradSlot{arr: out, sh: sh}, true } return gradSlot{}, false } // mulConjGradSlot writes g·conj(b) into a borrowed buffer shaped sh, // the complex product rule's adjoint. The conjugate is formed per // element exactly as conjArray forms it, and the product is the same // one core.Mul computes on the conjugated operand. func mulConjGradSlot(ar *gradArena, g gradSlot, b *core.Array, sh []int) (gradSlot, bool) { if g.arr == nil || b == nil || g.arr.Strided() || b.Strided() || len(sh) == 0 || len(sh) > 6 || g.arr.Len() != b.Len() || g.arr.Dtype() != core.Complex || b.Dtype() != core.Complex { return gradSlot{}, false } n := g.arr.Len() out := ar.borrowGrad(core.Complex, sh) gs, bs, os := g.arr.RawComplexes()[:n], b.RawComplexes()[:n], out.RawComplexes()[:n] for i := range os { z := bs[i] os[i] = gs[i] * complex(real(z), -imag(z)) } return gradSlot{arr: out, sh: sh}, true } // divGradSlot writes g/b into a borrowed buffer shaped sh with the // per-element quotient core.Div computes. The second return is false // outside the dense same-dtype real and complex paths. func divGradSlot(ar *gradArena, g gradSlot, b *core.Array, sh []int) (gradSlot, bool) { if g.arr == nil || b == nil || g.arr.Strided() || b.Strided() || len(sh) == 0 || len(sh) > 6 || g.arr.Len() != b.Len() { return gradSlot{}, false } dt := g.arr.Dtype() if dt != b.Dtype() { return gradSlot{}, false } n := g.arr.Len() switch dt { case core.Float: out := ar.borrowGrad(dt, sh) gs, bs, os := g.arr.RawFloats()[:n], b.RawFloats()[:n], out.RawFloats()[:n] if n < elemSweepMin { for i := range os { os[i] = gs[i] / bs[i] } } else { elemPass(n, func(s, e int) { for i := s; i < e; i++ { os[i] = gs[i] / bs[i] } }) } return gradSlot{arr: out, sh: sh}, true case core.Float32: out := ar.borrowGrad(dt, sh) gs, bs, os := g.arr.RawFloat32s()[:n], b.RawFloat32s()[:n], out.RawFloat32s()[:n] if n < elemSweepMin { for i := range os { os[i] = float32(float64(gs[i]) / float64(bs[i])) } } else { elemPass(n, func(s, e int) { for i := s; i < e; i++ { os[i] = float32(float64(gs[i]) / float64(bs[i])) } }) } return gradSlot{arr: out, sh: sh}, true case core.Complex: out := ar.borrowGrad(dt, sh) gs, bs, os := g.arr.RawComplexes()[:n], b.RawComplexes()[:n], out.RawComplexes()[:n] for i := range os { os[i] = gs[i] / bs[i] } return gradSlot{arr: out, sh: sh}, true } return gradSlot{}, false } // divGradReal writes the real Div backward, ga = g/b and gb = −g·a/b², // into two borrowed buffers shaped sh. gb's sub-expressions run in the // order the staged core chain produced them in (neg = g·(−1), // num = neg·a, bb = b·b, gb = num/bb), so every element rounds exactly // as the chain rounded it. The second return is false outside the // dense float64 and float32 paths. func divGradReal(ar *gradArena, g gradSlot, a, b *core.Array, sh []int) (gradSlot, gradSlot, bool) { if g.arr == nil || a == nil || b == nil || g.arr.Strided() || a.Strided() || b.Strided() || len(sh) == 0 || len(sh) > 6 || g.arr.Len() != a.Len() || g.arr.Len() != b.Len() { return gradSlot{}, gradSlot{}, false } dt := g.arr.Dtype() if dt != a.Dtype() || dt != b.Dtype() || (dt != core.Float && dt != core.Float32) { return gradSlot{}, gradSlot{}, false } n := g.arr.Len() ga := ar.borrowGrad(dt, sh) gb := ar.borrowGrad(dt, sh) if dt == core.Float { gs, as, bs := g.arr.RawFloats()[:n], a.RawFloats()[:n], b.RawFloats()[:n] da, db := ga.RawFloats()[:n], gb.RawFloats()[:n] if n < elemSweepMin { for i := range n { da[i] = gs[i] / bs[i] neg := gs[i] * -1 num := neg * as[i] bb := bs[i] * bs[i] db[i] = num / bb } return gradSlot{arr: ga, sh: sh}, gradSlot{arr: gb, sh: sh}, true } elemPass(n, func(s, e int) { for i := s; i < e; i++ { da[i] = gs[i] / bs[i] neg := gs[i] * -1 num := neg * as[i] bb := bs[i] * bs[i] db[i] = num / bb } }) return gradSlot{arr: ga, sh: sh}, gradSlot{arr: gb, sh: sh}, true } gs, as, bs := g.arr.RawFloat32s()[:n], a.RawFloat32s()[:n], b.RawFloat32s()[:n] da, db := ga.RawFloat32s()[:n], gb.RawFloat32s()[:n] if n < elemSweepMin { for i := range n { da[i] = float32(float64(gs[i]) / float64(bs[i])) neg := float32(float64(gs[i]) * -1) num := float32(float64(neg) * float64(as[i])) bb := float32(float64(bs[i]) * float64(bs[i])) db[i] = float32(float64(num) / float64(bb)) } return gradSlot{arr: ga, sh: sh}, gradSlot{arr: gb, sh: sh}, true } elemPass(n, func(s, e int) { for i := s; i < e; i++ { da[i] = float32(float64(gs[i]) / float64(bs[i])) neg := float32(float64(gs[i]) * -1) num := float32(float64(neg) * float64(as[i])) bb := float32(float64(bs[i]) * float64(bs[i])) db[i] = float32(float64(num) / float64(bb)) } }) return gradSlot{arr: ga, sh: sh}, gradSlot{arr: gb, sh: sh}, true } // transposeGradSlot writes the transpose of the rank-2 dense g into a // borrowed buffer shaped sh: pure data motion, the same elements // core.Transpose copies, only into recycled storage. The caller shapes // g as (sh[1], sh[0]), the transpose node's own output shape, so out // element (i, j) is g element (j, i). The second return is false // elsewhere and the caller keeps core.Transpose. func transposeGradSlot(ar *gradArena, g gradSlot, sh []int) (gradSlot, bool) { if g.arr == nil || g.arr.Strided() || g.arr.NDim() != 2 || len(sh) != 2 { return gradSlot{}, false } dt := g.arr.Dtype() switch dt { case core.Float, core.Float32, core.Complex: default: return gradSlot{}, false } m, n := sh[0], sh[1] if g.arr.Len() != m*n { return gradSlot{}, false } out := ar.borrowGrad(dt, sh) switch dt { case core.Float: gs, os := g.arr.RawFloats(), out.RawFloats() for i := range m { for j := range n { os[i*n+j] = gs[j*m+i] } } case core.Float32: gs, os := g.arr.RawFloat32s(), out.RawFloat32s() for i := range m { for j := range n { os[i*n+j] = gs[j*m+i] } } default: gs, os := g.arr.RawComplexes(), out.RawComplexes() for i := range m { for j := range n { os[i*n+j] = gs[j*m+i] } } } return gradSlot{arr: out, sh: sh}, true } // zeros allocates a zeroed array, ignoring the constructor's error: // every shape here comes from existing arrays, so it cannot be invalid. func zeros(dt core.Dtype, shape []int) *core.Array { a, _ := core.Zeros(dt, shape...) return a } // checkFloat rejects non-float tensors for differentiation. func (t *Tensor) checkFloat(name string) error { if t.data.Dtype() != core.Float && t.data.Dtype() != core.Float32 { return errf("autograd: %s needs a float or float32 tensor, got %s", name, t.data.Dtype()) } return nil } // unaryResult wraps out as the result of a one-input op on t and, when // the graph reaches t, records the backward node. The grad closure must // capture the operand arrays it needs: Tensor.data is mutable through // ReplaceWith, so nothing may be read off t at backward time. The // closure writes its answers into dst and never returns the incoming // gradient buffer itself: the sweep recycles buffers on the strength of // that invariant. Tensor and node share one allocation. func (t *Tensor) unaryResult(op string, out *core.Array, grad func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error) *Tensor { if !t.reqGrad { // No graph reaches this result, so it needs no node. return &Tensor{data: out} } res := &opResult{t: Tensor{data: out, reqGrad: true}} res.n = gradNode{op: op, grad: grad, in: [2]*Tensor{t}, out: &res.t, arity: 1} res.t.node = &res.n return &res.t } // binaryResult is unaryResult for a two-input op over t and u. func binaryResult(op string, t, u *Tensor, out *core.Array, grad func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error) *Tensor { if !t.reqGrad && !u.reqGrad { return &Tensor{data: out} } res := &opResult{t: Tensor{data: out, reqGrad: true}} res.n = gradNode{op: op, grad: grad, in: [2]*Tensor{t, u}, out: &res.t, arity: 2} res.t.node = &res.n return &res.t } // binary is the shared plumbing for element-wise differentiable ops. // The four arithmetic ops all differentiate complex input, so the gate // here is checkDiff. func (t *Tensor) binary(name string, u *Tensor, forward func(a, b *core.Array) (*core.Array, error), backward func(g gradSlot, a, b *core.Array, dst *[2]gradSlot, ar *gradArena) error, ) (*Tensor, error) { if err := t.checkDiff(name); err != nil { return nil, err } if err := u.checkDiff(name); err != nil { return nil, err } out, err := forward(t.data, u.data) if err != nil { return nil, err } at, au := t.data, u.data return binaryResult(name, t, u, out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { return backward(g, at, au, dst, ar) }), nil } // Add returns t + u. func (t *Tensor) Add(u *Tensor) (*Tensor, error) { return t.binary("Add", u, core.Add, func(g gradSlot, _, _ *core.Array, dst *[2]gradSlot, ar *gradArena) error { // The gradient is the seed for both inputs, and the two leaves // must not share one buffer: a shared instance would make a // write through one leaf's Grad corrupt the other's, and the // commits into leaf.grad would alias. Both answers are // independent copies with identical bits, which also keeps each // gradient entry exclusively owned for the sweep's recycling. sh := g.sh c0, err := copyGradSlot(ar, g, sh) if err != nil { return err } c1, err := copyGradSlot(ar, g, sh) if err != nil { return err } dst[0], dst[1] = c0, c1 return nil }) } // Sub returns t - u. func (t *Tensor) Sub(u *Tensor) (*Tensor, error) { return t.binary("Sub", u, core.Sub, func(g gradSlot, _, _ *core.Array, dst *[2]gradSlot, ar *gradArena) error { sh := g.sh c0, err := copyGradSlot(ar, g, sh) if err != nil { return err } c1, err := negGradSlot(ar, g, sh) if err != nil { return err } dst[0], dst[1] = c0, c1 return nil }) } // Mul returns the element-wise product t·u. The complex adjoint // conjugates the other factor: dz = g·w̄, the Wirtinger rule for a // holomorphic product. func (t *Tensor) Mul(u *Tensor) (*Tensor, error) { return t.binary("Mul", u, core.Mul, func(g gradSlot, a, b *core.Array, dst *[2]gradSlot, ar *gradArena) error { sh := g.sh if eitherComplex(a, b) { ga, ok := mulConjGradSlot(ar, g, b, sh) if !ok { var err error if ga.arr, err = core.Mul(g.arr, conjArray(b)); err != nil { return err } ga.sh = sh } gb, ok := mulConjGradSlot(ar, g, a, sh) if !ok { var err error if gb.arr, err = core.Mul(g.arr, conjArray(a)); err != nil { return err } gb.sh = sh } dst[0], dst[1] = ga, gb return nil } ga, ok := mulGradSlot(ar, g, b, sh) if !ok { var err error if ga.arr, err = core.Mul(g.arr, b); err != nil { return err } ga.sh = sh } gb, ok := mulGradSlot(ar, g, a, sh) if !ok { var err error if gb.arr, err = core.Mul(g.arr, a); err != nil { return err } gb.sh = sh } dst[0], dst[1] = ga, gb return nil }) } // Div returns the element-wise true division t/u. The complex adjoint // is da = g/conj(b), db = −g·conj(a)/conj(b)². func (t *Tensor) Div(u *Tensor) (*Tensor, error) { return t.binary("Div", u, core.Div, func(g gradSlot, a, b *core.Array, dst *[2]gradSlot, ar *gradArena) error { sh := g.sh if eitherComplex(a, b) { cb := conjArray(b) ga, err := core.Div(g.arr, cb) if err != nil { return err } cbsq, err := core.Mul(cb, cb) if err != nil { return err } ca := conjArray(a) num, err := core.Mul(g.arr, ca) if err != nil { return err } gb, err := core.Div(core.MulI(num, -1), cbsq) if err != nil { return err } dst[0] = gradSlot{arr: ga, sh: sh} dst[1] = gradSlot{arr: gb, sh: sh} return nil } if ga, gb, ok := divGradReal(ar, g, a, b, sh); ok { dst[0], dst[1] = ga, gb return nil } ga, err := core.Div(g.arr, b) if err != nil { return err } bb, err := core.Mul(b, b) if err != nil { return err } neg := core.MulI(g.arr, -1) num, err := core.Mul(neg, a) if err != nil { return err } gb, err := core.Div(num, bb) if err != nil { return err } dst[0] = gradSlot{arr: ga, sh: sh} dst[1] = gradSlot{arr: gb, sh: sh} return nil }) } // MatMul is the differentiable matrix product (2-D×2-D, 2-D×1-D and // 1-D×2-D, mirroring MatMul2D). It rides the same plumbing as the // element-wise binary ops; only the forward and backward differ. The // complex adjoint conjugates and transposes: dA = g·Bᴴ, dB = Aᴴ·g.arr. func (t *Tensor) MatMul(u *Tensor) (*Tensor, error) { if err := t.checkDiff("MatMul"); err != nil { return nil, err } if err := u.checkDiff("MatMul"); err != nil { return nil, err } out, err := core.MatMul2D(t.data, u.data) if err != nil { return nil, err } at, au := t.data, u.data return binaryResult("MatMul", t, u, out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { var ga, gb gradSlot var err error if eitherComplex(at, au) { ga, gb, err = matmulGradComplex(ar, g, at, au) } else { ga, gb, err = matmulGrad(ar, g, at, au) } if err != nil { return err } dst[0], dst[1] = ga, gb return nil }), nil } // matmulGradComplex is the Wirtinger backward of MatMul2D: every // operand enters the adjoint conjugated, so the plain transposes of // the real backward become conjugate transposes. The transposes and // reshapes are closure-local scratch, released back to the arena once // the products are formed. func matmulGradComplex(ar *gradArena, g gradSlot, a, b *core.Array) (gradSlot, gradSlot, error) { ca, cb := conjArray(a), conjArray(b) shA, shB := a.Shape(), b.Shape() switch { case a.NDim() == 2 && b.NDim() == 2: bh, ok := transposeGradSlot(ar, gradSlot{arr: cb}, []int{shB[1], shB[0]}) if !ok { bh = gradSlot{arr: core.Transpose(cb), sh: []int{shB[1], shB[0]}} } da, err := core.MatMul2D(g.arr, bh.arr) if err != nil { return gradSlot{}, gradSlot{}, err } ah, ok := transposeGradSlot(ar, gradSlot{arr: ca}, []int{shA[1], shA[0]}) if !ok { ah = gradSlot{arr: core.Transpose(ca), sh: []int{shA[1], shA[0]}} } db, err := core.MatMul2D(ah.arr, g.arr) ar.releaseGrad(bh.arr, bh.sh) ar.releaseGrad(ah.arr, ah.sh) // conjArray answers a real operand unchanged, so the conjugate // is released only when it is a buffer this call built. if ca != a { ar.releaseGrad(ca, shA) } if cb != b { ar.releaseGrad(cb, shB) } if err != nil { return gradSlot{}, gradSlot{}, err } return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil case a.NDim() == 2 && b.NDim() == 1: // da[i, j] = g[i]·b̄[j], the outer product against the // conjugated vector. gr, err := core.Reshape(g.arr, shA[0], 1) if err != nil { return gradSlot{}, gradSlot{}, err } br, err := core.Reshape(cb, 1, b.Len()) if err != nil { return gradSlot{}, gradSlot{}, err } da, err := core.MatMul2D(gr, br) if err != nil { return gradSlot{}, gradSlot{}, err } ah, ok := transposeGradSlot(ar, gradSlot{arr: ca}, []int{shA[1], shA[0]}) if !ok { ah = gradSlot{arr: core.Transpose(ca), sh: []int{shA[1], shA[0]}} } db, err := core.MatMul2D(ah.arr, g.arr) ar.releaseGrad(ah.arr, ah.sh) if ca != a { ar.releaseGrad(ca, shA) } if cb != b { ar.releaseGrad(cb, shB) } if err != nil { return gradSlot{}, gradSlot{}, err } return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil case a.NDim() == 1 && b.NDim() == 2: bh, ok := transposeGradSlot(ar, gradSlot{arr: cb}, []int{shB[1], shB[0]}) if !ok { bh = gradSlot{arr: core.Transpose(cb), sh: []int{shB[1], shB[0]}} } da, err := core.MatMul2D(g.arr, bh.arr) if err != nil { return gradSlot{}, gradSlot{}, err } ar2, err := core.Reshape(ca, a.Len(), 1) if err != nil { return gradSlot{}, gradSlot{}, err } gr, err := core.Reshape(g.arr, 1, g.arr.Len()) if err != nil { return gradSlot{}, gradSlot{}, err } db, err := core.MatMul2D(ar2, gr) ar.releaseGrad(bh.arr, bh.sh) if ca != a { ar.releaseGrad(ca, shA) } if cb != b { ar.releaseGrad(cb, shB) } if err != nil { return gradSlot{}, gradSlot{}, err } return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil } return gradSlot{}, gradSlot{}, errf("autograd MatMul: unsupported shapes %s and %s", prettyShape(a.Shape()), prettyShape(b.Shape())) } // matmulGrad is the backward of MatMul2D for every supported shape // combination. The transposes of the operands are closure-local // scratch: a dense rank-2 operand transposes into an arena buffer, and // whichever route produced it, the buffer returns to the arena once // the two products are formed. func matmulGrad(ar *gradArena, g gradSlot, a, b *core.Array) (gradSlot, gradSlot, error) { shA, shB := a.Shape(), b.Shape() switch { case a.NDim() == 2 && b.NDim() == 2: bt, ok := transposeGradSlot(ar, gradSlot{arr: b}, []int{shB[1], shB[0]}) if !ok { bt = gradSlot{arr: core.Transpose(b), sh: []int{shB[1], shB[0]}} } da, err := core.MatMul2D(g.arr, bt.arr) if err != nil { return gradSlot{}, gradSlot{}, err } at, ok := transposeGradSlot(ar, gradSlot{arr: a}, []int{shA[1], shA[0]}) if !ok { at = gradSlot{arr: core.Transpose(a), sh: []int{shA[1], shA[0]}} } db, err := core.MatMul2D(at.arr, g.arr) ar.releaseGrad(bt.arr, bt.sh) ar.releaseGrad(at.arr, at.sh) if err != nil { return gradSlot{}, gradSlot{}, err } return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil case a.NDim() == 2 && b.NDim() == 1: // da[i, j] = g[i]·b[j], the outer product as a 1×k by n×1 // matrix product. gr, err := core.Reshape(g.arr, shA[0], 1) if err != nil { return gradSlot{}, gradSlot{}, err } br, err := core.Reshape(b, 1, b.Len()) if err != nil { return gradSlot{}, gradSlot{}, err } da, err := core.MatMul2D(gr, br) if err != nil { return gradSlot{}, gradSlot{}, err } at, ok := transposeGradSlot(ar, gradSlot{arr: a}, []int{shA[1], shA[0]}) if !ok { at = gradSlot{arr: core.Transpose(a), sh: []int{shA[1], shA[0]}} } db, err := core.MatMul2D(at.arr, g.arr) ar.releaseGrad(at.arr, at.sh) if err != nil { return gradSlot{}, gradSlot{}, err } return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil case a.NDim() == 1 && b.NDim() == 2: bt, ok := transposeGradSlot(ar, gradSlot{arr: b}, []int{shB[1], shB[0]}) if !ok { bt = gradSlot{arr: core.Transpose(b), sh: []int{shB[1], shB[0]}} } da, err := core.MatMul2D(g.arr, bt.arr) if err != nil { return gradSlot{}, gradSlot{}, err } ar_, err := core.Reshape(a, a.Len(), 1) if err != nil { return gradSlot{}, gradSlot{}, err } gr, err := core.Reshape(g.arr, 1, g.arr.Len()) if err != nil { return gradSlot{}, gradSlot{}, err } db, err := core.MatMul2D(ar_, gr) ar.releaseGrad(bt.arr, bt.sh) if err != nil { return gradSlot{}, gradSlot{}, err } return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil } return gradSlot{}, gradSlot{}, errf("autograd MatMul: unsupported shapes %s and %s", prettyShape(a.Shape()), prettyShape(b.Shape())) } // Sum reduces the tensor to a single-element tensor holding the sum of // all elements. A complex tensor sums in complex. func (t *Tensor) Sum() (*Tensor, error) { if err := t.checkDiff("Sum"); err != nil { return nil, err } a := t.data if isComplexArr(a) { out := scalarComplex(core.Sum(a).Complex()) return t.unaryResult("Sum", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := a.Shape() s := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh} fillGradSlotC(s, g.arr.ComplexAt(0)) dst[0] = s return nil }), nil } out := scalarLike(t.data, core.Sum(t.data).Float()) return t.unaryResult("Sum", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := a.Shape() s := gradSlot{arr: ar.borrowGrad(a.Dtype(), sh), sh: sh} fillGradSlot(s, g.arr.FloatAt(0)) dst[0] = s return nil }), nil } // Mean reduces the tensor to a single-element tensor holding the mean // of all elements. A complex tensor means in complex (the core Mean // is real-only, so the complex path divides the sum directly). func (t *Tensor) Mean() (*Tensor, error) { if err := t.checkDiff("Mean"); err != nil { return nil, err } a := t.data if isComplexArr(a) { if a.Len() == 0 { // The real branch's core.Mean refuses an empty reduction; // this branch divides by the element count, so an empty // tensor would answer 0/0 = NaN instead. return nil, errf("Mean: an empty array has no mean") } s := core.Sum(a).Complex() n := complex(float64(a.Len()), 0) out := scalarComplex(s / n) return t.unaryResult("Mean", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := a.Shape() da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh} fillGradSlotC(da, g.arr.ComplexAt(0)/n) dst[0] = da return nil }), nil } m, err := core.Mean(t.data) if err != nil { return nil, err } out := scalarLike(t.data, m) n := float64(a.Len()) return t.unaryResult("Mean", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := a.Shape() da := gradSlot{arr: ar.borrowGrad(a.Dtype(), sh), sh: sh} fillGradSlot(da, g.arr.FloatAt(0)/n) dst[0] = da return nil }), nil } // Exp returns e raised to each element. Complex input differentiates // too: exp is holomorphic, so the Wirtinger adjoint multiplies by the // conjugated output. func (t *Tensor) Exp() (*Tensor, error) { if err := t.checkDiff("Exp"); err != nil { return nil, err } if isComplexArr(t.data) { a := t.data out := zeros(core.Complex, a.Shape()) cs := out.RawComplexes() if a.Strided() { for i := range cs { z := a.ComplexAt(i) cs[i] = complex(math.Exp(real(z))*math.Cos(imag(z)), math.Exp(real(z))*math.Sin(imag(z))) } } else { as := a.RawComplexes() for i := range cs { z := as[i] cs[i] = complex(math.Exp(real(z))*math.Cos(imag(z)), math.Exp(real(z))*math.Sin(imag(z))) } } return t.unaryResult("Exp", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh} ds := da.arr.RawComplexes()[:da.arr.Len()] if g.arr.Strided() { for i := range ds { ds[i] = g.arr.ComplexAt(i) * conj(cs[i]) } dst[0] = da return nil } gs := g.arr.RawComplexes() for i := range ds { ds[i] = gs[i] * conj(cs[i]) } dst[0] = da return nil }), nil } out, err := core.Exp(t.data) if err != nil { return nil, err } a := t.data return t.unaryResult("Exp", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) da, ok := mulGradSlot(ar, g, out, sh) if !ok { arr, nerr := core.Mul(g.arr, out) if nerr != nil { return nerr } da = gradSlot{arr: arr, sh: sh} } dst[0] = da return nil }), nil } // Log returns the natural logarithm of each element. The domain is the // caller's: a non-positive input is not refused, and the NaN or ±Inf // the forward produces flows into the gradient without a report, so // shift the input inside its domain before it reaches the graph. func (t *Tensor) Log() (*Tensor, error) { if err := t.checkFloat("Log"); err != nil { return nil, err } out, err := core.Log(t.data) if err != nil { return nil, err } a := t.data return t.unaryResult("Log", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) da, ok := divGradSlot(ar, g, a, sh) if !ok { arr, nerr := core.Div(g.arr, a) if nerr != nil { return nerr } da = gradSlot{arr: arr, sh: sh} } dst[0] = da return nil }), nil } // Sigmoid returns the logistic sigmoid of each element. func (t *Tensor) Sigmoid() (*Tensor, error) { if err := t.checkFloat("Sigmoid"); err != nil { return nil, err } out, err := core.Sigmoid(t.data) if err != nil { return nil, err } return t.unaryResult("Sigmoid", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { // dz = g·σ·(1−σ) in one pass; the fallback keeps the staged // core chain for anything the raw paths cannot see through. sh := gradShape(g, out) if g.arr.Dtype() != out.Dtype() || g.arr.Strided() || out.Strided() { comp, err := core.Sub(fillLike(out, 1), out) if err != nil { return err } so, err := core.Mul(out, comp) if err != nil { return err } da, err := core.Mul(g.arr, so) if err != nil { return err } dst[0] = gradSlot{arr: da, sh: sh} return nil } da := gradSlot{arr: ar.borrowGrad(out.Dtype(), sh), sh: sh} if out.Dtype() == core.Float32 { gs, os := g.arr.RawFloat32s(), out.RawFloat32s() ds := da.arr.RawFloat32s()[:da.arr.Len()] elemPass(len(ds), func(s, e int) { for i := s; i < e; i++ { cp := float32(1 - float64(os[i])) so := float32(float64(os[i]) * float64(cp)) ds[i] = float32(float64(gs[i]) * float64(so)) } }) dst[0] = da return nil } gs, os := g.arr.RawFloats(), out.RawFloats() ds := da.arr.RawFloats()[:da.arr.Len()] if len(ds) < elemSweepMin { for i := range ds { ds[i] = gs[i] * (os[i] * (1 - os[i])) } } else { elemPass(len(ds), func(s, e int) { for i := s; i < e; i++ { ds[i] = gs[i] * (os[i] * (1 - os[i])) } }) } dst[0] = da return nil }), nil } // Tanh returns the hyperbolic tangent of each element. func (t *Tensor) Tanh() (*Tensor, error) { if err := t.checkFloat("Tanh"); err != nil { return nil, err } out, err := core.Tanh(t.data) if err != nil { return nil, err } return t.unaryResult("Tanh", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { // dz = g·(1−tanh²) in one pass; the float32 branch rounds at the // same three points the Mul/Sub chain it replaces did, so the // payload is unchanged element for element. sh := gradShape(g, out) if g.arr.Dtype() != out.Dtype() || g.arr.Strided() || out.Strided() { sq, err := core.Mul(out, out) if err != nil { return err } comp, err := core.Sub(fillLike(out, 1), sq) if err != nil { return err } da, err := core.Mul(g.arr, comp) if err != nil { return err } dst[0] = gradSlot{arr: da, sh: sh} return nil } da := gradSlot{arr: ar.borrowGrad(out.Dtype(), sh), sh: sh} if out.Dtype() == core.Float32 { gs, os := g.arr.RawFloat32s(), out.RawFloat32s() ds := da.arr.RawFloat32s()[:da.arr.Len()] elemPass(len(ds), func(s, e int) { for i := s; i < e; i++ { sq := float32(float64(os[i]) * float64(os[i])) cp := float32(1 - float64(sq)) ds[i] = float32(float64(gs[i]) * float64(cp)) } }) dst[0] = da return nil } gs, os := g.arr.RawFloats(), out.RawFloats() ds := da.arr.RawFloats()[:da.arr.Len()] if len(ds) < elemSweepMin { for i := range ds { ds[i] = gs[i] * (1 - os[i]*os[i]) } } else { elemPass(len(ds), func(s, e int) { for i := s; i < e; i++ { ds[i] = gs[i] * (1 - os[i]*os[i]) } }) } dst[0] = da return nil }), nil } // Neg returns -t. func (t *Tensor) Neg() (*Tensor, error) { if err := t.checkDiff("Neg"); err != nil { return nil, err } a := t.data return t.unaryResult("Neg", core.MulI(t.data, -1), func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) da, err := negGradSlot(ar, g, sh) if err != nil { return err } dst[0] = da return nil }), nil } // Transpose reverses the dimensions; the gradient transposes back. func (t *Tensor) Transpose() (*Tensor, error) { if err := t.checkDiff("Transpose"); err != nil { return nil, err } a := t.data return t.unaryResult("Transpose", core.Transpose(t.data), func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := a.Shape() da, ok := transposeGradSlot(ar, g, sh) if !ok { da = gradSlot{arr: core.Transpose(g.arr), sh: sh} } dst[0] = da return nil }), nil } // Squeeze removes size-1 dimensions; the gradient unsqueezes back. func (t *Tensor) Squeeze(dim int) (*Tensor, error) { if err := t.checkDiff("Squeeze"); err != nil { return nil, err } out, err := core.Squeeze(t.data, dim) if err != nil { return nil, err } orig := t.data.Shape() return t.unaryResult("Squeeze", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { gr, err := core.Reshape(g.arr, orig...) if err != nil { return err } dst[0] = gradSlot{arr: gr, sh: orig} return nil }), nil } // Unsqueeze inserts a size-1 dimension; the gradient squeezes back. func (t *Tensor) Unsqueeze(dim int) (*Tensor, error) { if err := t.checkDiff("Unsqueeze"); err != nil { return nil, err } out, err := core.Unsqueeze(t.data, dim) if err != nil { return nil, err } orig := t.data.Shape() return t.unaryResult("Unsqueeze", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { gr, err := core.Reshape(g.arr, orig...) if err != nil { return err } dst[0] = gradSlot{arr: gr, sh: orig} return nil }), nil } // Clip clamps each element into [lo, hi]; the gradient passes through // only where the value was inside the range. func (t *Tensor) Clip(lo, hi float64) (*Tensor, error) { if err := t.checkFloat("Clip"); err != nil { return nil, err } out, err := core.ClipF(t.data, lo, hi) if err != nil { return nil, err } a := t.data return t.unaryResult("Clip", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { // d/dx clip(x) = 1 where lo ≤ x ≤ hi, 0 elsewhere. The // comparisons answer bool masks, so their composition is the // logic conjunction; And keeps the mask a bool payload. sh := gradShape(g, a) ge, err := core.GeF(a, lo) if err != nil { return err } le, err := core.LeF(a, hi) if err != nil { return err } mask, err := core.And(ge, le) if err != nil { return err } // Mul reads the bool mask through its exact 0/1 widening: g·1 // is g bit for bit and g·0 is zero, the values the int mask's // product used to write. da, err := core.Mul(g.arr, mask) if err != nil { return err } dst[0] = gradSlot{arr: da, sh: sh} return nil }), nil } // scalarLike builds a 1-element array of a's dtype holding v. func scalarLike(a *core.Array, v float64) *core.Array { out := zeros(a.Dtype(), []int{1}) out.SetFloatAt(0, v) return out } // errf formats an error message with the package prefix. func errf(format string, args ...any) error { return fmt.Errorf("tensor: "+format, args...) } // fillLike returns an array shaped like a with every element set to v. // One pass over a fresh payload, unlike the OnesLike-then-MulF walk it // replaces: ones·v multiplies to exactly v per element (1·v is the // IEEE identity), so the bits are unchanged. func fillLike(a *core.Array, v float64) *core.Array { return fillConst(a.Dtype(), a.Shape(), v) } // fillConst returns a fresh array of dt and shape with every element // set to v, written straight into the raw payload. An int dtype // promotes to float, mirroring the scalar-multiply semantics the // accessor path had. The sweep stays on the calling goroutine: it is a // plain store stream over one buffer, so a split buys no bandwidth and // measures slower than the serial store. func fillConst(dt core.Dtype, shape []int, v float64) *core.Array { if dt == core.Int { dt = core.Float } out := zeros(dt, shape) switch dt { case core.Float: fs := out.RawFloats() for i := range fs { fs[i] = v } case core.Float32: fs := out.RawFloat32s() fv := float32(v) for i := range fs { fs[i] = fv } default: cs := out.RawComplexes() z := complex(v, 0) for i := range cs { cs[i] = z } } return out } // elemPass runs fn over n elements under the elementwise spawn policy: // ranges too small to amortise a worker stay whole on the calling // goroutine, large ones split into disjoint chunks. The split cannot // move a bit: every element is computed from its own operands alone. func elemPass(n int, fn func(s, e int)) { engine.ParallelMin(n, 1024, fn) } // mathElemMinPerWorker is the per-worker chunk floor for the sweeps // whose per-element cost is a math.Pow call: tens of cycles each, so a // spawned worker amortises its own start-up at a far smaller chunk // than the arithmetic maps' floor, which holds every sweep below // roughly 32k elements on the calling goroutine. A divide is cheap // enough per element that the arithmetic floor stays the right one for // it: measured at 20k elements, moving the square-root backward to // this floor turned a 63µs serial sweep into a 94µs parallel one, // while the two power sweeps halved. const mathElemMinPerWorker = 256 // mathElemPass is elemPass for the costly per-element sweeps: same // split, same arithmetic, only the spawn floor differs. func mathElemPass(n int, fn func(s, e int)) { engine.ParallelMin(n, mathElemMinPerWorker, fn) } // l2LinesPerWorker is the per-worker line floor for the L2 norm // backward: one line costs two passes over the reduced axis, so the // element-wise floor of about a thousand elements per worker is that // many line slots, and never below one line. func l2LinesPerWorker(size int) int { if size < 1 { return 1 } return max(1, 1024/(2*size)) } // flatFloats copies a's elements into a fresh float64 slice, reading // the contiguous payload directly when possible and falling back to // the widening accessor for views and other dtypes. func flatFloats(a *core.Array) []float64 { out := make([]float64, a.Len()) if !a.Strided() && a.Dtype() == core.Float { copy(out, a.RawFloats()) return out } for i := range out { out[i] = a.FloatAt(i) } return out } // prettyShape renders a shape for diagnostics as "(2, 3)". func prettyShape(shape []int) string { parts := make([]string, len(shape)) for i, d := range shape { parts[i] = strconv.Itoa(d) } return "(" + strings.Join(parts, ", ") + ")" } // ReplaceWith swaps the underlying data array, the one mutable point // the optimisers rely on for parameter updates. func (t *Tensor) ReplaceWith(a *core.Array) { t.data = a } // SumAxis sums the elements along dim, dropping it; the backward // broadcasts the incoming gradient back over the reduced dimension. func (t *Tensor) SumAxis(dim int) (*Tensor, error) { if err := t.checkDiff("SumAxis"); err != nil { return nil, err } out, err := core.SumAxis(t.data, dim) if err != nil { return nil, err } shape := append([]int{}, t.data.Shape()...) return t.unaryResult("SumAxis", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { expanded, err := expandToShape(g.arr, dim, shape) if err != nil { return err } dst[0] = gradSlot{arr: expanded, sh: shape} return nil }), nil } // MeanAxis averages along dim, dropping it; the backward is the sum // backward scaled by 1/size(dim). func (t *Tensor) MeanAxis(dim int) (*Tensor, error) { if err := t.checkFloat("MeanAxis"); err != nil { return nil, err } out, err := core.MeanAxis(t.data, dim) if err != nil { return nil, err } shape := append([]int{}, t.data.Shape()...) n := float64(shape[dim]) return t.unaryResult("MeanAxis", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { // Scale the incoming gradient before the broadcast: the division // hits every source element once either way, and the broadcast // copies the quotient exactly. The elements are independent, so // the split cannot move a bit, and a divide is cheap enough for // the arithmetic maps' spawn floor. gsh := g.sh scaled := gradSlot{arr: ar.borrowGrad(g.arr.Dtype(), gsh), sh: gsh} if !g.arr.Strided() && g.arr.Dtype() == core.Float32 { gs := g.arr.RawFloat32s() ss := scaled.arr.RawFloat32s()[:scaled.arr.Len()] elemPass(len(ss), func(s, e int) { for i := s; i < e; i++ { ss[i] = float32(float64(gs[i]) / n) } }) } else if !g.arr.Strided() && g.arr.Dtype() == core.Float { gs := g.arr.RawFloats() ss := scaled.arr.RawFloats()[:scaled.arr.Len()] elemPass(len(ss), func(s, e int) { for i := s; i < e; i++ { ss[i] = gs[i] / n } }) } else { for i := range g.arr.Len() { scaled.arr.SetFloatAt(i, g.arr.FloatAt(i)/n) } } expanded, err := expandToShape(scaled.arr, dim, shape) ar.releaseGrad(scaled.arr, gsh) if err != nil { return err } dst[0] = gradSlot{arr: expanded, sh: shape} return nil }), nil } // L2NormAxis computes the L2 norm along dim, dropping it. The backward // divides each input element by the line norm: dx = g·x/‖x‖. func (t *Tensor) L2NormAxis(dim int) (*Tensor, error) { if err := t.checkFloat("L2NormAxis"); err != nil { return nil, err } out, err := core.Norm(t.data, 2, dim, false) if err != nil { return nil, err } x := t.data size := x.Shape()[dim] stride := 1 for k := dim + 1; k < x.NDim(); k++ { stride *= x.Shape()[k] } // An empty dimension (or an empty trailing axis) empties the whole // tensor; the gradient of an empty tensor is empty, and the line // count below must not divide by a zero stride. perLine := size * stride lines := 0 if perLine > 0 { lines = x.Len() / perLine } const tiny = 1e-12 xShape := x.Shape() return t.unaryResult("L2NormAxis", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { dx := gradSlot{arr: ar.borrowGrad(core.Float, xShape), sh: xShape} // The gradient arrives as float64 (the norm's dtype). The raw // sweeps below widen a float32 input exactly where FloatAt did, // and the fallback widens a strided operand once into a linear // buffer instead of calling the accessor per element. fast64 := !x.Strided() && !g.arr.Strided() && x.Dtype() == core.Float && g.arr.Dtype() == core.Float fast32 := !x.Strided() && !g.arr.Strided() && x.Dtype() == core.Float32 && g.arr.Dtype() == core.Float linesPerWorker := l2LinesPerWorker(size) switch { case fast64: gs, xs, ds := g.arr.RawFloats(), x.RawFloats(), dx.arr.RawFloats()[:dx.arr.Len()] engine.ParallelMin(lines*stride, linesPerWorker, func(s, e int) { for key := s; key < e; key++ { line := key / stride off := key % stride var sq float64 for k := range size { v := xs[line*size*stride+k*stride+off] sq += v * v } denom := math.Sqrt(sq) if denom < tiny { denom = tiny } gv := gs[key] for k := range size { pos := line*size*stride + k*stride + off ds[pos] = gv * xs[pos] / denom } } }) case fast32: gs, xs32, ds := g.arr.RawFloats(), x.RawFloat32s(), dx.arr.RawFloats()[:dx.arr.Len()] engine.ParallelMin(lines*stride, linesPerWorker, func(s, e int) { for key := s; key < e; key++ { line := key / stride off := key % stride var sq float64 for k := range size { v := float64(xs32[line*size*stride+k*stride+off]) sq += v * v } denom := math.Sqrt(sq) if denom < tiny { denom = tiny } gv := gs[key] for k := range size { pos := line*size*stride + k*stride + off ds[pos] = gv * float64(xs32[pos]) / denom } } }) default: xs, gs, ds := flatFloats(x), flatFloats(g.arr), dx.arr.RawFloats()[:dx.arr.Len()] engine.ParallelMin(lines*stride, linesPerWorker, func(s, e int) { for key := s; key < e; key++ { line := key / stride off := key % stride var sq float64 for k := range size { v := xs[line*size*stride+k*stride+off] sq += v * v } denom := math.Sqrt(sq) if denom < tiny { denom = tiny } gv := gs[key] for k := range size { pos := line*size*stride + k*stride + off ds[pos] = gv * xs[pos] / denom } } }) } dst[0] = dx return nil }), nil } // expandToShape reinserts a dropped axis at dim (size 1) and broadcasts // g to shape, the inverse of an axis reduction. The reshape and the // broadcast validate their shapes, so a mismatch travels back to the // backward pass as an error instead of a nil array. func expandToShape(g *core.Array, dim int, shape []int) (*core.Array, error) { gShape := g.Shape() newShape := make([]int, 0, len(shape)) newShape = append(newShape, gShape[:dim]...) newShape = append(newShape, 1) newShape = append(newShape, gShape[dim:]...) r, err := core.Reshape(g, newShape...) if err != nil { return nil, err } return core.BroadcastTo(r, shape...) } // BroadcastTo expands the tensor to shape under broadcasting rules // (size-1 dimensions replicate). The backward sums the incoming // gradient over every replicated dimension, collapsing it back to the // original shape. Only float, float32 and complex tensors carry a // graph, so an int tensor is refused; the shape is validated first, so // an impossible target keeps reporting the shape mismatch. func (t *Tensor) BroadcastTo(shape ...int) (*Tensor, error) { out, err := core.BroadcastTo(t.data, shape...) if err != nil { return nil, err } if err := t.checkDiff("BroadcastTo"); err != nil { return nil, err } orig := t.data.Shape() return t.unaryResult("BroadcastTo", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { red, err := sumToShape(g.arr, orig) if err != nil { return err } dst[0] = gradSlot{arr: red, sh: orig} return nil }), nil } // sumToShape reduces g back to target by summing across every expanded // dimension: leading extra axes are collapsed one by one, then any dim // whose size grew (target == 1) is summed at its own position. A rank-1 // target of size 1 has no summable axis left, so the whole gradient // collapses to a single sum. func sumToShape(g *core.Array, target []int) (*core.Array, error) { for g.NDim() > len(target) { g2, err := core.SumAxis(g, 0) if err != nil { return nil, err } g = g2 } for i := range len(target) { if g.Shape()[i] == target[i] { continue } if g.NDim() == 1 { // The only dimension left: under the broadcasting rules a // mismatch here means target[0] == 1, so the full sum is // the correct reduction. if isComplexArr(g) { return scalarComplex(core.Sum(g).Complex()), nil } return scalarLike(g, core.Sum(g).Float()), nil } g2, err := core.SumAxis(g, i) if err != nil { return nil, err } ns := make([]int, 0, g.NDim()) ns = append(ns, g.Shape()[:i]...) ns = append(ns, 1) ns = append(ns, g.Shape()[i+1:]...) if g, err = core.Reshape(g2, ns...); err != nil { return nil, err } } return g, nil } // Pow raises each element to an integer exponent n ≥ 0 on the graph; // backward multiplies by n·x^(n−1). The exponent-0 backward is the // zero gradient everywhere, including at x = 0, where n·x^(n−1) would // evaluate 0·∞ and come out NaN. A complex tensor conjugates the // derivative, per the Wirtinger convention. func (t *Tensor) Pow(n int64) (*Tensor, error) { if n < 0 { return nil, errf("Pow: negative exponent has no general real gradient") } if err := t.checkDiff("Pow"); err != nil { return nil, err } if isComplexArr(t.data) { a := t.data out := powComplexForward(a, n) return t.unaryResult("Pow", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) dst[0] = gradSlot{arr: powComplexGrad(ar, g, a, n, sh), sh: sh} return nil }), nil } out := zeros(t.Data().Dtype(), t.Data().Shape()) a := t.Data() exp := float64(n) switch { case a.Strided(): // The destination is fresh and dense, so its payload takes the // power directly while the strided operand keeps the accessor // read that rebases the index. if out.Dtype() == core.Float32 { os := out.RawFloat32s() engine.Parallel(len(os), func(s, e int) { for i := s; i < e; i++ { os[i] = float32(math.Pow(a.FloatAt(i), exp)) } }) } else { os := out.RawFloats() engine.Parallel(len(os), func(s, e int) { for i := s; i < e; i++ { os[i] = math.Pow(a.FloatAt(i), exp) } }) } case a.Dtype() == core.Float32: as, os := a.RawFloat32s(), out.RawFloat32s() mathElemPass(len(os), func(s, e int) { for i := s; i < e; i++ { os[i] = float32(math.Pow(float64(as[i]), exp)) } }) default: as, os := a.RawFloats(), out.RawFloats() mathElemPass(len(os), func(s, e int) { for i := s; i < e; i++ { os[i] = math.Pow(as[i], exp) } }) } return t.unaryResult("Pow", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) dx := gradSlot{arr: ar.borrowGrad(core.Float, sh), sh: sh} if exp == 0 { // d/dx x⁰ = 0, including at x = 0. dst[0] = dx return nil } // The backward keeps the gradient sweep on raw payloads; dx // stays float64 whatever the operand width, as the accessor // path's zeros(core.Float) always did. if !g.arr.Strided() && g.arr.Dtype() == core.Float32 && !a.Strided() && a.Dtype() == core.Float32 { gs, as, ds := g.arr.RawFloat32s(), a.RawFloat32s(), dx.arr.RawFloats()[:dx.arr.Len()] mathElemPass(len(ds), func(s, e int) { for i := s; i < e; i++ { ds[i] = float64(gs[i]) * exp * math.Pow(float64(as[i]), exp-1) } }) dst[0] = dx return nil } if !g.arr.Strided() && g.arr.Dtype() == core.Float && !a.Strided() && a.Dtype() == core.Float { gs, as, ds := g.arr.RawFloats(), a.RawFloats(), dx.arr.RawFloats()[:dx.arr.Len()] mathElemPass(len(ds), func(s, e int) { for i := s; i < e; i++ { ds[i] = gs[i] * exp * math.Pow(as[i], exp-1) } }) dst[0] = dx return nil } // Mixed widths or a strided operand: the accessors convert and // rebase, the destination is dx's own dense float64 payload. ds := dx.arr.RawFloats()[:dx.arr.Len()] engine.Parallel(a.Len(), func(s, e int) { for i := s; i < e; i++ { x := a.FloatAt(i) ds[i] = g.arr.FloatAt(i) * exp * math.Pow(x, exp-1) } }) dst[0] = dx return nil }), nil } // Abs returns |x|; the real subgradient is sign(x) (0 at zero), the // complex one z/(2|z|). func (t *Tensor) Abs() (*Tensor, error) { if err := t.checkDiff("Abs"); err != nil { return nil, err } if isComplexArr(t.data) { return t.absComplex() } a := t.data return t.unaryResult("Abs", core.Abs(a), func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) dx := gradSlot{arr: ar.borrowGrad(core.Float, sh), sh: sh} // dx = g·sign(x) on raw payloads; the gradient is float64 here // whatever the operand width, so the dtype branch sits outside // the loop. if !g.arr.Strided() && !a.Strided() && g.arr.Dtype() == core.Float && a.Dtype() == core.Float { gs, as, ds := g.arr.RawFloats(), a.RawFloats(), dx.arr.RawFloats()[:dx.arr.Len()] elemPass(len(ds), func(s, e int) { for i := s; i < e; i++ { dv := 0.0 if v := as[i]; v > 0 { dv = 1 } else if v < 0 { dv = -1 } ds[i] = gs[i] * dv } }) dst[0] = dx return nil } if !g.arr.Strided() && !a.Strided() && g.arr.Dtype() == core.Float32 && a.Dtype() == core.Float32 { gs, as, ds := g.arr.RawFloat32s(), a.RawFloat32s(), dx.arr.RawFloats()[:dx.arr.Len()] elemPass(len(ds), func(s, e int) { for i := s; i < e; i++ { dv := 0.0 if v := as[i]; v > 0 { dv = 1 } else if v < 0 { dv = -1 } ds[i] = float64(gs[i]) * dv } }) dst[0] = dx return nil } ds := dx.arr.RawFloats()[:dx.arr.Len()] engine.Parallel(a.Len(), func(s, e int) { for i := s; i < e; i++ { v := a.FloatAt(i) dv := 0.0 if v > 0 { dv = 1 } else if v < 0 { dv = -1 } ds[i] = g.arr.FloatAt(i) * dv } }) dst[0] = dx return nil }), nil } // Sqrt returns √x; backward is 1/(2√x). The domain is the caller's: a // negative input is not refused, and the NaN the forward produces // flows into the gradient without a report, so clip or shift the // input before it reaches the graph. func (t *Tensor) Sqrt() (*Tensor, error) { if err := t.checkFloat("Sqrt"); err != nil { return nil, err } out, err := core.Sqrt(t.data) if err != nil { return nil, err } const eps = 1e-12 return t.unaryResult("Sqrt", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, out) dx := gradSlot{arr: ar.borrowGrad(core.Float, sh), sh: sh} // dx = g/(2√x) with the same epsilon floor, swept from the raw // payloads; dx stays float64 whatever the operand width. if !g.arr.Strided() && !out.Strided() && g.arr.Dtype() == core.Float && out.Dtype() == core.Float { gs, os, ds := g.arr.RawFloats(), out.RawFloats(), dx.arr.RawFloats()[:dx.arr.Len()] elemPass(len(ds), func(s, e int) { for i := s; i < e; i++ { sv := os[i] if sv < eps { sv = eps } ds[i] = gs[i] / (2 * sv) } }) dst[0] = dx return nil } if !g.arr.Strided() && !out.Strided() && g.arr.Dtype() == core.Float32 && out.Dtype() == core.Float32 { gs, os, ds := g.arr.RawFloat32s(), out.RawFloat32s(), dx.arr.RawFloats()[:dx.arr.Len()] elemPass(len(ds), func(s, e int) { for i := s; i < e; i++ { sv := float64(os[i]) if sv < eps { sv = eps } ds[i] = float64(gs[i]) / (2 * sv) } }) dst[0] = dx return nil } ds := dx.arr.RawFloats()[:dx.arr.Len()] engine.Parallel(out.Len(), func(s, e int) { for i := s; i < e; i++ { sv := out.FloatAt(i) if sv < eps { sv = eps } ds[i] = g.arr.FloatAt(i) / (2 * sv) } }) dst[0] = dx return nil }), nil } // Floor rounds down; its derivative is zero almost everywhere, so // Backward on its output contributes no gradient. func (t *Tensor) Floor() (*Tensor, error) { if err := t.checkFloat("Floor"); err != nil { return nil, err } out, err := core.Floor(t.data) if err != nil { return nil, err } return t.unaryResult("Floor", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { // All-zero gradient; the borrowed buffer arrives zeroed. sh := g.sh dst[0] = gradSlot{arr: ar.borrowGrad(g.arr.Dtype(), sh), sh: sh} return nil }), nil } // SetGrad replaces the accumulated gradient array: a low-level hook // for gradient-manipulation utilities (clipping, accumulation resets, // test harnesses) to drive through directly. func (t *Tensor) SetGrad(g *core.Array) { t.grad = g } // scaleGradSlot writes g·f into a borrowed buffer shaped sh with the // expressions mapReal writes per dtype, so the bits are MulF's own. // The second return is false outside the dense float, float32 and // complex paths, and the caller keeps core.MulF. func scaleGradSlot(ar *gradArena, g gradSlot, sh []int, f float64) (gradSlot, bool) { if g.arr == nil || g.arr.Strided() || len(sh) == 0 || len(sh) > 6 { return gradSlot{}, false } n := g.arr.Len() switch g.arr.Dtype() { case core.Float: out := ar.borrowGrad(core.Float, sh) gs, os := g.arr.RawFloats()[:n], out.RawFloats()[:n] if n < elemSweepMin { for i := range os { os[i] = gs[i] * f } } else { elemPass(n, func(s, e int) { for i := s; i < e; i++ { os[i] = gs[i] * f } }) } return gradSlot{arr: out, sh: sh}, true case core.Float32: out := ar.borrowGrad(core.Float32, sh) gs, os := g.arr.RawFloat32s()[:n], out.RawFloat32s()[:n] if n < elemSweepMin { for i := range os { os[i] = float32(float64(gs[i]) * f) } } else { elemPass(n, func(s, e int) { for i := s; i < e; i++ { os[i] = float32(float64(gs[i]) * f) } }) } return gradSlot{arr: out, sh: sh}, true case core.Complex: out := ar.borrowGrad(core.Complex, sh) gs, os := g.arr.RawComplexes()[:n], out.RawComplexes()[:n] for i := range os { os[i] = gs[i] * complex(f, 0) } return gradSlot{arr: out, sh: sh}, true } return gradSlot{}, false } // Scale multiplies every element by a constant factor; backward // multiplies the incoming gradient by the same factor. func (t *Tensor) Scale(f float64) (*Tensor, error) { if err := t.checkDiff("Scale"); err != nil { return nil, err } a := t.data return t.unaryResult("Scale", core.MulF(t.data, f), func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) da, ok := scaleGradSlot(ar, g, sh, f) if !ok { da = gradSlot{arr: core.MulF(g.arr, f), sh: sh} } dst[0] = da return nil }), nil }