Files

99 lines
2.4 KiB
Go
Raw Permalink 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 "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)
}
}
}