2304 lines
69 KiB
Go
2304 lines
69 KiB
Go
// 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
|
||
}
|