// Copyright (c) 2026 Petr BalvĂ­n (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) } }) }