feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,639 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package core
|
||||
|
||||
import (
|
||||
"math"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/engine"
|
||||
)
|
||||
|
||||
// The general engine's slot walk and the axis folds the reduction-only
|
||||
// specs dispatch to: their benchmarks, and the oracles that pin the
|
||||
// walk's bits.
|
||||
|
||||
// einsumSlotSumOdometer steps every summed axis through one shared
|
||||
// odometer. It is the oracle for the peeled kernel: a slot must hold
|
||||
// exactly what this walk puts there, bit for bit.
|
||||
func einsumSlotSumOdometer[T int64 | float64 | complex128](rv [][]T, out []T, slot int, base, off, coord, delta, sizes []int, total int) {
|
||||
nSum := len(sizes)
|
||||
copy(off, base)
|
||||
clear(coord[:nSum])
|
||||
acc := T(1)
|
||||
for p := range rv {
|
||||
acc *= rv[p][off[p]]
|
||||
}
|
||||
out[slot] += acc
|
||||
for range total - 1 {
|
||||
t := nSum - 1
|
||||
for t >= 0 {
|
||||
coord[t]++
|
||||
for p := range rv {
|
||||
off[p] += delta[p*nSum+t]
|
||||
}
|
||||
if coord[t] < sizes[t] {
|
||||
break
|
||||
}
|
||||
coord[t] = 0
|
||||
for p := range rv {
|
||||
off[p] -= delta[p*nSum+t] * sizes[t]
|
||||
}
|
||||
t--
|
||||
}
|
||||
acc = T(1)
|
||||
for p := range rv {
|
||||
acc *= rv[p][off[p]]
|
||||
}
|
||||
out[slot] += acc
|
||||
}
|
||||
}
|
||||
|
||||
// einsumSlotSumF32Odometer is einsumSlotSumOdometer for a float32
|
||||
// result, narrowing every addend into the slot as it arrives.
|
||||
func einsumSlotSumF32Odometer(rv [][]float64, out []float32, slot int, base, off, coord, delta, sizes []int, total int) {
|
||||
nSum := len(sizes)
|
||||
copy(off, base)
|
||||
clear(coord[:nSum])
|
||||
acc := 1.0
|
||||
for p := range rv {
|
||||
acc *= rv[p][off[p]]
|
||||
}
|
||||
out[slot] = float32(float64(out[slot]) + acc)
|
||||
for range total - 1 {
|
||||
t := nSum - 1
|
||||
for t >= 0 {
|
||||
coord[t]++
|
||||
for p := range rv {
|
||||
off[p] += delta[p*nSum+t]
|
||||
}
|
||||
if coord[t] < sizes[t] {
|
||||
break
|
||||
}
|
||||
coord[t] = 0
|
||||
for p := range rv {
|
||||
off[p] -= delta[p*nSum+t] * sizes[t]
|
||||
}
|
||||
t--
|
||||
}
|
||||
acc = 1.0
|
||||
for p := range rv {
|
||||
acc *= rv[p][off[p]]
|
||||
}
|
||||
out[slot] = float32(float64(out[slot]) + acc)
|
||||
}
|
||||
}
|
||||
|
||||
// einsumProbeFloats fills n elements with distinct values of mixed
|
||||
// magnitude: any visit order other than the reference's moves addends
|
||||
// across a rounding boundary and changes the low bits of the sum.
|
||||
func einsumProbeFloats(n int) []float64 {
|
||||
v := make([]float64, n)
|
||||
for i := range v {
|
||||
v[i] = float64(i%11)*1.5 - 4 + math.Ldexp(1, i%17-8) + float64(i)/512
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
// einsumSlotConfig is one slot-walk configuration: the axis extents, the
|
||||
// per-operand cursor and the per-operand, per-axis strides.
|
||||
type einsumSlotConfig struct {
|
||||
sizes []int
|
||||
base []int
|
||||
delta []int
|
||||
reach []int
|
||||
}
|
||||
|
||||
// configs walks the operand counts and the summed-axis counts, giving
|
||||
// every kernel shape a configuration with in-bounds cursors; the
|
||||
// zero-summed-axis shapes come last.
|
||||
func einsumSlotConfigs() []einsumSlotConfig {
|
||||
var out []einsumSlotConfig
|
||||
for nOps := 1; nOps <= 5; nOps++ {
|
||||
for nSum := 1; nSum <= 3; nSum++ {
|
||||
cfg := einsumSlotConfig{
|
||||
sizes: make([]int, nSum),
|
||||
base: make([]int, nOps),
|
||||
delta: make([]int, nOps*nSum),
|
||||
reach: make([]int, nOps),
|
||||
}
|
||||
for t := range nSum {
|
||||
cfg.sizes[t] = 2 + (t*5+nOps)%3
|
||||
}
|
||||
for p := range nOps {
|
||||
cfg.base[p] = p
|
||||
cfg.reach[p] = p
|
||||
for t := range nSum {
|
||||
cfg.delta[p*nSum+t] = 1 + (p*3+t)%4
|
||||
cfg.reach[p] += (cfg.sizes[t] - 1) * cfg.delta[p*nSum+t]
|
||||
}
|
||||
}
|
||||
out = append(out, cfg)
|
||||
}
|
||||
}
|
||||
for nOps := 1; nOps <= 4; nOps++ {
|
||||
cfg := einsumSlotConfig{
|
||||
base: make([]int, nOps),
|
||||
delta: make([]int, nOps),
|
||||
reach: make([]int, nOps),
|
||||
}
|
||||
for p := range nOps {
|
||||
cfg.base[p] = p
|
||||
cfg.reach[p] = p
|
||||
}
|
||||
out = append(out, cfg)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// slotTotal is the number of visits one slot's walk takes.
|
||||
func (c einsumSlotConfig) slotTotal() int {
|
||||
total := 1
|
||||
for _, s := range c.sizes {
|
||||
total *= s
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
// name spells the configuration for a failure message.
|
||||
func (c einsumSlotConfig) name() string {
|
||||
return "ops=" + strconv.Itoa(len(c.base)) + "/sum=" + strconv.Itoa(len(c.sizes))
|
||||
}
|
||||
|
||||
// einsumCheckSlotKernel runs one configuration through the peeled
|
||||
// kernel and the odometer reference for a payload type and compares the
|
||||
// two slots by the given identity, so every operand count and cursor
|
||||
// reaches both loops.
|
||||
func einsumCheckSlotKernel[T int64 | float64 | complex128](
|
||||
t *testing.T, label string, cfg einsumSlotConfig,
|
||||
mk func(reach []int) [][]T, same func(T, T) bool,
|
||||
) {
|
||||
t.Helper()
|
||||
rv := mk(cfg.reach)
|
||||
got := make([]T, 1)
|
||||
want := make([]T, 1)
|
||||
nSum := len(cfg.sizes)
|
||||
einsumSlotSum(rv, got, 0, cfg.base, make([]int, len(cfg.base)), make([]int, nSum), cfg.delta, cfg.sizes, cfg.slotTotal())
|
||||
einsumSlotSumOdometer(rv, want, 0, cfg.base, make([]int, len(cfg.base)), make([]int, nSum), cfg.delta, cfg.sizes, cfg.slotTotal())
|
||||
if !same(got[0], want[0]) {
|
||||
t.Fatalf("%s %s: slot %v, want %v", label, cfg.name(), got[0], want[0])
|
||||
}
|
||||
}
|
||||
|
||||
// einssumCheckSlotKernelF32 pins the float32 slot kernel, which narrows
|
||||
// every addend into the slot, against its own odometer form.
|
||||
func einssumCheckSlotKernelF32(t *testing.T, cfg einsumSlotConfig, mk func(reach []int) [][]float64) {
|
||||
t.Helper()
|
||||
rv := mk(cfg.reach)
|
||||
got := make([]float32, 1)
|
||||
want := make([]float32, 1)
|
||||
nSum := len(cfg.sizes)
|
||||
einsumSlotSumF32(rv, got, 0, cfg.base, make([]int, len(cfg.base)), make([]int, nSum), cfg.delta, cfg.sizes, cfg.slotTotal())
|
||||
einsumSlotSumF32Odometer(rv, want, 0, cfg.base, make([]int, len(cfg.base)), make([]int, nSum), cfg.delta, cfg.sizes, cfg.slotTotal())
|
||||
if math.Float32bits(got[0]) != math.Float32bits(want[0]) {
|
||||
t.Fatalf("float32 %s: slot %v, want %v", cfg.name(), got[0], want[0])
|
||||
}
|
||||
}
|
||||
|
||||
// floatsPerReach builds one probe reader per operand, each long enough
|
||||
// for that operand's whole walk.
|
||||
func floatsPerReach(reach []int) [][]float64 {
|
||||
rv := make([][]float64, len(reach))
|
||||
for p, r := range reach {
|
||||
rv[p] = einsumProbeFloats(r + 1)
|
||||
}
|
||||
return rv
|
||||
}
|
||||
|
||||
// TestEinsumSlotSumMatchesOdometer pins the peeled slot kernels against
|
||||
// the odometer form across operand counts, summed-axis counts, axis
|
||||
// extents and cursors, on every payload type the kernels serve. The
|
||||
// distinct probe values make a moved addend visible in the low bits.
|
||||
func TestEinsumSlotSumMatchesOdometer(t *testing.T) {
|
||||
configs := einsumSlotConfigs()
|
||||
for _, cfg := range configs {
|
||||
// The generic kernel, instantiated for each payload type.
|
||||
einsumCheckSlotKernel(t, "float", cfg, floatsPerReach,
|
||||
func(a, b float64) bool { return math.Float64bits(a) == math.Float64bits(b) })
|
||||
einsumCheckSlotKernel(t, "int", cfg,
|
||||
func(reach []int) [][]int64 {
|
||||
rv := make([][]int64, len(reach))
|
||||
for p, r := range reach {
|
||||
rv[p] = make([]int64, r+1)
|
||||
for i := range rv[p] {
|
||||
rv[p][i] = int64(i%13)*7 - 9
|
||||
}
|
||||
}
|
||||
return rv
|
||||
}, func(a, b int64) bool { return a == b })
|
||||
einsumCheckSlotKernel(t, "complex", cfg,
|
||||
func(reach []int) [][]complex128 {
|
||||
rv := make([][]complex128, len(reach))
|
||||
for p, r := range reach {
|
||||
rv[p] = make([]complex128, r+1)
|
||||
for i := range rv[p] {
|
||||
rv[p][i] = complex(float64(i%7)-3, float64(i%5)-2)
|
||||
}
|
||||
}
|
||||
return rv
|
||||
}, func(a, b complex128) bool {
|
||||
return math.Float64bits(real(a)) == math.Float64bits(real(b)) &&
|
||||
math.Float64bits(imag(a)) == math.Float64bits(imag(b))
|
||||
})
|
||||
einssumCheckSlotKernelF32(t, cfg, floatsPerReach)
|
||||
}
|
||||
}
|
||||
|
||||
// einsumOperand builds a deterministic contiguous operand of the given
|
||||
// dtype.
|
||||
func einsumOperand(tb testing.TB, dt Dtype, shape ...int) *Array {
|
||||
tb.Helper()
|
||||
n := 1
|
||||
for _, d := range shape {
|
||||
n *= d
|
||||
}
|
||||
var a *Array
|
||||
var err error
|
||||
switch dt {
|
||||
case Int:
|
||||
v := make([]int64, n)
|
||||
for i := range v {
|
||||
v[i] = int64(i%9)*3 - 4 + int64(i)*7
|
||||
}
|
||||
a, err = FromInts(v, shape...)
|
||||
case Float32:
|
||||
v := make([]float32, n)
|
||||
for i := range v {
|
||||
v[i] = float32(i%7)*0.5 - 2 + float32(i)/64
|
||||
}
|
||||
a, err = FromFloat32s(v, shape...)
|
||||
case Complex:
|
||||
v := make([]complex128, n)
|
||||
for i := range v {
|
||||
v[i] = complex(float64(i%5)-2, float64(i%3)-1)
|
||||
}
|
||||
a, err = FromComplexes(v, shape...)
|
||||
default:
|
||||
a, err = FromFloats(einsumProbeFloats(n), shape...)
|
||||
}
|
||||
if err != nil {
|
||||
tb.Fatalf("operand %s %v: %v", dt, shape, err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// einsumStridedOperand builds a read-only view of the given shape whose
|
||||
// axes step by the given strides over a longer payload.
|
||||
func einsumStridedOperand(t testing.TB, dt Dtype, shape, strides []int) *Array {
|
||||
t.Helper()
|
||||
n := 1
|
||||
for d, s := range shape {
|
||||
n += (s - 1) * strides[d]
|
||||
}
|
||||
flat := einsumOperand(t, dt, n)
|
||||
view := &Array{shape: slices.Clone(shape), dt: dt, strides: slices.Clone(strides)}
|
||||
switch dt {
|
||||
case Int:
|
||||
view.ints = flat.RawInts()
|
||||
case Float32:
|
||||
view.floats32 = flat.RawFloat32s()
|
||||
case Complex:
|
||||
view.complexes = flat.RawComplexes()
|
||||
default:
|
||||
view.floats = flat.RawFloats()
|
||||
}
|
||||
return view
|
||||
}
|
||||
|
||||
// einsumBitsEqual compares two results by dtype, shape and raw payload
|
||||
// bits, so a NaN matches only its own bit pattern.
|
||||
func einsumBitsEqual(a, b *Array) bool {
|
||||
if a.Dtype() != b.Dtype() || !slices.Equal(a.Shape(), b.Shape()) {
|
||||
return false
|
||||
}
|
||||
switch a.dt {
|
||||
case Int:
|
||||
return slices.Equal(a.RawInts(), b.RawInts())
|
||||
case Float32:
|
||||
return slices.EqualFunc(a.RawFloat32s(), b.RawFloat32s(), func(x, y float32) bool {
|
||||
return math.Float32bits(x) == math.Float32bits(y)
|
||||
})
|
||||
case Float:
|
||||
return slices.EqualFunc(a.RawFloats(), b.RawFloats(), func(x, y float64) bool {
|
||||
return math.Float64bits(x) == math.Float64bits(y)
|
||||
})
|
||||
case Complex:
|
||||
return slices.EqualFunc(a.RawComplexes(), b.RawComplexes(), func(x, y complex128) bool {
|
||||
return math.Float64bits(real(x)) == math.Float64bits(real(y)) &&
|
||||
math.Float64bits(imag(x)) == math.Float64bits(imag(y))
|
||||
})
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// TestEinsumSlotSplitMatchesSerial runs the whole dispatch with the
|
||||
// worker pool as it stands and with it pinned to one worker, which
|
||||
// walks every slot in one chunk. The slot walk must write the same bits
|
||||
// either way, for every operand count, dtype and stride pattern the
|
||||
// engine accepts.
|
||||
func TestEinsumSlotSplitMatchesSerial(t *testing.T) {
|
||||
cases := []struct {
|
||||
spec string
|
||||
shapes [][]int
|
||||
}{
|
||||
{"...ij,jk->...ik", [][]int{{4, 6, 5}, {5, 3}}},
|
||||
{"...ij,...jk->...ik", [][]int{{3, 4, 5}, {3, 5, 2}}},
|
||||
{"ik,kj,jl->il", [][]int{{6, 5}, {5, 4}, {4, 7}}},
|
||||
{"ijk->kji", [][]int{{3, 4, 5}}},
|
||||
{"ij,jk,kl->il", [][]int{{5, 4}, {4, 6}, {6, 3}}},
|
||||
{"ij,jk,lk->il", [][]int{{5, 4}, {4, 6}, {3, 6}}},
|
||||
{"i,...i->...", [][]int{{5}, {3, 5}}},
|
||||
{"...i,...i->...", [][]int{{2, 4}, {4}}},
|
||||
{"kji->k", [][]int{{3, 4, 5}}},
|
||||
{"iij->ij", [][]int{{4, 4, 3}}},
|
||||
{"ii->i", [][]int{{5, 5}}},
|
||||
{"...->...", [][]int{{3, 4}}},
|
||||
{"ij->i", [][]int{{7, 5}}},
|
||||
{"ab,bc->ac", [][]int{{6, 5}, {5, 4}}},
|
||||
{"ij,ij->", [][]int{{6, 5}, {6, 5}}},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
for _, dt := range []Dtype{Int, Float32, Float, Complex} {
|
||||
ops := make([]*Array, len(tc.shapes))
|
||||
for i, shape := range tc.shapes {
|
||||
ops[i] = einsumOperand(t, dt, shape...)
|
||||
}
|
||||
want, werr := Einsum(tc.spec, ops...)
|
||||
restore := engine.SetNumWorkers(1)
|
||||
got, gerr := Einsum(tc.spec, ops...)
|
||||
engine.SetNumWorkers(restore)
|
||||
if werr != nil || gerr != nil {
|
||||
t.Fatalf("%s/%s: split error %v, serial error %v", tc.spec, dt, werr, gerr)
|
||||
}
|
||||
if !einsumBitsEqual(got, want) {
|
||||
t.Fatalf("%s/%s: split result %s differs from the serial walk %s",
|
||||
tc.spec, dt, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
// A strided operand is gathered through the accessor on every dtype
|
||||
// the engine widens; the walk must read the view's own elements.
|
||||
for _, dt := range []Dtype{Int, Float32, Float, Complex} {
|
||||
ops := []*Array{
|
||||
einsumStridedOperand(t, dt, []int{3, 4}, []int{8, 2}),
|
||||
einsumOperand(t, dt, 3, 4),
|
||||
}
|
||||
want, werr := Einsum("ij,ij->i", ops...)
|
||||
restore := engine.SetNumWorkers(1)
|
||||
got, gerr := Einsum("ij,ij->i", ops...)
|
||||
engine.SetNumWorkers(restore)
|
||||
if werr != nil || gerr != nil {
|
||||
t.Fatalf("strided/%s: split error %v, serial error %v", dt, werr, gerr)
|
||||
}
|
||||
if !einsumBitsEqual(got, want) {
|
||||
t.Fatalf("strided/%s: split result %s differs from the serial walk %s", dt, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestEinsumSumAxesMatchesEngine pins the reduction-only dispatch to the
|
||||
// general engine: "ij->i" and its relatives must fold to the same bits
|
||||
// the engine's walk produced, for every shape, axis order and dtype. The
|
||||
// float32 and complex operands stay with the engine by design and are
|
||||
// compared here as well, so a later change to the dispatch cannot move
|
||||
// them silently.
|
||||
func TestEinsumSumAxesMatchesEngine(t *testing.T) {
|
||||
shapes := map[int][]int{2: {6, 5}, 3: {4, 6, 5}, 4: {3, 4, 6, 2}}
|
||||
specs := []string{
|
||||
"ij->i", "ij->j",
|
||||
"ijk->ij", "ijk->ik", "ijk->ji", "ijk->i", "ijk->j", "ijk->k",
|
||||
"kji->k", "kji->ki", "kji->jk",
|
||||
"ijkl->il", "ijkl->lj", "ijkl->ji", "ijkl->i", "ijkl->l",
|
||||
}
|
||||
for _, spec := range specs {
|
||||
lhsStr, rhsStr, ok := strings.Cut(spec, "->")
|
||||
shape, ok2 := shapes[len(lhsStr)]
|
||||
if !ok || !ok2 {
|
||||
continue
|
||||
}
|
||||
for _, dt := range []Dtype{Int, Float32, Float, Complex} {
|
||||
a := einsumOperand(t, dt, shape...)
|
||||
want, werr := einsumGeneral([]string{lhsStr}, rhsStr, true, []*Array{a})
|
||||
got, gerr := Einsum(spec, a)
|
||||
if werr != nil || gerr != nil {
|
||||
t.Fatalf("%s/%s %v: dispatch error %v, engine error %v", spec, dt, shape, gerr, werr)
|
||||
}
|
||||
if !einsumBitsEqual(got, want) {
|
||||
t.Fatalf("%s/%s %v: dispatched result %s differs from the engine %s",
|
||||
spec, dt, shape, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// einsumDenseOf copies a view's logical elements into a contiguous
|
||||
// array of the same dtype.
|
||||
func einsumDenseOf(t *testing.T, a *Array) *Array {
|
||||
t.Helper()
|
||||
n := a.Len()
|
||||
var out *Array
|
||||
var err error
|
||||
switch a.dt {
|
||||
case Int:
|
||||
v := make([]int64, n)
|
||||
for i := range v {
|
||||
v[i] = a.ints[a.physIndex(i)]
|
||||
}
|
||||
out, err = FromInts(v, a.Shape()...)
|
||||
case Float32:
|
||||
v := make([]float32, n)
|
||||
for i := range v {
|
||||
v[i] = a.floats32[a.physIndex(i)]
|
||||
}
|
||||
out, err = FromFloat32s(v, a.Shape()...)
|
||||
case Complex:
|
||||
v := make([]complex128, n)
|
||||
for i := range v {
|
||||
v[i] = a.complexAt(i)
|
||||
}
|
||||
out, err = FromComplexes(v, a.Shape()...)
|
||||
default:
|
||||
v := make([]float64, n)
|
||||
for i := range v {
|
||||
v[i] = a.floatAt(i)
|
||||
}
|
||||
out, err = FromFloats(v, a.Shape()...)
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("dense copy of %s: %v", a, err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestEinsumStridedViewMatchesDense pins the engine's readers against a
|
||||
// view's own elements: a strided operand must contract exactly as a
|
||||
// dense array holding the same logical elements, for every dtype the
|
||||
// engine widens. The patterns are the ones the engine owns; the table
|
||||
// paths read payload windows and take contiguous operands only, as the
|
||||
// Array contract states.
|
||||
func TestEinsumStridedViewMatchesDense(t *testing.T) {
|
||||
for _, dt := range []Dtype{Int, Float32, Float, Complex} {
|
||||
view := einsumStridedOperand(t, dt, []int{3, 4}, []int{8, 2})
|
||||
dense := einsumDenseOf(t, view)
|
||||
other := einsumOperand(t, dt, 3, 4)
|
||||
vec := einsumOperand(t, dt, 4)
|
||||
for _, tc := range []struct {
|
||||
spec string
|
||||
which []bool // true takes the strided view, false the dense operand
|
||||
}{
|
||||
{"ij->i", []bool{true}},
|
||||
{"ij->j", []bool{true}},
|
||||
{"ij,ij->i", []bool{true, false}},
|
||||
{"ij,ij->ji", []bool{true, false}},
|
||||
{"ij,kj->ik", []bool{true, false}},
|
||||
{"...i,i->...", []bool{true, false}},
|
||||
} {
|
||||
withView := make([]*Array, len(tc.which))
|
||||
withDense := make([]*Array, len(tc.which))
|
||||
for i, isView := range tc.which {
|
||||
if isView {
|
||||
withView[i], withDense[i] = view, dense
|
||||
continue
|
||||
}
|
||||
if tc.spec == "...i,i->..." {
|
||||
withView[i], withDense[i] = vec, vec
|
||||
continue
|
||||
}
|
||||
withView[i], withDense[i] = other, other
|
||||
}
|
||||
got, gerr := Einsum(tc.spec, withView...)
|
||||
want, werr := Einsum(tc.spec, withDense...)
|
||||
if werr != nil || gerr != nil {
|
||||
t.Fatalf("%s/%s: dense error %v, view error %v", tc.spec, dt, werr, gerr)
|
||||
}
|
||||
if !einsumBitsEqual(got, want) {
|
||||
t.Fatalf("%s/%s: view result %s differs from the dense array %s", tc.spec, dt, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// benchEinsum measures one spec with the worker pool as it stands and
|
||||
// with it pinned to one worker: the in-process A/B of the slot walk's
|
||||
// split, where the pinned pool walks every slot in one chunk.
|
||||
func benchEinsum(b *testing.B, spec string, ops ...*Array) {
|
||||
b.Run("pool", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
if _, err := Einsum(spec, ops...); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
b.Run("one worker", func(b *testing.B) {
|
||||
restore := engine.SetNumWorkers(1)
|
||||
defer engine.SetNumWorkers(restore)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
if _, err := Einsum(spec, ops...); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// BenchmarkEinsumEllipsisBatch is a contraction carrying a batch axis:
|
||||
// the pattern the dispatch table cannot take, walked slot by slot.
|
||||
func BenchmarkEinsumEllipsisBatch(b *testing.B) {
|
||||
a := benchMat(b, 11, 8, 48, 48)
|
||||
c := benchMat(b, 12, 48, 48)
|
||||
benchEinsum(b, "...ij,jk->...ik", a, c)
|
||||
}
|
||||
|
||||
// BenchmarkEinsumEllipsisFloat32 is the same contraction on float32
|
||||
// payloads: the operands widen into the walk and every addend narrows
|
||||
// into the float32 slot.
|
||||
func BenchmarkEinsumEllipsisFloat32(b *testing.B) {
|
||||
a := einsumOperand(b, Float32, 8, 48, 48)
|
||||
c := einsumOperand(b, Float32, 48, 48)
|
||||
benchEinsum(b, "...ij,jk->...ik", a, c)
|
||||
}
|
||||
|
||||
// BenchmarkEinsumThreeOperandChain sums one shared label across three
|
||||
// operands, the shape attention and tensordot patterns reduce to.
|
||||
func BenchmarkEinsumThreeOperandChain(b *testing.B) {
|
||||
a := benchMat(b, 13, 16, 32)
|
||||
c := benchMat(b, 14, 32, 32)
|
||||
d := benchMat(b, 15, 32, 16)
|
||||
benchEinsum(b, "ik,kj,jl->il", a, c, d)
|
||||
}
|
||||
|
||||
// BenchmarkEinsumSmallBatchedProduct is the small batched product the
|
||||
// dispatch table's own batched kernel takes, kept as the fixed-cost
|
||||
// control beside the general engine's shapes.
|
||||
func BenchmarkEinsumSmallBatchedProduct(b *testing.B) {
|
||||
a := benchMat(b, 16, 4, 32, 32)
|
||||
c := benchMat(b, 17, 4, 32, 32)
|
||||
benchEinsum(b, "bij,bjk->bik", a, c)
|
||||
}
|
||||
|
||||
// BenchmarkEinsumReduceOnly measures a reduction-only spec through the
|
||||
// axis fold the dispatch now takes, against the same spec on the
|
||||
// general engine, which is where it went before.
|
||||
func BenchmarkEinsumReduceOnly(b *testing.B) {
|
||||
a := benchMat(b, 18, 512, 512)
|
||||
b.Run("axis fold", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
if _, err := Einsum("ij->i", a); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
b.Run("general engine", func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
if _, err := einsumGeneral([]string{"ij"}, "i", true, []*Array{a}); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
b.Run("general engine, one worker", func(b *testing.B) {
|
||||
restore := engine.SetNumWorkers(1)
|
||||
defer engine.SetNumWorkers(restore)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
if _, err := einsumGeneral([]string{"ij"}, "i", true, []*Array{a}); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// BenchmarkEinsumSlotKernelWalk is one slot's walk shaped like a
|
||||
// contraction's, two operands over a summed axis of 48 with the second
|
||||
// operand's cursor stepping a row: the peeled kernel against the
|
||||
// odometer oracle, in one process.
|
||||
func BenchmarkEinsumSlotKernelWalk(b *testing.B) {
|
||||
sizes := []int{48}
|
||||
base := []int{0, 0}
|
||||
delta := []int{1, 48}
|
||||
rv := [][]float64{einsumProbeFloats(48), einsumProbeFloats(48 * 48)}
|
||||
out := make([]float64, 1)
|
||||
off := make([]int, 2)
|
||||
coord := make([]int, 1)
|
||||
b.Run("peeled", func(b *testing.B) {
|
||||
for b.Loop() {
|
||||
einsumSlotSum(rv, out, 0, base, off, coord, delta, sizes, 48)
|
||||
}
|
||||
})
|
||||
b.Run("odometer", func(b *testing.B) {
|
||||
for b.Loop() {
|
||||
einsumSlotSumOdometer(rv, out, 0, base, off, coord, delta, sizes, 48)
|
||||
}
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user