99 lines
2.4 KiB
Go
99 lines
2.4 KiB
Go
// 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)
|
|
}
|
|
}
|
|
}
|