Files
tensor/grad/tensor.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

2304 lines
69 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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
}