Files

119 lines
2.7 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 spmd
import (
"strconv"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The collective benchmarks time one world of four ranks running the
// named collective back to back: every rank loops b.N times, so the
// per-operation number is one round of the collective across the
// world. In-process links put the channel and the arithmetic on the
// clock, not the network.
func BenchmarkAllReduceShards(b *testing.B) {
const gn = 1 << 20
whole := fixtureArray(gn)
if err := Launch(4, func(w *World) error {
span := mustPartition(b, gn, w.Size(), w.Rank())
local, err := dealRows(whole, span)
if err != nil {
return err
}
for i := 0; i < b.N; i++ {
if _, err := w.AllReduceShards(local, span, Sum); err != nil {
return err
}
}
return w.Barrier()
}); err != nil {
b.Fatal(err)
}
}
func BenchmarkAllReduce(b *testing.B) {
const per = 1 << 18
a, err := core.FromFloats(fixture(per), per)
if err != nil {
b.Fatal(err)
}
if err := Launch(4, func(w *World) error {
for i := 0; i < b.N; i++ {
if _, err := w.AllReduce(a, Sum); err != nil {
return err
}
}
return w.Barrier()
}); err != nil {
b.Fatal(err)
}
}
// BenchmarkReduce times the same-shaped-arrays fold on one million
// float64 elements across world sizes 2, 4 and 8 for Sum and Min: the
// fold runs as the chunks land, each arriving chunk joining the
// running accumulator at its rank's turn.
func BenchmarkReduce(b *testing.B) {
const n = 1 << 20
a, err := core.FromFloats(fixture(n), n)
if err != nil {
b.Fatal(err)
}
for _, size := range []int{2, 4, 8} {
for _, op := range []Op{Sum, Min} {
b.Run(op.String()+"/"+strconv.Itoa(size), func(b *testing.B) {
b.ReportAllocs()
if err := Launch(size, func(w *World) error {
for i := 0; i < b.N; i++ {
if _, err := w.Reduce(a, op, 0); err != nil {
return err
}
}
return w.Barrier()
}); err != nil {
b.Fatal(err)
}
})
}
}
}
func BenchmarkBroadcast(b *testing.B) {
a := fixtureArray(1 << 18)
if err := Launch(4, func(w *World) error {
for i := 0; i < b.N; i++ {
if _, err := w.Broadcast(a, 0); err != nil {
return err
}
}
return w.Barrier()
}); err != nil {
b.Fatal(err)
}
}
func BenchmarkScatterGather(b *testing.B) {
const gn = 1 << 20
a := fixtureArray(gn)
if err := Launch(4, func(w *World) error {
for i := 0; i < b.N; i++ {
local, span, err := w.Scatter(a, 0)
if err != nil {
return err
}
if _, err := w.Gather(local, 0); err != nil {
return err
}
_ = span
}
return w.Barrier()
}); err != nil {
b.Fatal(err)
}
}