Files

166 lines
5.1 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"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))
}
}
}