// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package spmd import ( "net" "sync" "testing" "time" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // The TCP benchmarks run one loopback world of four ranks over real // framed connections, every rank inside this process: the network path // is the transport the routed frame pool lives on, and the scheduler // noise of separate processes stays out. Each benchmark runs twice, as // alternating sub-benchmarks with the pool on and off through the // package switch, so both variants of the A/B share one binary and one // machine state; the verdict comes from the medians across repeated // rounds. // benchmarkTCPBothWays runs one TCP benchmark as two sub-benchmarks, // the routed frame pool on, then off. func benchmarkTCPBothWays(b *testing.B, warm func(w *World) (func() error, error)) { for _, way := range []struct { name string pooled bool }{ {"pool", true}, {"base", false}, } { b.Run(way.name, func(b *testing.B) { was := framePoolEnabled framePoolEnabled = way.pooled defer func() { framePoolEnabled = was }() benchmarkTCP(b, warm) }) } } // benchmarkTCP assembles the loopback world, runs warm once on every // rank before the clock starts, then the round it answers b.N times on // every rank, so one operation is one round of the world. func benchmarkTCP(b *testing.B, warm func(w *World) (func() error, error)) { b.Helper() b.ReportAllocs() ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { b.Fatal(err) } defer ln.Close() const size = 4 var wg sync.WaitGroup errs := make([]error, size) wg.Go(func() { errs[0] = benchTCPRank(b, func() (*World, error) { return listen(ln, size, Options{Timeout: 2 * time.Minute}) }, warm) }) for r := 1; r < size; r++ { wg.Go(func() { errs[r] = benchTCPRank(b, func() (*World, error) { return Join(ln.Addr().String(), Options{Timeout: 2 * time.Minute}) }, warm) }) } wg.Wait() for r, err := range errs { if err != nil { b.Fatalf("rank %d: %v", r, err) } } } func benchTCPRank(b *testing.B, build func() (*World, error), warm func(w *World) (func() error, error)) error { w, err := build() if err != nil { return err } defer w.Close() round, err := warm(w) if err != nil { return err } if err := w.Barrier(); err != nil { return err } for range b.N { if err := round(); err != nil { return err } } return w.Barrier() } // BenchmarkTCPAllReduceShards is the measurement the pool item names: // four ranks folding one million float64s to one shared scalar. Its // frames stay small and every one of them is addressed to rank 0, so // none of them routes through the hub's outbox: it is the battery's // control, and the pool should leave it untouched. func BenchmarkTCPAllReduceShards(b *testing.B) { const gn = 1 << 20 whole := fixtureArray(gn) benchmarkTCPBothWays(b, func(w *World) (func() error, error) { span := mustPartition(b, gn, w.Size(), w.Rank()) local, err := dealRows(whole, span) if err != nil { return nil, err } return func() error { _, err := w.AllReduceShards(local, span, Sum) return err }, nil }) } // BenchmarkTCPMovement is the movement battery of the TCP tests: // Scatter, AllGather and a Broadcast whose root is a non-hub rank, so // its 560 kilobyte frames route through the hub's pump and outbox, the // path the pool serves. The canonical pieces of 70001 rows are the // ones that partition gives, empty pieces included, exactly as the // test the battery mirrors runs them. func BenchmarkTCPMovement(b *testing.B) { const gn = 70001 whole := fixtureArray(gn) benchmarkTCPBothWays(b, func(w *World) (func() error, error) { return func() error { local, _, err := w.Scatter(whole, 0) if err != nil { return err } got, err := w.AllGather(local) if err != nil { return err } _, err = w.Broadcast(got, 2) return err }, nil }) } // BenchmarkTCPReduce walks the array reduce, whose chunks travel rank // to owner through the hub: six routed frames of half a mebibyte a // round, the heaviest routed traffic in the battery. func BenchmarkTCPReduce(b *testing.B) { const per = 1 << 18 a, err := core.FromFloats(fixture(per), per) if err != nil { b.Fatal(err) } benchmarkTCPBothWays(b, func(w *World) (func() error, error) { return func() error { _, err := w.AllReduce(a, Sum) return err }, nil }) } // BenchmarkTCPExchangeHalos serves the stencil workloads: edge rows // travel to both neighbours through the hub. The grid sits on the // canonical partition's own boundaries, four blocks of 65536 rows, so // every rank holds a real piece, and the halo width puts the edges at // 128 kilobytes, a size the pool retains: four routed frames a round. func BenchmarkTCPExchangeHalos(b *testing.B) { const rows, width, halos = 262144, 8, 2048 whole := globalRows(rows, width) benchmarkTCPBothWays(b, func(w *World) (func() error, error) { span := mustPartition(b, rows, w.Size(), w.Rank()) local := dealRows2D(b, whole, span.Lo, span.Hi) return func() error { _, _, err := w.ExchangeHalos(local, halos) return err }, nil }) }