Files
tensor/internal/core/mat.go
T

1346 lines
43 KiB
Go
Raw 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 core
import (
"sourcedock.dev/petrbalvin/tensor/internal/engine"
"sync"
)
// Split thresholds for the matmul, transpose and gather kernels: below
// them the work is smaller than the goroutine fan-out it would pay
// for. Every floor is a measured constraint: the sweep benchmarks pin
// the last size where the split loses to the calling goroutine and the
// first where it wins, and the constant sits between the two.
const (
// matVecParallelMin is the minimum element work (n·k or k·m) one
// worker must carry before a 2-D×1-D or 1-D×2-D product spreads
// across workers, so the fan-out is capped by how much work each
// worker gets rather than by the total. Measured on the vector
// sweep: a dot product costs a couple of cycles per element, and a
// worker handed a few hundred of them pays more for its goroutine
// than it saves, so a 16,384-element product split thirty-two ways
// measures up to three times slower than the lone walk, while the
// same product on four workers matches it.
matVecParallelMin = 16_384
// matMulParallelMin is the minimum element-update work (n·k·m)
// before the Int, Float32 and Float 2-D×2-D products spread their
// output rows across workers. Measured on the square sweep: at
// 48 cubed (110,592 updates) the 32-way row split loses 2 to 22
// percent against the lone wide-panel walk, at 64 cubed (262,144)
// it wins 5 to 11 percent, and larger products keep widening the
// gap; 131,072 sits between the losing and winning sizes.
matMulParallelMin = 131_072
// matMulComplexParallelMin is the Complex path's floor on the
// same n·k·m measure. A complex element-update costs about four
// real multiplies, so the fan-out pays for itself four times
// earlier: at 16 cubed (4,096 updates) the split still loses 47
// percent, at 32 cubed (32,768) it already wins 33 percent.
matMulComplexParallelMin = 32_768
// colParallelMin is the row count from which a column gather runs
// across workers. The gather reads one strided element per row,
// so the row count is the measure: shapes with equal rows·cols
// work split opposite ways (32,768 rows of 8 columns wins, 4,096
// rows of 64 columns loses 3.7 times). Measured down the row
// axis: 16,384 rows still lose 7 to 31 percent, 32,768 rows win
// 9 to 16 percent.
colParallelMin = 32_768
// transposeSplitMin is the minimum element work (rows·cols)
// before the tiled transpose spreads its tiles across workers.
// Measured on the square and rectangular sweep: serial tiles win
// 1.5 to 2.3 times at and below 128×128, 256×256 and 512×512
// measure even, and of the two 65,536-element rectangles one
// wins and one loses by a fifth, so the split engages from that
// size up.
transposeSplitMin = 65_536
// transposeTileMin is the row and column count from which the
// transpose takes the tiled 2-D walk; smaller matrices, and every
// rank above two, stay on the odometer.
transposeTileMin = 8
// widenParallelMin is the minimum element count before the exact
// widening walks of denseFloatsLocal and complexPayload spread
// across workers. Measured on the conversion sweeps: 4,096
// elements run twice as fast serial, 16,384 elements win about
// 20 percent parallel, and both float64 and complex128 widening
// widen that gap with size.
widenParallelMin = 16_384
)
// Matrix operations with complex support. MatMul
// supports exactly 2-D×2-D, 2-D×1-D and 1-D×2-D; there is no batch
// broadcasting: the inner dimensions must agree or the error
// names both shapes. Transpose, Reshape, Row and Col all copy: results
// never alias their receiver.
// narrowRefused reports whether dt is one of the narrow element types
// the arithmetic surfaces do not take: bool and the small integers.
// Every entry point without a kernel for them refuses them loudly by
// name, so the caller converts with Astype instead of reading a
// payload nothing computed.
func narrowRefused(dt Dtype) bool {
return dt == Bool || narrowIntClass(dt)
}
// Identity builds the n×n identity matrix of the given dtype.
func Identity(dt Dtype, n int) (*Array, error) {
if n < 0 {
return nil, errf("Identity: n must be zero or greater, got %d", n)
}
out, err := Zeros(dt, n, n)
if err != nil {
return nil, err
}
for i := range n {
switch dt {
case Int:
out.ints[i*n+i] = 1
case Bool:
out.bools[i*n+i] = true
case Int8:
out.i8s[i*n+i] = 1
case Uint8:
out.u8s[i*n+i] = 1
case Int16:
out.i16s[i*n+i] = 1
case Uint16:
out.u16s[i*n+i] = 1
case Int32:
out.i32s[i*n+i] = 1
case Uint32:
out.u32s[i*n+i] = 1
case Float16:
out.halves[i*n+i] = halfOne
case Float32:
out.floats32[i*n+i] = 1
case Float:
out.floats[i*n+i] = 1
default:
out.complexes[i*n+i] = 1
}
}
return out, nil
}
// MatMul returns the matrix product of a and b: 2-D×2-D, a matrix times a
// vector, or a vector times a matrix. The result dtype follows the
// promotion ladder: int with int stays int, any float or complex operand
// promotes. Float16 operands are refused loudly: the product kernels are
// not offered for the half dtype yet, and Astype conversion is cheap. The
// narrow element types are refused the same way until their kernels
// arrive.
func MatMul2D(a, b *Array) (*Array, error) {
if a.Dtype() == Float16 || b.Dtype() == Float16 {
return nil, errf("MatMul: float16 operands are not supported; convert with Astype")
}
for _, op := range []*Array{a, b} {
if narrowRefused(op.Dtype()) {
return nil, errf("MatMul: dtype %s is not supported; convert with Astype", op.Dtype())
}
}
switch {
case a.NDim() == 2 && b.NDim() == 2:
n, k := a.Shape()[0], a.Shape()[1]
if b.Shape()[0] != k {
return nil, matmulMismatch(a, b)
}
m := b.Shape()[1]
out := &Array{shape: []int{n, m}}
// Every 2-D kernel preserves the plain i-p-j walk per cell:
// each output element receives exactly its products in
// ascending p order, so any split over rows is bit-identical
//. All dtypes take the row-panel kernels, which share
// one pass over b across a panel of output rows. Scalar
// mul-then-add sustains at most two floating-point operations
// per cycle per core, and the panel walk reaches that rate
// while streaming b linearly; measured against it, 2-D
// output-tile grids over register accumulators lose to
// per-element spill and bound-check overhead, and the row
// split's re-reads of b stay well inside the L3 bandwidth.
// The vector products split only above matVecParallelMin,
// where the fan-out pays for itself; the 2-D products
// gate their row split the same way through matMulSplit, so a
// small product runs whole on the calling goroutine, whose
// lone-worker walk reaches the wide panels a split would keep
// narrow.
work := n * k * m
switch promote(a.Dtype(), b.Dtype()) {
case Int:
out.dt = Int
out.ints = make([]int64, n*m)
matMulSplit(n, work, matMulParallelMin, func(rs, re int) {
matMulIntRows(a.ints, b.ints, out.ints, rs, re, k, m)
})
case Float32:
// Accumulate in float64 and round once: float32
// products are exact in float64, so the sums are never
// less accurate than a float32 kernel. The payloads are
// read directly, widening exactly where floatAt would, and
// each worker narrows its own finished rows once.
out.dt = Float32
out.floats32 = make([]float32, n*m)
aIs32 := a.Dtype() == Float32
bIs32 := b.Dtype() == Float32
matMulSplit(n, work, matMulParallelMin, func(rs, re int) {
panel := engine.GetFloat64Buf(4 * m)
defer engine.PutFloat64Buf(panel)
switch {
case aIs32 && bIs32:
matMulF32Rows(a.floats32, b.floats32, out.floats32, panel, rs, re, k, m)
case aIs32:
matMulF32Rows(a.floats32, b.ints, out.floats32, panel, rs, re, k, m)
case bIs32:
matMulF32Rows(a.ints, b.floats32, out.floats32, panel, rs, re, k, m)
default:
matMulF32Rows(a.ints, b.ints, out.floats32, panel, rs, re, k, m)
}
})
case Float:
out.dt = Float
out.floats = make([]float64, n*m)
// Mixed real operands convert once at entry, in the
// same element order, so the kernel streams plain rows.
// A float64 operand's payload already is that row
// stream, so it skips the conversion copy.
af := a.floats
if a.Dtype() != Float {
af = denseFloatsLocal(a, n, k)
}
bf := b.floats
if b.Dtype() != Float {
bf = denseFloatsLocal(b, k, m)
}
// A lone worker runs the wide 8-row panel; a split
// range runs the narrow one (see matMulF64Rows). A
// product below the floor is a lone worker too.
wide := work < matMulParallelMin || workersFor(n) == 1
matMulSplit(n, work, matMulParallelMin, func(rs, re int) {
matMulF64Rows(af, bf, out.floats, rs, re, k, m, wide)
})
default:
out.dt = Complex
out.complexes = make([]complex128, n*m)
ac := complexPayload(a)
bc := complexPayload(b)
matMulSplit(n, work, matMulComplexParallelMin, func(rs, re int) {
matMulComplexRows(ac, bc, out.complexes, rs, re, k, m)
})
}
return out, nil
case a.NDim() == 2 && b.NDim() == 1:
n, k := a.Shape()[0], a.Shape()[1]
if b.Len() != k {
return nil, matmulMismatch(a, b)
}
out := &Array{shape: []int{n}}
switch promote(a.Dtype(), b.Dtype()) {
case Int:
out.dt = Int
out.ints = make([]int64, n)
parallelSplit(n, n*k, func(rs, re int) {
matVecRows(a.ints, b.ints, out.ints, rs, re, k)
})
case Float32:
// Accumulate in float64 and round once. The raw
// payload walk applies for every real payload: each
// widening is exact, so the values and their order match
// accessor reads bit for bit.
out.dt = Float32
out.floats32 = make([]float32, n)
aIs32 := a.Dtype() == Float32
bIs32 := b.Dtype() == Float32
parallelSplit(n, n*k, func(rs, re int) {
switch {
case aIs32 && bIs32:
matVecF32Rows(a.floats32, b.floats32, out.floats32, rs, re, k)
case aIs32:
matVecF32Rows(a.floats32, b.ints, out.floats32, rs, re, k)
case bIs32:
matVecF32Rows(a.ints, b.floats32, out.floats32, rs, re, k)
default:
matVecF32Rows(a.ints, b.ints, out.floats32, rs, re, k)
}
})
case Float:
out.dt = Float
out.floats = make([]float64, n)
if a.Dtype() == Float && b.Dtype() == Float {
parallelSplit(n, n*k, func(rs, re int) {
matVecF64Rows(a.floats, b.floats, out.floats, rs, re, k)
})
} else {
// The mixed walk converts through the same accessor the
// scalar loop read, so the widened values and their
// order are unchanged.
af, bf := matVecF64Operands(a, n, k, b)
parallelSplit(n, n*k, func(rs, re int) {
matVecF64Rows(af, bf, out.floats, rs, re, k)
})
}
default:
out.dt = Complex
out.complexes = make([]complex128, n)
ac := complexPayload(a)
bc := complexPayload(b)
parallelSplit(n, n*k, func(rs, re int) {
matVecRows(ac, bc, out.complexes, rs, re, k)
})
}
return out, nil
case a.NDim() == 1 && b.NDim() == 2:
k := a.Len()
if b.Shape()[0] != k {
return nil, matmulMismatch(a, b)
}
m := b.Shape()[1]
out := &Array{shape: []int{m}}
switch promote(a.Dtype(), b.Dtype()) {
case Int:
out.dt = Int
out.ints = make([]int64, m)
parallelSplit(m, k*m, func(js, je int) {
vecMatCols(a.ints[:k], b.ints, out.ints[js:je], m, js, je)
})
case Float32:
// Accumulate in float64 and round once, exactly as
// the 2-D×1-D path does.
out.dt = Float32
out.floats32 = make([]float32, m)
aIs32 := a.Dtype() == Float32
bIs32 := b.Dtype() == Float32
parallelSplit(m, k*m, func(js, je int) {
acc := engine.GetFloat64Buf(je - js)
defer engine.PutFloat64Buf(acc)
window := out.floats32[js:je]
switch {
case aIs32 && bIs32:
vecMatF32Cols(a.floats32[:k], b.floats32, window, acc, m, js, je)
case aIs32:
vecMatF32Cols(a.floats32[:k], b.ints, window, acc, m, js, je)
case bIs32:
vecMatF32Cols(a.ints, b.floats32, window, acc, m, js, je)
default:
vecMatF32Cols(a.ints, b.ints, window, acc, m, js, je)
}
})
case Float:
out.dt = Float
out.floats = make([]float64, m)
if a.Dtype() == Float && b.Dtype() == Float {
parallelSplit(m, k*m, func(js, je int) {
vecMatF64Cols(a.floats[:k], b.floats, out.floats[js:je], m, js, je)
})
} else {
af, bf := vecMatF64Operands(a, k, b, m)
parallelSplit(m, k*m, func(js, je int) {
vecMatF64Cols(af, bf, out.floats[js:je], m, js, je)
})
}
default:
out.dt = Complex
out.complexes = make([]complex128, m)
ac := complexPayload(a)
bc := complexPayload(b)
parallelSplit(m, k*m, func(js, je int) {
vecMatCols(ac[:k], bc, out.complexes[js:je], m, js, je)
})
}
return out, nil
}
return nil, errf("MatMul: unsupported shapes %s and %s", shapeText(a.Shape()), shapeText(b.Shape()))
}
func matmulMismatch(a, b *Array) error {
return errf("MatMul: shape mismatch %s vs %s: the inner dimensions must agree",
shapeText(a.Shape()), shapeText(b.Shape()))
}
// parallelSplit runs fn over [0, n) with the vector products' worker
// policy: each worker must carry at least matVecParallelMin elements,
// and a workload that cannot fill one worker that far runs whole on the
// calling goroutine. The floor counts elements per worker rather than
// elements in total: a dot product is a couple of cycles per element,
// so a worker needs a few thousand of them before the goroutine's
// creation and its share of the join amortise, and a total-work floor
// would hand a bare 16,384-element product to thirty-two workers at a
// few hundred elements each, which measures up to three times slower
// than the lone walk.
func parallelSplit(n, work int, fn func(start, end int)) {
splitCapped(n, work, matVecParallelMin, fn)
}
// matMulSplit is the 2-D twin of parallelSplit: it runs fn over the n
// output rows of a 2-D×2-D product serially while the element-update
// work (n·k·m) stays below floor and across workers above it. The
// Complex path passes its own, lower floor because a complex
// element-update costs about four times the arithmetic of a real one.
func matMulSplit(n, work, floor int, fn func(start, end int)) {
if work < floor {
fn(0, n)
return
}
engine.Parallel(n, fn)
}
// splitCapped runs fn over the [0, items) range across at most
// work/floor workers, serially when the whole workload sits below floor
// or the capped fan-out falls to one. It is the scheduling primitive of
// the line-based folds (axis reductions, tiled transpose): those kernels
// stream about one element per instruction, so a spawned goroutine must
// carry tens of thousands of elements before its creation and
// synchronisation amortise, and uncapped fan-out pays that spawn cost
// core-count times on workloads a few workers swallow whole (measured on
// the axis and transpose sweeps: the 32-way split of a 256×256
// transpose or a 544-line sum costs more than the walk itself).
// Chunks stay contiguous and every item belongs to exactly one chunk,
// so a kernel whose accumulator slots are single-writer produces
// identical results under any split, including none.
func splitCapped(items, work, floor int, fn func(start, end int)) {
w := workersFor(items)
if work < floor {
w = 1
} else if byWork := work / floor; w > byWork {
w = max(byWork, 1)
}
if w <= 1 {
fn(0, items)
return
}
chunk := (items + w - 1) / w
var wg sync.WaitGroup
for ls := 0; ls < items; ls += chunk {
le := min(ls+chunk, items)
wg.Go(func() {
fn(ls, le)
})
}
wg.Wait()
}
// matMulF64Rows multiplies rows [rs, re) of the n×k matrix a by the
// k×m matrix b into the same rows of out, which must arrive zeroed.
// The kernel walks the p dimension with a panel of output rows sharing
// one pass over b: the innermost loop is a range walk over the b row,
// so the compiled loop is pure pointer arithmetic, and each output
// element still receives exactly the products a scalar walk would add,
// in ascending p order. The result is bit-identical down to signed
// zeros and NaN payloads. The wide flag picks the 8-row panel for a
// lone worker, where it has the whole L1 to itself; under a real
// worker split the SMT siblings share that L1, so the narrow 4-row
// panels win. Remainder rows fall to 2- and 1-row panels with the same
// per-element order.
func matMulF64Rows(a, b, out []float64, rs, re, k, m int, wide bool) {
matMulF64Tile(a, b, out, rs, re, k, m, 0, m, wide)
}
// matMulF64Tile multiplies rows [rs, re) of a by columns [js, je) of b
// into the same rows and columns of out, which must arrive zeroed. The
// whole product is the [0, n) × [0, m) case, and the tile form is what
// keeps a worker's slice of the output independent of every other
// worker's: rows are walked in 8-, 4-, 2- and 1-row panels and the
// columns of the tile in blocks of matMulColBlockMax, so the panel's
// output lines and one b row stay resident while the p walk streams b
// once per panel. Every output element receives its addends in
// ascending p order whatever tile carries it, so the result does not
// depend on the split.
func matMulF64Tile(a, b, out []float64, rs, re, k, m, js, je int, wide bool) {
step := min(je-js, matMulColBlockMax)
i := rs
if wide {
for ; i+8 <= re; i += 8 {
for j0 := js; j0 < je; j0 += step {
matMulF64Panel8(a, b, out, i, k, m, j0, min(j0+step, je))
}
}
}
for ; i+4 <= re; i += 4 {
for j0 := js; j0 < je; j0 += step {
matMulF64Panel4(a, b, out, i, k, m, j0, min(j0+step, je))
}
}
for ; i+2 <= re; i += 2 {
for j0 := js; j0 < je; j0 += step {
matMulF64Panel2(a, b, out, i, k, m, j0, min(j0+step, je))
}
}
for ; i < re; i++ {
for j0 := js; j0 < je; j0 += step {
matMulF64Panel1(a, b, out, i, k, m, j0, min(j0+step, je))
}
}
}
// matMulColBlockMax bounds the column span of one 8-row panel pass.
// Eight output rows occupy 64 bytes per column, so a row count beyond
// matMulColBlockMax would push the panel out of L1 even for a lone
// worker; wider matrices sweep m in blocks and keep a block of every
// row resident instead.
const matMulColBlockMax = 512
// matMulF64Panel8 multiplies the 8-row tile at i by the column block
// [j0, j1), accumulating straight into the zeroed output rows. The j
// walk runs over the symbolic block width w, which is also the length
// every row slice expression above declares, so the compiler drops
// every per-element bounds check.
func matMulF64Panel8(a, b, out []float64, i, k, m, j0, j1 int) {
a0 := a[i*k : i*k+k]
a1 := a[(i+1)*k : (i+2)*k]
a2 := a[(i+2)*k : (i+3)*k]
a3 := a[(i+3)*k : (i+4)*k]
a4 := a[(i+4)*k : (i+5)*k]
a5 := a[(i+5)*k : (i+6)*k]
a6 := a[(i+6)*k : (i+7)*k]
a7 := a[(i+7)*k : (i+8)*k]
o0 := out[i*m : (i+1)*m]
o1 := out[(i+1)*m : (i+2)*m]
o2 := out[(i+2)*m : (i+3)*m]
o3 := out[(i+3)*m : (i+4)*m]
o4 := out[(i+4)*m : (i+5)*m]
o5 := out[(i+5)*m : (i+6)*m]
o6 := out[(i+6)*m : (i+7)*m]
o7 := out[(i+7)*m : (i+8)*m]
w := j1 - j0
r0 := o0[j0:j1]
r1 := o1[j0:j1]
r2 := o2[j0:j1]
r3 := o3[j0:j1]
r4 := o4[j0:j1]
r5 := o5[j0:j1]
r6 := o6[j0:j1]
r7 := o7[j0:j1]
for p, av0 := range a0 {
av1, av2, av3 := a1[p], a2[p], a3[p]
av4, av5, av6, av7 := a4[p], a5[p], a6[p], a7[p]
brow := b[p*m+j0 : p*m+j1]
for j := range w {
bv := brow[j]
r0[j] = float64(av0*bv) + r0[j]
r1[j] = float64(av1*bv) + r1[j]
r2[j] = float64(av2*bv) + r2[j]
r3[j] = float64(av3*bv) + r3[j]
r4[j] = float64(av4*bv) + r4[j]
r5[j] = float64(av5*bv) + r5[j]
r6[j] = float64(av6*bv) + r6[j]
r7[j] = float64(av7*bv) + r7[j]
}
}
}
// matMulF64Panel4 is the 4-row tile: its j walk runs over the symbolic
// block width, so its per-element bounds checks drop the same way.
func matMulF64Panel4(a, b, out []float64, i, k, m, j0, j1 int) {
a0 := a[i*k : i*k+k]
a1 := a[(i+1)*k : (i+2)*k]
a2 := a[(i+2)*k : (i+3)*k]
a3 := a[(i+3)*k : (i+4)*k]
o0 := out[i*m : (i+1)*m]
o1 := out[(i+1)*m : (i+2)*m]
o2 := out[(i+2)*m : (i+3)*m]
o3 := out[(i+3)*m : (i+4)*m]
w := j1 - j0
r0 := o0[j0:j1]
r1 := o1[j0:j1]
r2 := o2[j0:j1]
r3 := o3[j0:j1]
for p, av0 := range a0 {
av1, av2, av3 := a1[p], a2[p], a3[p]
brow := b[p*m+j0 : p*m+j1]
for j := range w {
bv := brow[j]
r0[j] = float64(av0*bv) + r0[j]
r1[j] = float64(av1*bv) + r1[j]
r2[j] = float64(av2*bv) + r2[j]
r3[j] = float64(av3*bv) + r3[j]
}
}
}
// matMulF64Panel2 is the 2-row tile.
func matMulF64Panel2(a, b, out []float64, i, k, m, j0, j1 int) {
a0 := a[i*k : i*k+k]
a1 := a[(i+1)*k : (i+2)*k]
o0 := out[i*m : (i+1)*m]
o1 := out[(i+1)*m : (i+2)*m]
w := j1 - j0
r0 := o0[j0:j1]
r1 := o1[j0:j1]
for p, av0 := range a0 {
av1 := a1[p]
brow := b[p*m+j0 : p*m+j1]
for j := range w {
bv := brow[j]
r0[j] = float64(av0*bv) + r0[j]
r1[j] = float64(av1*bv) + r1[j]
}
}
}
// matMulF64Panel1 is the single-row walk. Every panel above writes the
// product through float64(a*b) + c: that spelling is what stops the
// compiler from contracting the multiply and the add into one fused
// operation at GOAMD64 v3 and later, which is what keeps the panels
// bit-for-bit identical on every machine and every pinned level.
func matMulF64Panel1(a, b, out []float64, i, k, m, j0, j1 int) {
arow := a[i*k : i*k+k]
orow := out[i*m : (i+1)*m]
r0 := orow[j0:j1]
for p, av := range arow {
brow := b[p*m+j0 : p*m+j1]
for j, bv := range brow {
r0[j] = float64(av*bv) + r0[j]
}
}
}
// matMulComplexRows multiplies rows [rs, re) of the complex n×k matrix
// a by the complex k×m matrix b into the same rows of out, with 4-row
// panels sharing one pass over b: the plain per-row walk reads the whole
// k×m b once for every output row, which for a complex payload is a
// sixteen-byte element stream four times the width of the arithmetic
// that consumes it. A panel quarters that stream and leaves the complex
// arithmetic as the limit. Every output element still receives its
// addends in ascending p order, so the sums are unchanged.
func matMulComplexRows(a, b, out []complex128, rs, re, k, m int) {
step := min(m, matMulComplexColBlockMax)
i := rs
for ; i+4 <= re; i += 4 {
for j0 := 0; j0 < m; j0 += step {
matMulComplexPanel4(a, b, out, i, k, m, j0, min(j0+step, m))
}
}
for ; i < re; i++ {
for j0 := 0; j0 < m; j0 += step {
matMulComplexPanel1(a, b, out, i, k, m, j0, min(j0+step, m))
}
}
}
// matMulComplexColBlockMax bounds the column span of one complex panel
// pass: a four-row panel of sixteen-byte elements fills 64 bytes per
// column, one cache line, so a wider block would push the accumulator
// rows out of L1 beside the b row band.
const matMulComplexColBlockMax = 256
// matMulComplexPanel4 is the 4-row complex tile.
func matMulComplexPanel4(a, b, out []complex128, i, k, m, j0, j1 int) {
a0 := a[i*k : i*k+k]
a1 := a[(i+1)*k : (i+2)*k]
a2 := a[(i+2)*k : (i+3)*k]
a3 := a[(i+3)*k : (i+4)*k]
o0 := out[i*m : (i+1)*m]
o1 := out[(i+1)*m : (i+2)*m]
o2 := out[(i+2)*m : (i+3)*m]
o3 := out[(i+3)*m : (i+4)*m]
r0 := o0[j0:j1]
r1 := o1[j0:j1]
r2 := o2[j0:j1]
r3 := o3[j0:j1]
for p, av0 := range a0 {
av1, av2, av3 := a1[p], a2[p], a3[p]
brow := b[p*m+j0 : p*m+j1]
for j, bv := range brow {
r0[j] += av0 * bv
r1[j] += av1 * bv
r2[j] += av2 * bv
r3[j] += av3 * bv
}
}
}
// matMulComplexPanel1 is the single-row complex walk.
func matMulComplexPanel1(a, b, out []complex128, i, k, m, j0, j1 int) {
arow := a[i*k : i*k+k]
orow := out[i*m : (i+1)*m]
r0 := orow[j0:j1]
for p, av := range arow {
brow := b[p*m+j0 : p*m+j1]
for j, bv := range brow {
r0[j] += av * bv
}
}
}
// matMulF32Rows is matMulF64Rows for a Float32 result. Float32 outputs
// cannot hold the running sums, so each tile accumulates in a float64
// panel buffer (one m-wide row per output row, supplied zeroed by the
// caller) and narrows each finished row once. The operands keep
// their real payloads; every widening here is the exact conversion the
// accessor performs, so the products and their order are unchanged. The
// panel sweeps the whole row rather than blocking its columns: the
// measured sweep keeps a full row for every panel that fits L1, and the
// float32 rows that do not fit lose more to the strided b walk a block
// would create than they gain from the smaller panel.
func matMulF32Rows[A, B int64 | float32](a []A, b []B, out []float32, panel []float64, rs, re, k, m int) {
i := rs
for ; i+4 <= re; i += 4 {
matMulF32Panel4(a, b, out, panel, i, k, m)
}
for ; i < re; i++ {
matMulF32Panel1(a, b, out, panel, i, k, m)
}
}
// matMulF32Panel4 multiplies the 4-row tile at i through the float64
// panel, then narrows the four rows in one pass.
func matMulF32Panel4[A, B int64 | float32](a []A, b []B, out []float32, panel []float64, i, k, m int) {
clear(panel)
a0 := a[i*k : i*k+k]
a1 := a[(i+1)*k : (i+2)*k]
a2 := a[(i+2)*k : (i+3)*k]
a3 := a[(i+3)*k : (i+4)*k]
s0 := panel[0:m]
s1 := panel[m : 2*m]
s2 := panel[2*m : 3*m]
s3 := panel[3*m : 4*m]
for p, av0 := range a0 {
av0f := float64(av0)
av1, av2, av3 := float64(a1[p]), float64(a2[p]), float64(a3[p])
brow := b[p*m : p*m+m]
for j, bv := range brow {
bvf := float64(bv)
s0[j] += av0f * bvf
s1[j] += av1 * bvf
s2[j] += av2 * bvf
s3[j] += av3 * bvf
}
}
o0 := out[i*m : (i+1)*m]
o1 := out[(i+1)*m : (i+2)*m]
o2 := out[(i+2)*m : (i+3)*m]
o3 := out[(i+3)*m : (i+4)*m]
for j, v := range s0 {
o0[j] = float32(v)
o1[j] = float32(s1[j])
o2[j] = float32(s2[j])
o3[j] = float32(s3[j])
}
}
// matMulF32Panel1 multiplies the single row at i through the float64
// panel and narrows it once.
func matMulF32Panel1[A, B int64 | float32](a []A, b []B, out []float32, panel []float64, i, k, m int) {
clear(panel)
arow := a[i*k : i*k+k]
s0 := panel[0:m]
for p, av := range arow {
avf := float64(av)
brow := b[p*m : p*m+m]
for j, bv := range brow {
s0[j] += avf * float64(bv)
}
}
orow := out[i*m : (i+1)*m]
for j, v := range s0 {
orow[j] = float32(v)
}
}
// matMulIntRows multiplies rows [rs, re) with 4-row panels sharing one
// pass over b. Int addition wraps exactly like the in-memory walk it
// replaces, and each column still receives its addends in ascending p
// order, so the wrapped sums are identical.
func matMulIntRows(a, b, out []int64, rs, re, k, m int) {
i := rs
for ; i+4 <= re; i += 4 {
a0 := a[i*k : i*k+k]
a1 := a[(i+1)*k : (i+2)*k]
a2 := a[(i+2)*k : (i+3)*k]
a3 := a[(i+3)*k : (i+4)*k]
o0 := out[i*m : (i+1)*m]
o1 := out[(i+1)*m : (i+2)*m]
o2 := out[(i+2)*m : (i+3)*m]
o3 := out[(i+3)*m : (i+4)*m]
for p, av0 := range a0 {
av1, av2, av3 := a1[p], a2[p], a3[p]
brow := b[p*m : p*m+m]
for j, bv := range brow {
o0[j] += av0 * bv
o1[j] += av1 * bv
o2[j] += av2 * bv
o3[j] += av3 * bv
}
}
}
for ; i < re; i++ {
arow := a[i*k : i*k+k]
orow := out[i*m : (i+1)*m]
for p, av := range arow {
brow := b[p*m : p*m+m]
for j, bv := range brow {
orow[j] += av * bv
}
}
}
}
// matVecF64Rows writes rows [rs, re) of a·b for the n×k matrix a and
// the k vector b, walking four rows at once: each row keeps its own
// ascending-p sum, so the four independent chains only add
// instruction-level parallelism, never reassociation.
func matVecF64Rows(a, b, out []float64, rs, re, k int) {
i := rs
for ; i+4 <= re; i += 4 {
a0 := a[i*k : i*k+k]
a1 := a[(i+1)*k : (i+2)*k]
a2 := a[(i+2)*k : (i+3)*k]
a3 := a[(i+3)*k : (i+4)*k]
var s0, s1, s2, s3 float64
for p := range k {
bv := b[p]
s0 += a0[p] * bv
s1 += a1[p] * bv
s2 += a2[p] * bv
s3 += a3[p] * bv
}
out[i], out[i+1], out[i+2], out[i+3] = s0, s1, s2, s3
}
for ; i < re; i++ {
arow := a[i*k : i*k+k]
var s float64
for p, av := range arow {
s += av * b[p]
}
out[i] = s
}
}
// matVecF32Rows is matVecF64Rows for a Float32 result: one float64 sum
// per row, narrowed once, read straight from the real
// payloads. Four rows walk together like the float64 twin, so the four
// independent sum chains fill the pipeline a single row's latency-bound
// chain leaves idle.
func matVecF32Rows[A, B int64 | float32](a []A, b []B, out []float32, rs, re, k int) {
i := rs
for ; i+4 <= re; i += 4 {
a0 := a[i*k : i*k+k]
a1 := a[(i+1)*k : (i+2)*k]
a2 := a[(i+2)*k : (i+3)*k]
a3 := a[(i+3)*k : (i+4)*k]
var s0, s1, s2, s3 float64
for p := range k {
bv := float64(b[p])
s0 += float64(a0[p]) * bv
s1 += float64(a1[p]) * bv
s2 += float64(a2[p]) * bv
s3 += float64(a3[p]) * bv
}
out[i], out[i+1], out[i+2], out[i+3] = float32(s0), float32(s1), float32(s2), float32(s3)
}
for ; i < re; i++ {
arow := a[i*k : i*k+k]
var s float64
for p, av := range arow {
s += float64(av) * float64(b[p])
}
out[i] = float32(s)
}
}
// matVecRows is the row-dot walk for the int and complex payloads;
// each row's sum is independent, so a worker split cannot move a
// single addend. The rows walk together, as in matVecF64Rows, but the
// int payload takes four rows and the complex payload two: a complex
// chain carries two register-wide values, so four of them fill the
// register file and spill exactly what the unroll was hiding.
func matVecRows[T int64 | complex128](a, b, out []T, rs, re, k int) {
var zero T
rows := 4
if _, isInt := any(zero).(int64); !isInt {
rows = 2
}
i := rs
if rows == 4 {
for ; i+4 <= re; i += 4 {
a0 := a[i*k : i*k+k]
a1 := a[(i+1)*k : (i+2)*k]
a2 := a[(i+2)*k : (i+3)*k]
a3 := a[(i+3)*k : (i+4)*k]
var s0, s1, s2, s3 T
for p := range k {
bv := b[p]
s0 += a0[p] * bv
s1 += a1[p] * bv
s2 += a2[p] * bv
s3 += a3[p] * bv
}
out[i], out[i+1], out[i+2], out[i+3] = s0, s1, s2, s3
}
} else {
for ; i+2 <= re; i += 2 {
a0 := a[i*k : i*k+k]
a1 := a[(i+1)*k : (i+2)*k]
var s0, s1 T
for p := range k {
bv := b[p]
s0 += a0[p] * bv
s1 += a1[p] * bv
}
out[i], out[i+1] = s0, s1
}
}
for ; i < re; i++ {
arow := a[i*k : i*k+k]
var s T
for p, av := range arow {
s += av * b[p]
}
out[i] = s
}
}
// matVecF64Operands resolves a 2-D×1-D product's operands to float64
// row streams: a float64 payload already is one, everything else
// converts once through the accessor's exact widening.
func matVecF64Operands(a *Array, n, k int, b *Array) ([]float64, []float64) {
var af, bf []float64
if a.Dtype() == Float {
af = a.floats
} else {
af = denseFloatsLocal(a, n, k)
}
if b.Dtype() == Float {
bf = b.floats
} else {
bf = denseFloatsLocal(b, 1, k)
}
return af, bf
}
// vecMatF64Cols multiplies a by columns [js, je) of b into out, which
// is the zeroed [js, je) window of the result: the b row is a range
// walk, the compiled loop is pure pointer arithmetic, and every output
// column receives its addends in ascending p order.
func vecMatF64Cols(a, b, out []float64, m, js, je int) {
if je-js <= vecMatColBandMax {
vecMatF64ColsNarrow(a, b, out, m, js, je)
return
}
for p, av := range a {
brow := b[p*m+js : p*m+je]
for j, bv := range brow {
out[j] += av * bv
}
}
}
// vecMatColBandMax is the widest output band the column walk accumulates
// in registers, one cache line of float64. A band that narrow leaves the
// b rows' lines half read whatever the walk does, and then the per-p
// load and store of every accumulator in out costs more than the walk's
// own arithmetic: measured across the vector sweep, a sixteen-column
// output of long vectors runs fifteen times slower through the memory
// accumulators than through four register ones.
const vecMatColBandMax = 8
// vecMatF64ColsNarrow walks a narrow band through register
// accumulators, four columns at a time: each column still sums its
// addends in ascending p order, so the values are those the memory walk
// produces.
func vecMatF64ColsNarrow(a, b, out []float64, m, js, je int) {
w := je - js
j := 0
for ; j+4 <= w; j += 4 {
var s0, s1, s2, s3 float64
for p, av := range a {
brow := b[p*m+js+j : p*m+js+j+4]
s0 += av * brow[0]
s1 += av * brow[1]
s2 += av * brow[2]
s3 += av * brow[3]
}
out[j], out[j+1], out[j+2], out[j+3] = s0, s1, s2, s3
}
for ; j < w; j++ {
var s float64
for p, av := range a {
s += av * b[p*m+js+j]
}
out[j] = s
}
}
// vecMatF32Cols is vecMatF64Cols for a Float32 result: the window
// accumulates in a float64 scratch (cleared by the caller) and narrows
// once at the end.
func vecMatF32Cols[A, B int64 | float32](a []A, b []B, out []float32, acc []float64, m, js, je int) {
clear(acc)
for p, av := range a {
avf := float64(av)
brow := b[p*m+js : p*m+je]
for j, bv := range brow {
acc[j] += avf * float64(bv)
}
}
for j, v := range acc {
out[j] = float32(v)
}
}
// vecMatCols is the column walk for the int and complex payloads;
// each output column receives its addends in ascending p order, and
// int wrapping is unchanged by the visit order.
func vecMatCols[T int64 | complex128](a, b, out []T, m, js, je int) {
for p, av := range a {
brow := b[p*m+js : p*m+je]
for j, bv := range brow {
out[j] += av * bv
}
}
}
// vecMatF64Operands resolves a 1-D×2-D product's operands to float64
// row streams, the 1-D twin of matVecF64Operands.
func vecMatF64Operands(a *Array, k int, b *Array, m int) ([]float64, []float64) {
var af, bf []float64
if a.Dtype() == Float {
af = a.floats
} else {
af = denseFloatsLocal(a, 1, k)
}
if b.Dtype() == Float {
bf = b.floats
} else {
bf = denseFloatsLocal(b, k, m)
}
return af, bf
}
// complexPayload returns the array's elements as complex values,
// reading a complex payload directly and converting everything else
// once; every array this package produces stores elements at their
// flat payload index. The conversion walk splits across workers from
// widenParallelMin up: each element widens exactly as the scalar read
// does and lands in its own slot, so the values are unchanged.
func complexPayload(a *Array) []complex128 {
if a.Dtype() == Complex {
return a.complexes
}
out := make([]complex128, a.Len())
if len(out) < widenParallelMin {
for i := range out {
out[i] = a.ComplexAt(i)
}
return out
}
engine.Parallel(len(out), func(ws, we int) {
for i := ws; i < we; i++ {
out[i] = a.ComplexAt(i)
}
})
return out
}
// Transpose returns a new array with the dimensions reversed; on 1-D it
// is a copy. It is infallible.
func Transpose(a *Array) *Array {
sh := a.Shape()
d := len(sh)
newShape := make([]int, d)
for i := range newShape {
newShape[i] = sh[d-1-i]
}
out := &Array{shape: newShape, dt: a.dt}
n := a.Len()
out.alloc(n)
// Rank-2 arrays take the tiled walk: cache-sized tiles keep both
// sides' lines alive and the tile range spreads across workers.
// The permutation is unchanged, only the visit order, so the copy
// stays bit-identical. Every other rank, and matrices below
// transposeTileMin in either dimension, keep the odometer below.
if d == 2 && sh[0] >= transposeTileMin && sh[1] >= transposeTileMin {
switch a.dt {
case Int:
transposeTiles(a.ints, out.ints, sh[0], sh[1])
case Bool:
transposeTiles(a.bools, out.bools, sh[0], sh[1])
case Int8:
transposeTiles(a.i8s, out.i8s, sh[0], sh[1])
case Uint8:
transposeTiles(a.u8s, out.u8s, sh[0], sh[1])
case Int16:
transposeTiles(a.i16s, out.i16s, sh[0], sh[1])
case Uint16:
transposeTiles(a.u16s, out.u16s, sh[0], sh[1])
case Int32:
transposeTiles(a.i32s, out.i32s, sh[0], sh[1])
case Uint32:
transposeTiles(a.u32s, out.u32s, sh[0], sh[1])
case Float16:
transposeTiles(a.halves, out.halves, sh[0], sh[1])
case Float32:
transposeTiles(a.floats32, out.floats32, sh[0], sh[1])
case Float:
transposeTiles(a.floats, out.floats, sh[0], sh[1])
default:
transposeTiles(a.complexes, out.complexes, sh[0], sh[1])
}
return out
}
// Under the reversed shape the destination flat index is the
// column-major index of the source coordinates: each advancing
// dimension contributes its source stride, so dst = Σ coord[k]·sw[k]
// is the same mapping the materialised index table used to hold,
// computed one pass earlier. The dtype dispatch keeps setFrom's
// per-element switch out of the copy loops.
sw := make([]int, d)
for k := range sw {
sw[k] = 1
for _, s := range sh[:k] {
sw[k] *= s
}
}
coord := make([]int, d)
switch a.dt {
case Int:
for src := range n {
dst := 0
for k := range coord {
dst += coord[k] * sw[k]
}
out.ints[dst] = a.ints[src]
advanceOdometer(coord, sh)
}
case Float16:
for src := range n {
dst := 0
for k := range coord {
dst += coord[k] * sw[k]
}
out.halves[dst] = a.halves[src]
advanceOdometer(coord, sh)
}
case Float32:
for src := range n {
dst := 0
for k := range coord {
dst += coord[k] * sw[k]
}
out.floats32[dst] = a.floats32[src]
advanceOdometer(coord, sh)
}
case Float:
for src := range n {
dst := 0
for k := range coord {
dst += coord[k] * sw[k]
}
out.floats[dst] = a.floats[src]
advanceOdometer(coord, sh)
}
case Complex:
for src := range n {
dst := 0
for k := range coord {
dst += coord[k] * sw[k]
}
out.complexes[dst] = a.complexes[src]
advanceOdometer(coord, sh)
}
default:
// Bool and the narrow integer widths carry no per-dtype loop
// here; setFrom writes them with the same values in the same
// order.
for src := range n {
dst := 0
for k := range coord {
dst += coord[k] * sw[k]
}
out.setFrom(dst, a, src)
advanceOdometer(coord, sh)
}
}
return out
}
// transposeTiles copies the rows×cols matrix src into its transpose dst
// through 32×32 tiles. The tile size keeps both sides' working set
// inside L1; the visit order is free because every destination slot is
// written exactly once. The tile range splits across workers only from
// transposeSplitMin up, and the fan-out is capped by splitCapped so
// every worker carries at least that many elements: below the floor the
// measured sweep keeps the whole walk faster on the calling goroutine,
// which still takes the tiled route, only alone.
func transposeTiles[T any](src, dst []T, rows, cols int) {
const tile = 32
colTiles := (cols + tile - 1) / tile
tiles := ((rows + tile - 1) / tile) * colTiles
walk := func(ts, te int) {
var stage [tile * tile]T
for t := ts; t < te; t++ {
i0 := (t / colTiles) * tile
j0 := (t % colTiles) * tile
i1 := min(i0+tile, rows)
j1 := min(j0+tile, cols)
w := j1 - j0
h := i1 - i0
// Both proofs sit outside the loops they serve: every stage
// read below lands inside the h×w tile and every drow store
// inside its h-element run, so the inner walks run without
// per-element bounds checks.
_ = stage[h*w-1]
for i := i0; i < i1; i++ {
srow := src[i*cols+j0 : i*cols+j1]
copy(stage[(i-i0)*w:(i-i0)*w+w], srow)
}
for j := j0; j < j1; j++ {
drow := dst[j*rows+i0 : j*rows+i1]
_ = drow[h-1]
// si steps down the stage column: the tile's transposed
// neighbours sit w slots apart.
si := j - j0
for i := range h {
drow[i] = stage[si]
si += w
}
}
}
}
splitCapped(tiles, rows*cols, transposeSplitMin, walk)
}
// Reshape returns a copy with a new shape of the same element count.
func Reshape(a *Array, shape ...int) (*Array, error) {
total, sh, err := checkedDims(shape)
if err != nil {
return nil, err
}
if total != a.Len() {
return nil, errf("Reshape: %d elements do not fill the shape %s", a.Len(), shapeText(sh))
}
// cloneArray carries every payload the dtype owns: the narrow
// element types have no five-slice clone to build from.
out := a.cloneArray()
out.shape = sh
return out, nil
}
// Row returns a 1-D copy of row i; the array must be 2-D.
func Row(a *Array, i int) (*Array, error) {
if a.NDim() != 2 {
return nil, errf("Row: needs a 2-D array, got shape %s", shapeText(a.Shape()))
}
if i < 0 || i >= a.Shape()[0] {
return nil, errf("Row: index %d is out of range for %d rows", i, a.Shape()[0])
}
cols := a.Shape()[1]
out := &Array{shape: []int{cols}, dt: a.Dtype()}
out.alloc(cols)
// A row is one contiguous payload block, so each dtype copies it in
// a single move.
switch a.dt {
case Int:
copy(out.ints, a.ints[i*cols:(i+1)*cols])
case Bool:
copy(out.bools, a.bools[i*cols:(i+1)*cols])
case Int8:
copy(out.i8s, a.i8s[i*cols:(i+1)*cols])
case Uint8:
copy(out.u8s, a.u8s[i*cols:(i+1)*cols])
case Int16:
copy(out.i16s, a.i16s[i*cols:(i+1)*cols])
case Uint16:
copy(out.u16s, a.u16s[i*cols:(i+1)*cols])
case Int32:
copy(out.i32s, a.i32s[i*cols:(i+1)*cols])
case Uint32:
copy(out.u32s, a.u32s[i*cols:(i+1)*cols])
case Float16:
copy(out.halves, a.halves[i*cols:(i+1)*cols])
case Float32:
copy(out.floats32, a.floats32[i*cols:(i+1)*cols])
case Float:
copy(out.floats, a.floats[i*cols:(i+1)*cols])
default:
copy(out.complexes, a.complexes[i*cols:(i+1)*cols])
}
return out, nil
}
// Col returns a 1-D copy of column j; the array must be 2-D.
func Col(a *Array, j int) (*Array, error) {
if a.NDim() != 2 {
return nil, errf("Col: needs a 2-D array, got shape %s", shapeText(a.Shape()))
}
if j < 0 || j >= a.Shape()[1] {
return nil, errf("Col: index %d is out of range for %d columns", j, a.Shape()[1])
}
rows, cols := a.Shape()[0], a.Shape()[1]
out := &Array{shape: []int{rows}, dt: a.Dtype()}
out.alloc(rows)
// Column reads stride by the row width; the dtype dispatch keeps
// the per-element accessor switch out of the gather, and tall
// matrices spread the strided reads across workers.
switch a.dt {
case Int:
colGather(a.ints, out.ints, rows, cols, j)
case Bool:
colGather(a.bools, out.bools, rows, cols, j)
case Int8:
colGather(a.i8s, out.i8s, rows, cols, j)
case Uint8:
colGather(a.u8s, out.u8s, rows, cols, j)
case Int16:
colGather(a.i16s, out.i16s, rows, cols, j)
case Uint16:
colGather(a.u16s, out.u16s, rows, cols, j)
case Int32:
colGather(a.i32s, out.i32s, rows, cols, j)
case Uint32:
colGather(a.u32s, out.u32s, rows, cols, j)
case Float16:
colGather(a.halves, out.halves, rows, cols, j)
case Float32:
colGather(a.floats32, out.floats32, rows, cols, j)
case Float:
colGather(a.floats, out.floats, rows, cols, j)
default:
colGather(a.complexes, out.complexes, rows, cols, j)
}
return out, nil
}
// colGather copies column j into dst, serially below colParallelMin
// and across workers above it: the gather is a pure permutation, so
// the split cannot move a single element. The floor counts rows, not
// rows·cols: the walk reads one strided element per row, and the sweep
// splits equal-work shapes opposite ways down the row axis.
func colGather[T any](src, dst []T, rows, cols, j int) {
if rows < colParallelMin {
for r := range rows {
dst[r] = src[r*cols+j]
}
return
}
engine.Parallel(rows, func(rs, re int) {
for r := rs; r < re; r++ {
dst[r] = src[r*cols+j]
}
})
}
// denseFloatsLocal copies an array's elements into a flat float64
// slice, widening int and float32 elements exactly. The walk splits
// across workers by rows from widenParallelMin up: every element
// widens exactly as the scalar read does and lands in its own slot,
// so the row streams are unchanged.
func denseFloatsLocal(a *Array, rows, cols int) []float64 {
out := make([]float64, rows*cols)
if rows*cols < widenParallelMin {
for i := range out {
out[i] = a.FloatAt(i)
}
return out
}
engine.Parallel(rows, func(rs, re int) {
for r := rs; r < re; r++ {
row := out[r*cols : (r+1)*cols]
for c := range row {
row[c] = a.FloatAt(r*cols + c)
}
}
})
return out
}