// Copyright (c) 2026 Petr Balvín (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)) } } }