Files
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

1346 lines
43 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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
}