// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import "testing" // Einsum benchmarks guard both the fast paths and the general engine. // The batched shapes are the training-loop workloads; the general // engine shapes are the ones the dispatch table declines. func benchVals(b *testing.B, seed, n int) []float64 { b.Helper() v := make([]float64, n) for i := range v { v[i] = float64(i%17)*float64(seed%5) + float64(i%7) - 6 } return v } func benchMat(b *testing.B, seed int, shape ...int) *Array { b.Helper() n := 1 for _, d := range shape { n *= d } a, err := FromFloats(benchVals(b, seed, n), shape...) if err != nil { b.Fatal(err) } return a } // BenchmarkEinsumMatMulPath exercises the dispatch table's MatMul fast // path as the control. func BenchmarkEinsumMatMulPath(b *testing.B) { a := benchMat(b, 1, 64, 64) c := benchMat(b, 2, 64, 64) b.ReportAllocs() for b.Loop() { if _, err := Einsum("ij,jk->ik", a, c); err != nil { b.Fatal(err) } } } // BenchmarkEinsumBatchedMatMul is the batched product the table has no // fast path for: it falls to the general engine. func BenchmarkEinsumBatchedMatMul(b *testing.B) { a := benchMat(b, 3, 16, 64, 64) c := benchMat(b, 4, 16, 64, 64) b.ReportAllocs() for b.Loop() { if _, err := Einsum("bij,bjk->bik", a, c); err != nil { b.Fatal(err) } } } // BenchmarkEinsumBatchedSmall keeps the general engine's fixed costs // visible at a size where they dominate. func BenchmarkEinsumBatchedSmall(b *testing.B) { a := benchMat(b, 5, 4, 32, 32) c := benchMat(b, 6, 4, 32, 32) b.ReportAllocs() for b.Loop() { if _, err := Einsum("bij,bjk->bik", a, c); err != nil { b.Fatal(err) } } } // BenchmarkEinsumEllipsisMatMul is the ellipsis spelling of the batched // product. func BenchmarkEinsumEllipsisMatMul(b *testing.B) { a := benchMat(b, 7, 8, 48, 48) c := benchMat(b, 8, 48, 48) b.ReportAllocs() for b.Loop() { if _, err := Einsum("...ij,jk->...ik", a, c); err != nil { b.Fatal(err) } } } // BenchmarkEinsumBilinear sums a shared label across three operands, // the pattern attention and tensordot shapes reduce to. func BenchmarkEinsumBilinear(b *testing.B) { a := benchMat(b, 9, 16, 32) c := benchMat(b, 10, 32, 32) d := benchMat(b, 11, 32, 16) b.ReportAllocs() for b.Loop() { if _, err := Einsum("ik,kj,jl->il", a, c, d); err != nil { b.Fatal(err) } } }