Files

2304 lines
69 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
}