feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,177 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user