1346 lines
43 KiB
Go
1346 lines
43 KiB
Go
// 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
|
||
}
|