166 lines
5.1 KiB
Go
166 lines
5.1 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package grad
|
||
|
|
|
||
|
|
import (
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
|
|
)
|
||
|
|
|
||
|
|
// MatMulBatched multiplies stacked 3-D matrices batch-wise:
|
||
|
|
// (N, M, K) · (N, K, P) gives (N, M, P). The backward runs the classic
|
||
|
|
// product rule inside every batch slot, dA·Bᵀ and Aᵀ·dB, so batched
|
||
|
|
// sequence models can one day drop their per-slice fan-out without
|
||
|
|
// leaving the graph.
|
||
|
|
func (t *Tensor) MatMulBatched(u *Tensor) (*Tensor, error) {
|
||
|
|
if err := t.checkFloat("MatMulBatched"); err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
if err := u.checkFloat("MatMulBatched"); err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
ta, tu := t.data.Shape(), u.data.Shape()
|
||
|
|
if len(ta) != 3 || len(tu) != 3 {
|
||
|
|
return nil, errf("MatMulBatched: needs rank-3 operands, got %s and %s",
|
||
|
|
prettyShape(ta), prettyShape(tu))
|
||
|
|
}
|
||
|
|
if ta[0] != tu[0] || ta[2] != tu[1] {
|
||
|
|
return nil, errf("MatMulBatched: batch or inner dimension mismatch for %s · %s",
|
||
|
|
prettyShape(ta), prettyShape(tu))
|
||
|
|
}
|
||
|
|
n := ta[0]
|
||
|
|
|
||
|
|
slices := make([]*core.Array, n)
|
||
|
|
// The dtype follows the promotion ladder even for an empty batch,
|
||
|
|
// where no product runs to derive it: an Int output for float
|
||
|
|
// inputs would leak the zero value.
|
||
|
|
dt := t.data.Dtype()
|
||
|
|
if u.data.Dtype() == core.Complex || dt == core.Complex {
|
||
|
|
dt = core.Complex
|
||
|
|
} else if u.data.Dtype() == core.Float && dt != core.Float {
|
||
|
|
dt = core.Float
|
||
|
|
}
|
||
|
|
for i := range n {
|
||
|
|
aSlot, _ := core.Slice(t.data, 0, i, i+1)
|
||
|
|
bSlot, _ := core.Slice(u.data, 0, i, i+1)
|
||
|
|
aMat, _ := core.Reshape(aSlot, ta[1], ta[2])
|
||
|
|
bMat, _ := core.Reshape(bSlot, tu[1], tu[2])
|
||
|
|
prod, err := core.MatMul2D(aMat, bMat)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
slices[i] = prod
|
||
|
|
dt = prod.Dtype()
|
||
|
|
}
|
||
|
|
out := zeros(dt, []int{n, ta[1], tu[2]})
|
||
|
|
slot := ta[1] * tu[2]
|
||
|
|
// Each slot lands in a fresh contiguous product, so the scatter is
|
||
|
|
// a raw slice move per batch row, widened nowhere: out carries the
|
||
|
|
// products' own dtype.
|
||
|
|
for i := range n {
|
||
|
|
switch dt {
|
||
|
|
case core.Float32:
|
||
|
|
copy(out.RawFloat32s()[i*slot:(i+1)*slot], slices[i].RawFloat32s())
|
||
|
|
case core.Float:
|
||
|
|
copy(out.RawFloats()[i*slot:(i+1)*slot], slices[i].RawFloats())
|
||
|
|
default:
|
||
|
|
for j := range slot {
|
||
|
|
out.SetFloatAt(i*slot+j, slices[i].FloatAt(j))
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
at, au := t.data, u.data
|
||
|
|
return binaryResult("MatMulBatched", t, u, out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||
|
|
da := gradSlot{arr: ar.borrowGrad(core.Float, ta), sh: ta}
|
||
|
|
db := gradSlot{arr: ar.borrowGrad(core.Float, tu), sh: tu}
|
||
|
|
m, k, p := ta[1], ta[2], tu[2]
|
||
|
|
slotLen := m * p
|
||
|
|
for i := range n {
|
||
|
|
gMat, err := core.Reshape(
|
||
|
|
mustWindow(g.arr, i*slotLen, slotLen), m, p)
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
bMat, err := windowMatrix(au, i, k, p)
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
bT := core.Transpose(bMat)
|
||
|
|
daPart, err := core.MatMul2D(gMat, bT)
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
copyInto(da.arr, daPart, i*m*k)
|
||
|
|
aMat, err := windowMatrix(at, i, m, k)
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
aT := core.Transpose(aMat)
|
||
|
|
dbPart, err := core.MatMul2D(aT, gMat)
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
copyInto(db.arr, dbPart, i*k*p)
|
||
|
|
}
|
||
|
|
dst[0], dst[1] = da, db
|
||
|
|
return nil
|
||
|
|
}), nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// mustWindow flattens row-major window [from, from+len) into an
|
||
|
|
// (rows, cols) matrix view materialisation. The window hands its slice
|
||
|
|
// to FloatsFromArray, which takes ownership, so contiguous gradients
|
||
|
|
// copy nothing at all.
|
||
|
|
func mustWindow(g *core.Array, from, length int) *core.Array {
|
||
|
|
if !g.Strided() && g.Dtype() == core.Float {
|
||
|
|
arr, _ := core.FloatsFromArray(g.RawFloats()[from:from+length], length)
|
||
|
|
return arr
|
||
|
|
}
|
||
|
|
vals := make([]float64, length)
|
||
|
|
for i := range vals {
|
||
|
|
vals[i] = g.FloatAt(from + i)
|
||
|
|
}
|
||
|
|
arr, _ := core.FromFloats(vals, length)
|
||
|
|
return arr
|
||
|
|
}
|
||
|
|
|
||
|
|
// windowMatrix reads batch slot n as an (r, c) float64 matrix, the
|
||
|
|
// backward's native arithmetic domain. A contiguous float64 operand is
|
||
|
|
// aliased rather than copied; anything else is read through the
|
||
|
|
// widening accessor.
|
||
|
|
func windowMatrix(a *core.Array, n, r, c int) (*core.Array, error) {
|
||
|
|
base := n * r * c
|
||
|
|
if !a.Strided() && a.Dtype() == core.Float {
|
||
|
|
arr, err := core.FloatsFromArray(a.RawFloats()[base:base+r*c], r, c)
|
||
|
|
return arr, err
|
||
|
|
}
|
||
|
|
vals := make([]float64, r*c)
|
||
|
|
for i := range vals {
|
||
|
|
vals[i] = a.FloatAt(base + i)
|
||
|
|
}
|
||
|
|
return core.FromFloats(vals, r, c)
|
||
|
|
}
|
||
|
|
|
||
|
|
// copyInto writes src's elements at dst's flat offset. The
|
||
|
|
// destination is a fresh contiguous float64 accumulator; a contiguous
|
||
|
|
// float64 source moves with one copy, a float32 one widens in place.
|
||
|
|
func copyInto(dst, src *core.Array, offset int) {
|
||
|
|
switch {
|
||
|
|
case !src.Strided() && src.Dtype() == core.Float:
|
||
|
|
copy(dst.RawFloats()[offset:], src.RawFloats())
|
||
|
|
case !src.Strided() && src.Dtype() == core.Float32:
|
||
|
|
// Bounded by the source: the destination tail runs on to the
|
||
|
|
// end of the accumulator, which is longer for every batch but
|
||
|
|
// the last.
|
||
|
|
ss, ds := src.RawFloat32s(), dst.RawFloats()[offset:]
|
||
|
|
for i := range src.Len() {
|
||
|
|
ds[i] = float64(ss[i])
|
||
|
|
}
|
||
|
|
default:
|
||
|
|
for i := range src.Len() {
|
||
|
|
dst.SetFloatAt(offset+i, src.FloatAt(i))
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|