Files
tensor/internal/core/bench_einsum2_test.go
T

640 lines
19 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"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)
}
})
}