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