feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+248
View File
@@ -0,0 +1,248 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The arg reductions answer where an extremum lives in the global
// array: the same global index the single-array ArgMax and ArgMin
// answer, ties resolved by the earliest index, NaN elements skipped
// as missing. The shards compare values exactly, so the answer is the
// single-array answer's at any world size.
// AllReduceArgShards answers the global index of the extremum (Max or
// Min) of one global array whose canonical pieces the ranks hold,
// every rank receiving the same index. Ties keep the earliest global
// index, the single-array walk's own rule; NaN elements never win,
// and an array with no candidate at all is an error, as the
// single-array reduction is.
func (w *World) AllReduceArgShards(local *core.Array, span Span, op Op) (int, error) {
ans, err := w.reduceArgShards(local, span, op, 0)
if err != nil {
return 0, err
}
s, err := w.broadcastScalar(ans, 0)
if err != nil {
return 0, err
}
return int(s.Int()), nil
}
// ReduceArgShards is AllReduceArgShards with the answer on the root
// alone; every other rank receives nil.
func (w *World) ReduceArgShards(local *core.Array, span Span, op Op, root int) (*int, error) {
if err := checkRoot(root, w.size); err != nil {
return nil, w.fail(err)
}
ans, err := w.reduceArgShards(local, span, op, root)
if err != nil {
return nil, err
}
if w.rank != root {
return nil, nil
}
i := int(ans.Int())
return &i, nil
}
// argCandidate is one rank's extremum candidate: the extremum's value
// and its global index. A slab with no candidate (all NaN, or no rows)
// carries ok=false.
type argCandidate struct {
ok bool
value float64
iv int64 // the exact value for the integer dtypes
idx int
isInt bool
}
// reduceArgShards runs the arg reduction; the answer exists at the
// root.
func (w *World) reduceArgShards(local *core.Array, span Span, op Op, root int) (core.Scalar, error) {
if err := w.status(); err != nil {
return core.Scalar{}, err
}
if op != Min && op != Max {
return core.Scalar{}, w.fail(base.Errf("spmd: arg reductions take Min or Max, got %s", op))
}
if err := w.checkOneDimSpan(local, span); err != nil {
return core.Scalar{}, w.fail(err)
}
if local.Dtype() == core.Bool || local.Dtype() == core.Complex {
return core.Scalar{}, w.fail(base.Errf("spmd: %s of dtype %s has no ordering", op, local.Dtype()))
}
if span.Global == 0 {
return core.Scalar{}, w.fail(base.Errf("spmd: %s of an empty array has no answer", op))
}
cand := argCandidateFor(local, span, op)
if w.rank != root {
if err := w.sendTo(root, tagShardsValues, encodeBlockValues(blockValues{
first: cand.idx, kind: kindArg, i: []int64{cand.iv, int64(cand.idx)}, f: []float64{cand.value}, oks: []bool{cand.ok, cand.isInt},
})); err != nil {
return core.Scalar{}, err
}
return core.Scalar{}, nil
}
all := make([]argCandidate, w.size)
all[w.rank] = cand
for r := range w.size {
if r == root {
continue
}
data, err := w.recvFrom(r, tagShardsValues)
if err != nil {
return core.Scalar{}, err
}
bv, err := decodeBlockValues(data)
if err != nil {
return core.Scalar{}, w.fail(err)
}
if bv.kind != kindArg || len(bv.i) != 2 || len(bv.f) != 1 || len(bv.oks) != 2 {
return core.Scalar{}, w.fail(base.Errf("spmd: arg shards disagree on the payload shape"))
}
all[r] = argCandidate{ok: bv.oks[0], value: bv.f[0], iv: bv.i[0], idx: int(bv.i[1]), isInt: bv.oks[1]}
}
greater := op == Max
best := -1
for r, c := range all {
if !c.ok {
continue
}
if best < 0 {
best = r
continue
}
b := all[best]
better := false
if c.isInt == b.isInt {
if c.isInt {
better = (greater && c.iv > b.iv) || (!greater && c.iv < b.iv)
} else {
better = (greater && c.value > b.value) || (!greater && c.value < b.value)
}
} else {
return core.Scalar{}, w.fail(base.Errf("spmd: arg shards disagree on the payload kind"))
}
if better {
best = r
}
// A tie keeps the earlier rank's candidate, whose global index
// is the earlier one: the pieces are contiguous and ascend.
}
if best < 0 {
return core.Scalar{}, w.fail(base.Errf("spmd: %s shards have no candidate", op))
}
winner := all[best]
return core.IntScalar(int64(winner.idx)), nil
}
// argCandidateFor finds the slab's own first extremum, the serial
// walk's rule: the first strictly better element, NaNs skipped. The
// integer dtypes compare exactly in int64, never through float64,
// the reason the core keeps its own int64 walk.
func argCandidateFor(local *core.Array, span Span, op Op) argCandidate {
greater := op == Max
best := -1
var bv float64
var biv int64
intExact := false
walkInt := func(v int64, i int) {
if best < 0 || (greater && v > biv) || (!greater && v < biv) {
best, biv, intExact = i, v, true
}
}
walkFloat := func(v float64, i int) {
if math.IsNaN(v) {
return
}
if best < 0 || (greater && v > bv) || (!greater && v < bv) {
best, bv, intExact = i, v, false
}
}
for i := range local.Len() {
switch local.Dtype() {
case core.Float:
walkFloat(local.FloatAt(i), i)
case core.Float32:
walkFloat(float64(local.RawFloat32s()[i]), i)
case core.Float16:
walkFloat(core.HalfToFloat64(local.RawHalves()[i]), i)
case core.Int:
walkInt(local.RawInts()[i], i)
case core.Int8:
walkInt(int64(local.RawInt8s()[i]), i)
case core.Uint8:
walkInt(int64(local.RawUint8s()[i]), i)
case core.Int16:
walkInt(int64(local.RawInt16s()[i]), i)
case core.Uint16:
walkInt(int64(local.RawUint16s()[i]), i)
case core.Int32:
walkInt(int64(local.RawInt32s()[i]), i)
case core.Uint32:
walkInt(int64(local.RawUint32s()[i]), i)
}
}
if best < 0 {
return argCandidate{}
}
if intExact {
return argCandidate{ok: true, iv: biv, idx: span.Lo + best, isInt: true}
}
return argCandidate{ok: true, value: bv, idx: span.Lo + best}
}
// AllReduceArgSortShards answers the global permutation that sorts the
// whole array ascending: an Int array of global indices, the
// single-array ArgSort's own answer with its own tie and NaN
// placement, on every rank. The shards' values gather in global order
// and the core sorts them, so the permutation is the single-array
// one by construction.
func (w *World) AllReduceArgSortShards(local *core.Array, span Span) (*core.Array, error) {
ans, err := w.reduceArgSortShards(local, span, 0)
if err != nil {
return nil, err
}
return w.Broadcast(ans, 0)
}
// ReduceArgSortShards is AllReduceArgSortShards with the answer on the
// root alone; every other rank receives nil.
func (w *World) ReduceArgSortShards(local *core.Array, span Span, root int) (*core.Array, error) {
if err := checkRoot(root, w.size); err != nil {
return nil, w.fail(err)
}
return w.reduceArgSortShards(local, span, root)
}
func (w *World) reduceArgSortShards(local *core.Array, span Span, root int) (*core.Array, error) {
if err := w.status(); err != nil {
return nil, err
}
if err := w.checkOneDimSpan(local, span); err != nil {
return nil, w.fail(err)
}
switch local.Dtype() {
case core.Int, core.Float32, core.Float16, core.Float:
default:
return nil, w.fail(base.Errf("spmd: ArgSort of dtype %s is not supported; convert with Astype", local.Dtype()))
}
gathered, err := w.gather(local, root)
if err != nil {
return nil, err
}
if w.rank != root {
return nil, nil
}
return core.ArgSort(gathered)
}
// nanValue is the quiet NaN the no-candidate tests build their
// fixtures from.
var nanValue = math.NaN()
+158
View File
@@ -0,0 +1,158 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The arg reductions answer the single-array walk's global index, ties
// by the earliest index, NaNs skipped, and the sharded sort answers
// the single-array permutation outright.
// TestShardedArgMatchesSingleArray pins ArgMax and ArgMin against the
// core's own walk across dtypes, NaNs included, with the tie at the
// earliest index.
func TestShardedArgMatchesSingleArray(t *testing.T) {
for _, size := range []int{1, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537} {
for dt, whole := range fixtureDtypes(t, gn) {
switch dt {
case core.Float, core.Float32, core.Float16, core.Int, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32:
default:
continue
}
for _, op := range []Op{Max, Min} {
var want int
var err error
if op == Max {
want, err = core.ArgMax(whole)
} else {
want, err = core.ArgMin(whole)
}
if err != nil {
t.Fatal(err)
}
err = Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceArgShards(local, span, op)
if err != nil {
return err
}
if got != want {
t.Fatalf("size %d gn %d %s %s: sharded index %d against single-array %d",
w.Size(), gn, dt, op, got, want)
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d %s %s: %v", size, gn, dt, op, err)
}
}
}
}
}
}
// TestShardedArgNoCandidate: an all-NaN array is an error on every
// rank, the single-array walk's own refusal.
func TestShardedArgNoCandidate(t *testing.T) {
nans := make([]float64, 100)
for i := range nans {
nans[i] = math2NaN()
}
whole, err := core.FromFloats(nans, 100)
if err != nil {
t.Fatal(err)
}
err = Launch(3, func(w *World) error {
span := mustPartition(t, 100, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
_, err := w.AllReduceArgShards(local, span, Max)
return err
})
if err == nil {
t.Fatal("an all-NaN array answered an index")
}
}
func math2NaN() float64 { return nanValue }
// TestShardedArgSortMatchesSingleArray: the sharded permutation is
// the single-array ArgSort's own, values, ties, NaN placement and all.
func TestShardedArgSortMatchesSingleArray(t *testing.T) {
for _, size := range []int{1, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537} {
for _, dt := range []core.Dtype{core.Float, core.Float32, core.Float16, core.Int} {
whole := fixtureDtypes(t, gn)[dt]
want, err := core.ArgSort(whole)
if err != nil {
t.Fatal(err)
}
err = Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceArgSortShards(local, span)
if err != nil {
return err
}
if got.Len() != want.Len() {
t.Fatalf("size %d gn %d %s: permutation of %d against %d",
w.Size(), gn, dt, got.Len(), want.Len())
}
for i := range want.Len() {
if got.RawInts()[i] != want.RawInts()[i] {
t.Fatalf("size %d gn %d %s: permutation differs at %d: %d against %d",
w.Size(), gn, dt, i, got.RawInts()[i], want.RawInts()[i])
}
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d %s: %v", size, gn, dt, err)
}
}
}
}
}
// TestShardedArgOverTCP runs the arg and the sort over real
// connections.
func TestShardedArgOverTCP(t *testing.T) {
const gn = 65537
whole := fixtureDtypes(t, gn)[core.Float]
wantMax, err := core.ArgMax(whole)
if err != nil {
t.Fatal(err)
}
wantSort, err := core.ArgSort(whole)
if err != nil {
t.Fatal(err)
}
runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
gotMax, err := w.AllReduceArgShards(local, span, Max)
if err != nil {
return err
}
if gotMax != wantMax {
t.Fatalf("rank %d: arg %d against %d", w.Rank(), gotMax, wantMax)
}
gotSort, err := w.AllReduceArgSortShards(local, span)
if err != nil {
return err
}
for i := range wantSort.Len() {
if gotSort.RawInts()[i] != wantSort.RawInts()[i] {
t.Fatalf("rank %d: the permutation differs at %d over TCP", w.Rank(), i)
}
}
return nil
})
}
+177
View File
@@ -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
})
}
+118
View File
@@ -0,0 +1,118 @@
// 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)
}
}
+285
View File
@@ -0,0 +1,285 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The movement collectives carry bits and carry nothing else: no
// arithmetic happens in transit, so their answers are deterministic by
// construction. The wire is the only form that travels, and every
// collective's join order is a function of the data, never of the
// order frames happen to arrive in.
// Broadcast delivers root's array to every rank of the world. The root
// gets the array it passed; every other rank gets an equal copy, bits
// included. A world of one rank hands the array back.
func (w *World) Broadcast(a *core.Array, root int) (*core.Array, error) {
if err := w.status(); err != nil {
return nil, err
}
if err := checkRoot(root, w.size); err != nil {
return nil, w.fail(err)
}
if w.size == 1 {
return a, nil
}
if w.rank == root {
wire, err := encodeArray(nil, a)
if err != nil {
return nil, w.fail(err)
}
for r := range w.size {
if r == root {
continue
}
if err := w.sendTo(r, tagBroadcast, wire); err != nil {
return nil, err
}
}
return a, nil
}
data, err := w.recvFrom(root, tagBroadcast)
if err != nil {
return nil, err
}
got, err := decodeWire(data)
if err != nil {
return nil, w.fail(err)
}
return got, nil
}
// Scatter deals root's array out along its first dimension: rank r
// receives exactly the canonical partition's piece of the axis, so the
// world's data lands on the boundaries the shard reductions compose
// on. The global shape travels first, so every rank can name its own
// piece before any payload moves; the span that comes back names it.
func (w *World) Scatter(global *core.Array, root int) (*core.Array, Span, error) {
if err := w.status(); err != nil {
return nil, Span{}, err
}
if err := checkRoot(root, w.size); err != nil {
return nil, Span{}, w.fail(err)
}
if global.NDim() == 0 {
return nil, Span{}, w.fail(base.Errf("spmd: Scatter needs a dimension to deal along"))
}
var shape []int
if w.rank == root {
shape = global.Shape()
head, err := encodeHead(nil, global.Dtype(), shape)
if err != nil {
return nil, Span{}, w.fail(err)
}
for r := range w.size {
if r == root {
continue
}
if err := w.sendTo(r, tagScatterHead, head); err != nil {
return nil, Span{}, err
}
}
} else {
data, err := w.recvFrom(root, tagScatterHead)
if err != nil {
return nil, Span{}, err
}
_, shape, err = decodeHead(data)
if err != nil {
return nil, Span{}, w.fail(err)
}
}
if len(shape) == 0 {
return nil, Span{}, w.fail(base.Errf("spmd: the dealt head names no dimension to deal along"))
}
gn := shape[0]
rest, ok := elementCount(shape[1:])
if !ok {
return nil, Span{}, w.fail(base.Errf("spmd: Scatter's shape %v overflows the element count", shape))
}
full, ok := elementCount(shape)
if !ok {
return nil, Span{}, w.fail(base.Errf("spmd: Scatter's shape %v overflows the element count", shape))
}
if w.rank == root && full != global.Len() {
return nil, Span{}, w.fail(base.Errf("spmd: Scatter's shape %s names %d elements against the array's %d",
base.ShapeText(shape), full, global.Len()))
}
span, err := Partition(gn, w.size, w.rank)
if err != nil {
return nil, Span{}, w.fail(err)
}
slabShape := make([]int, 0, len(shape))
slabShape = append(slabShape, span.Len())
slabShape = append(slabShape, shape[1:]...)
local, err := func() (*core.Array, error) {
if w.rank == root {
wire, err := encodePart(nil, global, slabShape, span.Lo*rest, span.Len()*rest)
if err != nil {
return nil, err
}
for r := range w.size {
if r == root {
continue
}
s, err := Partition(gn, w.size, r)
if err != nil {
return nil, err
}
wire, err := encodePart(nil, global,
append([]int{s.Len()}, shape[1:]...), s.Lo*rest, s.Len()*rest)
if err != nil {
return nil, err
}
if err := w.sendTo(r, tagScatter, wire); err != nil {
return nil, err
}
}
return decodeWire(wire)
}
wire, err := w.recvFrom(root, tagScatter)
if err != nil {
return nil, err
}
return decodeWire(wire)
}()
if err != nil {
return nil, Span{}, w.fail(err)
}
return local, span, nil
}
// Gather raises the dealt pieces back into the global array on root,
// joining them in rank order, which is the canonical partition's order.
// Ranks other than root answer with nil: they hold their piece, the
// root holds the whole.
func (w *World) Gather(local *core.Array, root int) (*core.Array, error) {
global, err := w.gather(local, root)
if err != nil {
return nil, err
}
if w.rank != root {
return nil, nil
}
return global, nil
}
// gather is Gather's body; AllGather shares it with a fixed root.
func (w *World) gather(local *core.Array, root int) (*core.Array, error) {
if err := w.status(); err != nil {
return nil, err
}
if err := checkRoot(root, w.size); err != nil {
return nil, w.fail(err)
}
if local.NDim() == 0 {
return nil, w.fail(base.Errf("spmd: Gather needs a dimension to raise along"))
}
wire, err := encodeArray(nil, local)
if err != nil {
return nil, w.fail(err)
}
// A world of one rank is already the whole.
if w.size == 1 {
return local, nil
}
// The pieces travel to the root; the root takes them in rank
// order, which is the partition's order, whatever order they
// arrive in. Its own piece sits at its own rank's place in the
// join.
if w.rank != root {
if err := w.sendTo(root, tagGather, wire); err != nil {
return nil, err
}
return nil, nil
}
pieces := make([][]byte, 0, w.size)
for r := range w.size {
if r == root {
pieces = append(pieces, wire)
continue
}
data, err := w.recvFrom(r, tagGather)
if err != nil {
return nil, err
}
pieces = append(pieces, data)
}
return w.joinPieces(pieces)
}
// joinPieces builds the global array from the pieces' wire forms: the
// heads must agree in dtype and in every dimension but the first, and
// the payloads concatenate in the order given.
func (w *World) joinPieces(pieces [][]byte) (*core.Array, error) {
var dt core.Dtype
var trailing []int
total := 0
var payload []byte
for i, piece := range pieces {
pdt, shape, err := decodeHead(piece)
if err != nil {
return nil, w.fail(err)
}
if len(shape) == 0 {
return nil, w.fail(base.Errf("spmd: piece %d has no dimension to raise along", i))
}
if i == 0 {
dt = pdt
trailing = shape[1:]
} else {
if pdt != dt {
return nil, w.fail(base.Errf("spmd: piece %d is %s against %s", i, pdt, dt))
}
if len(shape) != len(trailing)+1 {
return nil, w.fail(base.Errf("spmd: piece %d has %d dimensions against %d", i, len(shape), len(trailing)+1))
}
for d := range trailing {
if shape[d+1] != trailing[d] {
return nil, w.fail(base.Errf("spmd: piece %d disagrees in dimension %d: %d against %d",
i, d+1, shape[d+1], trailing[d]))
}
}
}
if shape[0] > math.MaxInt-total {
return nil, w.fail(base.Errf("spmd: joining the pieces overflows the first dimension at piece %d", i))
}
total += shape[0]
payload = append(payload, piece[2+8*len(shape):]...)
}
globalShape := make([]int, 0, len(trailing)+1)
globalShape = append(globalShape, total)
globalShape = append(globalShape, trailing...)
head, err := encodeHead(nil, dt, globalShape)
if err != nil {
return nil, w.fail(err)
}
got, err := decodeWire(append(head, payload...))
if err != nil {
return nil, w.fail(err)
}
return got, nil
}
// AllGather raises every dealt piece into the whole on every rank: one
// world, one array, identical bits everywhere.
func (w *World) AllGather(local *core.Array) (*core.Array, error) {
whole, err := w.gather(local, 0)
if err != nil {
return nil, err
}
return w.Broadcast(whole, 0)
}
func checkRoot(root, size int) error {
if root < 0 || root >= size {
return base.Errf("spmd: root %d is outside the world of %d ranks", root, size)
}
return nil
}
+211
View File
@@ -0,0 +1,211 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"math"
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// fixture builds a deterministic float64 slice whose values carry
// magnitude spread and a fixed pattern, so a shuffled or truncated
// movement shows up in the bits.
func fixture(n int) []float64 {
v := make([]float64, n)
for i := range v {
v[i] = float64((i*7919)%211-105) / 7.0
if i%97 == 0 {
v[i] = math.Inf(1)
}
if i%89 == 0 {
v[i] = math.Copysign(0, -1)
}
}
return v
}
func fixtureArray(n int) *core.Array {
a, err := core.FromFloats(fixture(n), n)
if err != nil {
panic(err)
}
return a
}
// TestBroadcastCarriesTheBits: whatever rank originates the broadcast,
// every rank ends holding the root's exact bits, over the in-process
// world.
func TestBroadcastCarriesTheBits(t *testing.T) {
const n = 1000
for _, size := range []int{1, 2, 3, 5, 8} {
for _, root := range []int{0, size - 1} {
err := Launch(size, func(w *World) error {
want := fixtureArray(n)
got, err := w.Broadcast(want, root)
if err != nil {
return err
}
if w.Rank() == root {
if got != want {
t.Fatalf("rank %d: the root did not keep its own array", w.Rank())
}
return nil
}
if !sameBits(want, got) {
t.Fatalf("rank %d: broadcast bits differ from rank %d's", w.Rank(), root)
}
return nil
})
if err != nil {
t.Fatalf("size %d root %d: %v", size, root, err)
}
}
}
}
// TestBroadcastCarriesEveryDtype walks the narrow, half and boolean
// element types through the movement path once: the wire is the only
// form that travels, and every dtype has to survive it.
func TestBroadcastCarriesEveryDtype(t *testing.T) {
cases := []*core.Array{
mk(core.FromBools([]bool{true, false, true}, 3)),
mk(core.HalvesFromArray([]uint16{0x0001, 0x7bff, 0xfc00}, 3)),
mk(core.FromInt8s([]int8{-128, 127, 0}, 3)),
mk(core.FromUint32s([]uint32{4294967295, 0, 7}, 3)),
mk(core.FromComplexes([]complex128{1 + 2i, -0i}, 2)),
}
err := Launch(3, func(w *World) error {
for i, want := range cases {
got, err := w.Broadcast(want, i%w.Size())
if err != nil {
return err
}
if !sameBits(want, got) {
t.Fatalf("rank %d case %d: bits differ", w.Rank(), i)
}
}
return nil
})
if err != nil {
t.Fatal(err)
}
}
// TestScatterGatherRoundTrip deals a global array out and raises it
// back: every rank's piece is the canonical partition's own cut, and
// the gathered whole is the original's exact bits.
func TestScatterGatherRoundTrip(t *testing.T) {
for _, size := range []int{1, 2, 3, 5, 8} {
for _, gn := range []int{0, 1, 100, 65537, 200001} {
for _, root := range []int{0, size - 1} {
err := Launch(size, func(w *World) error {
src, err := core.FromFloats(fixture(gn*3), gn, 3)
if err != nil {
return err
}
want := fixture(gn * 3)
local, span, err := w.Scatter(src, root)
if err != nil {
return err
}
if span != mustPartition(t, gn, w.Size(), w.Rank()) {
t.Fatalf("rank %d: span %v against the partition's %v",
w.Rank(), span, mustPartition(t, gn, w.Size(), w.Rank()))
}
if local.Len() != span.Len()*3 {
t.Fatalf("rank %d: slab of %d elements for a span of %d",
w.Rank(), local.Len(), span.Len())
}
// Every element the slab carries is the fixture's own
// value at its global index.
for i := 0; i < local.Len(); i++ {
if got := local.FloatAt(i); got != want[span.Lo*3+i] {
t.Fatalf("rank %d element %d: %v", w.Rank(), i, got)
}
}
back, err := w.Gather(local, root)
if err != nil {
return err
}
if w.Rank() == root {
if !sameBits(src, back) {
t.Fatalf("rank %d: the gathered whole differs from the dealt array", w.Rank())
}
} else if back != nil {
t.Fatalf("rank %d: gather returned a whole to a non-root", w.Rank())
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d root %d: %v", size, gn, root, err)
}
}
}
}
}
// TestAllGatherRebuildsEverywhere: one deal, one raise, and every rank
// holds the whole.
func TestAllGatherRebuildsEverywhere(t *testing.T) {
const gn = 1000
err := Launch(4, func(w *World) error {
local, span, err := w.Scatter(fixtureArray(gn), 0)
if err != nil {
return err
}
whole, err := w.AllGather(local)
if err != nil {
return err
}
want := fixtureArray(gn)
if !sameBits(want, whole) {
t.Fatalf("rank %d: the rebuilt whole differs", w.Rank())
}
if whole.Len() != gn || span.Global != gn {
t.Fatalf("rank %d: whole %d against global %d", w.Rank(), whole.Len(), span.Global)
}
return nil
})
if err != nil {
t.Fatal(err)
}
}
// TestMovementOverTCP runs the same movement battery over real
// connections, because the contract says the two transports are one
// machine.
func TestMovementOverTCP(t *testing.T) {
for _, gn := range []int{0, 100, 70001} {
t.Run("", func(t *testing.T) {
runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error {
src := fixtureArray(gn)
local, span, err := w.Scatter(src, 0)
if err != nil {
return err
}
if span != mustPartition(t, gn, w.Size(), w.Rank()) {
t.Fatalf("rank %d: span %v", w.Rank(), span)
}
whole, err := w.AllGather(local)
if err != nil {
return err
}
if !sameBits(src, whole) {
t.Fatalf("rank %d: the whole differs over TCP", w.Rank())
}
round, err := w.Broadcast(whole, 2)
if err != nil {
return err
}
if !sameBits(whole, round) {
t.Fatalf("rank %d: broadcast over TCP differs", w.Rank())
}
return nil
})
})
}
}
+42
View File
@@ -0,0 +1,42 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Package spmd runs one program on many ranks: explicit SPMD worlds
// with the collectives of scientific distributed computing, over TCP
// between machines or in process within one.
//
// # The determinism contract
//
// Every collective answers bit-identical results whatever the world
// size, the machine, the worker count or the order messages arrive in.
// The rule that buys this is one: the order of a reduction is a
// function of the data, never of the topology. Family B reductions
// (ReduceShards, AllReduceShards) cut the global length at the fold's
// own block boundaries, fold each block where its elements live, and
// combine the block partials with the same balanced tree the
// single-array fold uses, so a sharded reduction and the single-array
// Sum are one computation and one answer, for one rank or for fifty.
// Family A reductions (Reduce, AllReduce) fold same-shaped arrays in
// rank index order, which the program itself fixes. No collective
// every combines partials in the order they happen to arrive.
//
// # Worlds
//
// Launch builds a world of size ranks in one process, one goroutine
// per rank, over channels: the development surface, the test surface
// and the single-machine surface. Listen and Join build the same world
// over TCP, rank 0 listening and every other rank dialling; ranks are
// assigned in dial order. The sharded reductions' results never
// depend on the assignment; the same-shaped arrays' fold order is the
// rank order the program fixes, so a program that wants bit-reproducible
// family A runs keeps its rank assignment stable. Collectives are bulk synchronous: every rank calls the
// same collective in the same order, as an MPI program does.
//
// Any failure, a deadline included, fails the whole world: the rank
// that saw it and every rank that then touches the world or waits on
// it receive an error, and no collective ever returns a partial
// numeric result. A failed world stays failed.
//
// The package is pure Go on the standard library: no cgo, no
// dependency, no build tag.
package spmd
+110
View File
@@ -0,0 +1,110 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import "sync"
// The routed frame pool recycles the payload buffers of the frames the
// hub routes on, the frames whose destination is another rank. Its
// safety rests on one ownership chain and nothing else: the pump takes
// a buffer when it reads a routed frame, the frame hands the buffer to
// the destination link's outbox channel, and the drain, the chain's
// single consumer, writes the frame and returns the buffer at once,
// on a landed write and on a benign skip at a departing peer alike.
// No frame travels a pending stash with the pool's mark on, no
// collective ever sees a pooled buffer, and no count, atom or second
// return point exists: a channel handoff is the ownership transfer.
// A buffer the chain drops on the way, a mislabelled frame or a world
// that ends mid-route, is left to the collector, which makes the loss
// a pool miss and never a double return. This single owner is why the
// first pool attempt's failure mode, a payload retired while a decode
// still read it, cannot arise here.
//
// Two bounds keep the pool from pinning memory. A buffer above the
// retention ceiling is never kept, so one enormous frame cannot pin
// the heap across rounds, and the pool holds at most a fixed number
// of buffers, so a flood of mid-sized ones cannot grow without end.
// Both refusals are pool misses: the buffer goes to the collector,
// which is exactly the behaviour the package had before the pool.
const (
// framePoolCeiling is the largest buffer the pool retains, one
// mebibyte. A payload beyond it is a pool miss by rule.
framePoolCeiling = int64(1 << 20)
// framePoolBound is the largest number of buffers the pool holds
// at once; a return that would exceed it is dropped.
framePoolBound = 32
)
// framePoolEnabled switches the pool at package level. It exists for
// the A/B measurement, which flips it between alternating
// sub-benchmarks of one binary; the library itself always runs pooled.
var framePoolEnabled = true
// framePool is the free list behind the pool: payload buffers waiting
// for the next routed frame whose length fits.
type framePool struct {
mu sync.Mutex
free [][]byte
}
// routedFrames is the process's one pool. The worlds of one process
// share it, which is what makes its bounds process-wide facts.
var routedFrames framePool
// takeFrameBuffer answers the buffer a routed frame's payload reads
// into, or nil when the frame allocates as usual, outside the pool's
// chain: when the pool is switched off, when the frame stays at the
// hub, or when the length the header announced is beyond the ceiling.
// A qualified frame rides the chain whatever the free list holds: the
// buffer comes from the pool when one fits, and a fresh one joins the
// chain in its place, so the drain has a buffer to return and the
// pool warms with the traffic it serves. The link read calls this
// with the header already parsed, before the payload moves.
func takeFrameBuffer(m message, length int64) []byte {
if !framePoolEnabled || m.dest == 0 || length <= 0 || length > framePoolCeiling {
return nil
}
if buf := routedFrames.take(length); buf != nil {
return buf
}
return make([]byte, length)
}
// take answers a buffer of exactly length bytes with at least that
// capacity, the smallest retained buffer that fits, or nil when the
// pool holds none.
func (p *framePool) take(length int64) []byte {
if length <= 0 || length > framePoolCeiling {
return nil
}
p.mu.Lock()
defer p.mu.Unlock()
best := -1
for i, buf := range p.free {
if int64(cap(buf)) >= length && (best < 0 || cap(buf) < cap(p.free[best])) {
best = i
}
}
if best < 0 {
return nil
}
buf := p.free[best]
p.free = append(p.free[:best], p.free[best+1:]...)
return buf[:length]
}
// retire returns one buffer to the pool under the two bounds: a
// buffer whose capacity exceeds the ceiling, and a return that would
// push the pool past its count bound, are dropped to the collector.
func (p *framePool) retire(data []byte) {
if !framePoolEnabled || cap(data) == 0 || int64(cap(data)) > framePoolCeiling {
return
}
p.mu.Lock()
if len(p.free) < framePoolBound {
p.free = append(p.free, data[:cap(data)])
}
p.mu.Unlock()
}
+212
View File
@@ -0,0 +1,212 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"encoding/binary"
"runtime"
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/base"
)
// The pool tests cover the two bounds, the exact-length answer, the
// switch, the bound's flatness in the process's heap, and the one
// ownership chain itself: routed frames carrying distinct bits through
// a real hub while the pool recycles every buffer between pump and
// drain.
// resetPool empties the shared pool so no test inherits another's
// buffers, pins the switch on, and restores both at the end.
func resetPool(t *testing.T) {
t.Helper()
was := framePoolEnabled
t.Cleanup(func() {
framePoolEnabled = was
emptyPool()
})
emptyPool()
framePoolEnabled = true
}
func emptyPool() {
routedFrames.mu.Lock()
routedFrames.free = nil
routedFrames.mu.Unlock()
}
// poolDepth answers how many buffers the pool holds.
func poolDepth() int {
routedFrames.mu.Lock()
defer routedFrames.mu.Unlock()
return len(routedFrames.free)
}
func TestFramePoolTakeAnswersExactLengths(t *testing.T) {
resetPool(t)
if buf := routedFrames.take(1000); buf != nil {
t.Fatal("an empty pool answered a buffer")
}
routedFrames.retire(make([]byte, 1000))
buf := routedFrames.take(600)
if buf == nil {
t.Fatal("a retained buffer was not answered")
}
if len(buf) != 600 || cap(buf) < 600 {
t.Fatalf("take answered len %d cap %d for a 600 byte payload", len(buf), cap(buf))
}
if buf := routedFrames.take(600); buf != nil {
t.Fatal("a taken buffer was answered twice")
}
}
func TestFramePoolTakesTheSmallestThatFits(t *testing.T) {
resetPool(t)
routedFrames.retire(make([]byte, 5000))
routedFrames.retire(make([]byte, 900))
buf := routedFrames.take(800)
if buf == nil {
t.Fatal("a retained buffer was not answered")
}
if cap(buf) != 900 {
t.Fatalf("take answered cap %d when a 900 byte buffer was retained", cap(buf))
}
buf = routedFrames.take(800)
if buf == nil || cap(buf) != 5000 {
t.Fatalf("the second take answered cap %d, want the 5000 byte buffer", cap(buf))
}
}
func TestFramePoolRefusesTheOversized(t *testing.T) {
resetPool(t)
if buf := routedFrames.take(framePoolCeiling + 1); buf != nil {
t.Fatal("a length beyond the ceiling was answered")
}
routedFrames.retire(make([]byte, framePoolCeiling+1))
if got := poolDepth(); got != 0 {
t.Fatalf("a buffer beyond the ceiling was retained, pool holds %d", got)
}
routedFrames.retire(make([]byte, framePoolCeiling))
if got := poolDepth(); got != 1 {
t.Fatalf("a buffer at the ceiling was refused, pool holds %d", got)
}
if buf := routedFrames.take(framePoolCeiling); buf == nil {
t.Fatal("a length at the ceiling was not answered")
}
}
func TestFramePoolHoldsTheBound(t *testing.T) {
resetPool(t)
for range framePoolBound + 5 {
routedFrames.retire(make([]byte, 1000))
}
if got := poolDepth(); got != framePoolBound {
t.Fatalf("the pool holds %d buffers beyond its bound of %d", got, framePoolBound)
}
if buf := routedFrames.take(1000); buf == nil {
t.Fatal("a bound-full pool answered nothing")
}
}
func TestFramePoolAnswersNilWhenDisabled(t *testing.T) {
resetPool(t)
framePoolEnabled = false
routedFrames.retire(make([]byte, 1000))
if got := poolDepth(); got != 0 {
t.Fatalf("a disabled pool retained %d buffers", got)
}
if buf := routedFrames.take(1000); buf != nil {
t.Fatal("a disabled pool answered a buffer")
}
}
// TestFramePoolStaysFlatOnRepeatedRounds is the bound's proof in the
// heap: the same work run again and again in one process cannot raise
// the heap in use once the pool has warmed, and a rising trend is a
// defect.
func TestFramePoolStaysFlatOnRepeatedRounds(t *testing.T) {
resetPool(t)
sizes := []int64{64 << 10, 256 << 10, framePoolCeiling}
round := func() {
for _, s := range sizes {
buf := routedFrames.take(s)
if buf == nil {
buf = make([]byte, s)
}
for i := range buf {
buf[i] = byte(i)
}
routedFrames.retire(buf)
}
}
for range 50 {
round()
}
runtime.GC()
var before runtime.MemStats
runtime.ReadMemStats(&before)
for range 200 {
round()
}
runtime.GC()
var after runtime.MemStats
runtime.ReadMemStats(&after)
if after.HeapInuse > before.HeapInuse+4<<20 {
t.Fatalf("the heap in use rose from %d to %d bytes across 200 rounds", before.HeapInuse, after.HeapInuse)
}
}
// routedProbeRounds is how many distinct payloads the ownership test
// pushes through the hub, every one recycled through the pool.
const routedProbeRounds = 300
// TestFramePoolCarriesTheBitsOverTCP is the ownership chain exercised:
// a non-hub rank's frames route through the hub's pump, outbox and
// drain, the drain returns every buffer the moment its write lands,
// and the pump reads the next frame into what comes back, so a premature
// return or a shared buffer would scramble the bits the receiver
// checks, most of all under the race detector.
func TestFramePoolCarriesTheBitsOverTCP(t *testing.T) {
resetPool(t)
runTCPWorld(t, 3, Options{Timeout: 30 * time.Second}, func(w *World) error {
if w.Rank() == 2 {
for i := range routedProbeRounds {
buf := make([]byte, 1<<18)
binary.LittleEndian.PutUint64(buf, uint64(i))
for j := 8; j < len(buf); j += 7 {
buf[j] = byte(i + j)
}
if err := w.sendTo(1, tagHalo, buf); err != nil {
return err
}
}
return nil
}
if w.Rank() != 1 {
return nil
}
for i := range routedProbeRounds {
data, err := w.recvFrom(2, tagHalo)
if err != nil {
return err
}
if len(data) != 1<<18 {
return base.Errf("spmd: probe %d arrived %d bytes long", i, len(data))
}
if got := binary.LittleEndian.Uint64(data); got != uint64(i) {
return base.Errf("spmd: probe %d arrived with the serial of %d", i, got)
}
for j := 8; j < len(data); j += 7 {
if data[j] != byte(i+j) {
return base.Errf("spmd: probe %d differs at byte %d", i, j)
}
}
}
return nil
})
if poolDepth() == 0 {
t.Fatal("routed traffic left the pool empty, so no buffer ever rode the chain")
}
}
+209
View File
@@ -0,0 +1,209 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// ExchangeHalos serves the domain-decomposed simulations: a rank whose
// piece is cut along the first axis hands its edge rows to its two
// neighbours and receives theirs. upper carries the rows that precede
// the piece in the global array (the lower edge of rank-1), lower the
// rows that follow it (the upper edge of rank+1); a rank at the
// world's edge receives nil for the side where no neighbour lives,
// and a zero halo width answers the same nil on both sides: nil is
// the answer every zero-length slab takes. The exchange moves bits
// and moves nothing else, so there is
// no order for it to get wrong: the runs in both directions are fixed
// by the rank indices, never by the arrival order. It is
// ExchangeHalosOnGrid on the one-dimensional grid of the whole world,
// cut along axis 0.
func (w *World) ExchangeHalos(local *core.Array, halos int) (*core.Array, *core.Array, error) {
return w.ExchangeHalosOnGrid(local, halos, 0, []int{w.size})
}
// ExchangeHalosOnGrid is the halo exchange of a domain-decomposed
// simulation on a process grid. grid lays the world's ranks out
// row-major, grid[a] naming how many positions the grid runs along
// axis a, and its product must be the world's size; rank r sits at
// coordinate (r/prod(grid[axis+1:]))%grid[axis] along axis, so its
// two neighbours there sit one grid step away, at rank-dist and
// rank+dist with dist = prod(grid[axis+1:]), and a rank at the grid's
// edge has no neighbour on that side and receives nil for it. Each
// rank hands the slab of width halos at its piece's edge along axis
// to the neighbour it borders there and receives theirs: upper
// carries the slab that precedes the piece along the axis, lower the
// slab that follows it, both keeping every other dimension of the
// piece's shape. A zero halo width answers nil on both sides, an
// empty piece joins with empty wires and answers no halos, and a
// neighbour with no piece to offer answers nil too, so the protocol
// stays symmetric. The exchange moves bits and moves nothing else, so
// there is no order
// for it to get wrong: the runs in both directions are fixed by the
// rank indices, never by the arrival order.
func (w *World) ExchangeHalosOnGrid(local *core.Array, halos, axis int, grid []int) (*core.Array, *core.Array, error) {
if err := w.status(); err != nil {
return nil, nil, err
}
if halos < 0 {
return nil, nil, w.fail(base.Errf("spmd: a negative halo width %d", halos))
}
if local.NDim() == 0 {
return nil, nil, w.fail(base.Errf("spmd: the halo exchange needs a dimension to cut the edges from"))
}
if axis < 0 || axis >= local.NDim() {
return nil, nil, w.fail(base.Errf("spmd: halo axis %d is outside the array's %d dimensions", axis, local.NDim()))
}
if axis >= len(grid) {
return nil, nil, w.fail(base.Errf("spmd: halo axis %d is outside the grid's %d axes", axis, len(grid)))
}
extent := 1
for _, d := range grid {
if d < 0 {
return nil, nil, w.fail(base.Errf("spmd: a grid cannot name a negative extent %d", d))
}
if extent > 0 && d > math.MaxInt/extent {
return nil, nil, w.fail(base.Errf("spmd: a grid of %s lays out more ranks than the world can hold", base.ShapeText(grid)))
}
extent *= d
}
if extent != w.size {
return nil, nil, w.fail(base.Errf("spmd: a grid of %s lays out %d ranks against the world's %d", base.ShapeText(grid), extent, w.size))
}
// The rank's place in the row-major grid: the coordinate along the
// cut axis, and the rank distance of one grid step along it.
dist := 1
for _, d := range grid[axis+1:] {
dist *= d
}
coord := (w.rank / dist) % grid[axis]
hasLower, hasUpper := coord > 0, coord+1 < grid[axis]
shape := local.Shape()
if local.Len() == 0 {
// An empty piece owns no slab, so it carries no edges: it
// still joins the exchange with empty wires, so the
// neighbours' protocol stays symmetric, and answers no halos.
emptyShape := make([]int, len(shape))
copy(emptyShape, shape)
emptyShape[axis] = 0
empty, err := encodeHead(nil, local.Dtype(), emptyShape)
if err != nil {
return nil, nil, w.fail(err)
}
if hasUpper {
if err := w.sendTo(w.rank+dist, tagHalo, empty); err != nil {
return nil, nil, err
}
}
if hasLower {
if err := w.sendTo(w.rank-dist, tagHalo, empty); err != nil {
return nil, nil, err
}
}
if hasUpper {
if _, err := w.recvFrom(w.rank+dist, tagHalo); err != nil {
return nil, nil, err
}
}
if hasLower {
if _, err := w.recvFrom(w.rank-dist, tagHalo); err != nil {
return nil, nil, err
}
}
return nil, nil, nil
}
span := shape[axis]
if halos > span {
return nil, nil, w.fail(base.Errf("spmd: a halo width of %d rows exceeds the piece's %d", halos, span))
}
lead, rest := 1, 1
for _, d := range shape[:axis] {
lead *= d
}
for _, d := range shape[axis+1:] {
rest *= d
}
edgeShape := make([]int, len(shape))
copy(edgeShape, shape)
edgeShape[axis] = halos
// The edge slab is one run of halos*rest elements per leading
// position, contiguous in the row-major layout; the wire carries
// the runs joined under the slab's own shape, whose extent along
// the axis is the halo width and whose other extents are the
// piece's own.
edge := func(high bool) ([]byte, error) {
wire, err := encodeHead(nil, local.Dtype(), edgeShape)
if err != nil {
return nil, err
}
count := halos * rest
for p := range lead {
first := p*span*rest + (span-halos)*rest
if !high {
first = p * span * rest
}
part, err := encodePart(nil, local, []int{count}, first, count)
if err != nil {
return nil, err
}
wire = append(wire, part[partHeadLen:]...)
}
return wire, nil
}
// Two phases: the lower halos travel to the upper neighbour
// first, the upper halos to the lower one second. The sends ride
// the links' buffers and the networked links drain through the
// hub, so no rank waits on a rank that is waiting on it.
if hasUpper {
highEdge, err := edge(true)
if err != nil {
return nil, nil, w.fail(err)
}
if err := w.sendTo(w.rank+dist, tagHalo, highEdge); err != nil {
return nil, nil, err
}
}
var lower *core.Array
if hasUpper {
data, err := w.recvFrom(w.rank+dist, tagHalo)
if err != nil {
return nil, nil, err
}
lower, err = decodeWire(data)
if err != nil {
return nil, nil, w.fail(err)
}
if lower.Len() == 0 {
lower = nil // the neighbour owns no slab
}
}
if hasLower {
lowEdge, err := edge(false)
if err != nil {
return nil, nil, w.fail(err)
}
if err := w.sendTo(w.rank-dist, tagHalo, lowEdge); err != nil {
return nil, nil, err
}
}
var upper *core.Array
if hasLower {
data, err := w.recvFrom(w.rank-dist, tagHalo)
if err != nil {
return nil, nil, err
}
upper, err = decodeWire(data)
if err != nil {
return nil, nil, w.fail(err)
}
if upper.Len() == 0 {
upper = nil // the neighbour owns no slab
}
}
return upper, lower, nil
}
+372
View File
@@ -0,0 +1,372 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"slices"
"strings"
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The grid halo tests pin ExchangeHalosOnGrid against the whole array:
// the rank's row-major place in the grid, its neighbours along the cut
// axis and the slabs that travel between them, with the serial stencil
// as the exacter judge of the geometry.
// gridFixture builds the [rows, cols] fixture the grid tests deal out.
func gridFixture(rows, cols int) *core.Array {
vals := make([]float64, rows*cols)
for i := range vals {
vals[i] = float64(i*37%(rows*cols)) * 0.5
}
return mk(core.FromFloats(vals, rows, cols))
}
// tileSpan cuts tile idx's run of an extent dealt into parts tiles,
// the deal that leaves the remainders on the last tiles.
func tileSpan(extent, parts, idx int) (int, int) {
return idx * extent / parts, (idx + 1) * extent / parts
}
// gridCoords maps a rank into a row-major grid, coordinate a being
// (rank/prod(grid[a+1:]))%grid[a].
func gridCoords(rank int, grid []int) []int {
coords := make([]int, len(grid))
rest := 1
for a := len(grid) - 1; a >= 0; a-- {
coords[a] = (rank / rest) % grid[a]
rest *= grid[a]
}
return coords
}
// gridRank is gridCoords' inverse: the rank a row-major grid puts on
// the named coordinates.
func gridRank(coords []int, grid []int) int {
rank := 0
for a := range grid {
rank = rank*grid[a] + coords[a]
}
return rank
}
// dealTile cuts rows [rlo, rhi) by cols [clo, chi) off a 2-D whole,
// an empty array for an empty range.
func dealTile(t *testing.T, whole *core.Array, rlo, rhi, clo, chi int) *core.Array {
t.Helper()
width := whole.Shape()[1]
wire, err := encodeHead(nil, whole.Dtype(), []int{rhi - rlo, chi - clo})
if err != nil {
t.Fatal(err)
}
for r := rlo; r < rhi; r++ {
part, err := encodePart(nil, whole, []int{chi - clo}, r*width+clo, chi-clo)
if err != nil {
t.Fatal(err)
}
wire = append(wire, part[2+8:]...)
}
a, err := decodeWire(wire)
if err != nil {
t.Fatal(err)
}
return a
}
// checkGridExchange runs one rank's grid exchange and judges its halos
// against the neighbouring tiles' own edge slabs, worked out from the
// row-major grid arithmetic beside the exchange: a halo keeps the
// piece's other dimensions, so the expected slab is the neighbour's
// tile narrowed to the halo width along the cut axis. Empty pieces
// answer nothing, and empty neighbours and grid edges hand nil back.
func checkGridExchange(t *testing.T, w *World, whole *core.Array, halos, axis int, grid []int) {
t.Helper()
two := grid
if len(two) == 1 {
two = []int{two[0], 1}
}
coords := gridCoords(w.Rank(), two)
rlo, rhi, clo, chi := rankTile(whole, two, w.Rank())
local := dealTile(t, whole, rlo, rhi, clo, chi)
upper, lower, err := w.ExchangeHalosOnGrid(local, halos, axis, grid)
if err != nil {
t.Fatalf("rank %d: %v", w.Rank(), err)
}
if local.Len() == 0 {
if upper != nil || lower != nil {
t.Fatalf("rank %d: an empty piece answered halos", w.Rank())
}
return
}
sides := []struct {
name string
got *core.Array
off int
}{
{"upper", upper, -1},
{"lower", lower, 1},
}
for _, side := range sides {
got := side.got
if got != nil && got.Len() == 0 {
got = nil // a zero-width slab answers nothing
}
c := coords[axis] + side.off
if c < 0 || c >= two[axis] {
if got != nil {
t.Fatalf("rank %d: a %s halo arrived where no neighbour lives", w.Rank(), side.name)
}
continue
}
sideCoords := slices.Clone(coords)
sideCoords[axis] = c
nrlo, nrhi, nclo, nchi := rankTile(whole, two, gridRank(sideCoords, two))
if nrhi == nrlo || nchi == nclo {
// The neighbour owns no elements, so it carried no slab.
if got != nil {
t.Fatalf("rank %d: a %s halo arrived from an empty neighbour", w.Rank(), side.name)
}
continue
}
// The neighbour's slab: its tile's edge along the cut axis,
// kept whole on every other dimension.
cutLo, cutHi := nclo, nchi
if axis == 0 {
cutLo, cutHi = nrlo, nrhi
}
sLo, sHi := cutLo, cutLo+halos
if side.off < 0 {
sLo, sHi = cutHi-halos, cutHi
}
var want *core.Array
if axis == 0 {
want = dealTile(t, whole, sLo, sHi, nclo, nchi)
} else {
want = dealTile(t, whole, nrlo, nrhi, sLo, sHi)
}
if want.Len() == 0 {
if got != nil {
t.Fatalf("rank %d: a %s halo arrived for a zero-width slab", w.Rank(), side.name)
}
continue
}
if got == nil {
t.Fatalf("rank %d: the %s halo is nil with a living neighbour", w.Rank(), side.name)
}
if !slices.Equal(got.Shape(), want.Shape()) || !sameBits(want, got) {
t.Fatalf("rank %d: the %s halo differs from the neighbouring tile's slab along axis %d",
w.Rank(), side.name, axis)
}
}
}
// rankTile is dealTile's own cut for one rank of a 2-D grid.
func rankTile(whole *core.Array, grid []int, rank int) (rlo, rhi, clo, chi int) {
coords := gridCoords(rank, grid)
rlo, rhi = tileSpan(whole.Shape()[0], grid[0], coords[0])
clo, chi = tileSpan(whole.Shape()[1], grid[1], coords[1])
return rlo, rhi, clo, chi
}
// TestExchangeHalosOnGridMatchesTheWhole walks the whole matrix the
// API claims: world sizes 1 to 8, halo widths 0, 1 and 3, every
// two-factor grid of each size beside the degenerate one-dimensional
// one, cut along every axis the grid names. The expected slabs come
// from the rank's grid coordinates worked out beside the exchange, so
// a wrong row-major mapping, a wrong neighbour distance or a wrong
// slab cannot pass: on the grid [2, 3] the axis 1 neighbours are
// rank+-1 and the axis 0 neighbours rank+-3, and the checks know it
// independently.
func TestExchangeHalosOnGridMatchesTheWhole(t *testing.T) {
// The extents divide out to tiles of at least three elements on
// every factor grid up to eight positions, so even the widest
// pinned halo fits the piece it travels from.
const rows, cols = 24, 24
whole := gridFixture(rows, cols)
for _, size := range []int{1, 2, 3, 5, 8} {
grids := [][]int{{size}}
for a := 1; a <= size; a++ {
if size%a == 0 {
grids = append(grids, []int{a, size / a})
}
}
for _, grid := range grids {
for _, halos := range []int{0, 1, 3} {
for axis := range len(grid) {
err := Launch(size, func(w *World) error {
checkGridExchange(t, w, whole, halos, axis, grid)
return nil
})
if err != nil {
t.Fatalf("size %d grid %v halos %d axis %d: %v", size, grid, halos, axis, err)
}
}
}
}
}
}
// TestExchangeHalosOnGridEmptyPieces: the pieces the tiling leaves
// empty join the exchange symmetrically, answering no halos, and the
// neighbours see nil from the empty side, in both grid orientations.
func TestExchangeHalosOnGridEmptyPieces(t *testing.T) {
whole := gridFixture(1, 2)
for _, grid := range [][]int{{2, 2}, {3, 2}, {2, 3}} {
for axis := range 2 {
err := Launch(grid[0]*grid[1], func(w *World) error {
checkGridExchange(t, w, whole, 1, axis, grid)
return nil
})
if err != nil {
t.Fatalf("grid %v axis %d: %v", grid, axis, err)
}
}
}
}
// TestExchangeHalosOnGridStencilMatchesSerial runs the five-point
// stencil over a field tiled 2 by 3, exchanging halos along both grid
// axes, and compares every tile element for element against the same
// step computed on the whole field. The comparison is the matrix
// claim made flesh: on the row-major grid [2, 3] the axis 1
// neighbours are rank+-1, the axis 0 neighbours rank+-3, and the
// distributed stencil may only agree bit for bit when the slabs
// deliver exactly the neighbours the serial walk sees.
func TestExchangeHalosOnGridStencilMatchesSerial(t *testing.T) {
const rows, cols = 6, 12
whole := gridFixture(rows, cols)
serial := make([]float64, rows*cols)
for r := range rows {
for c := range cols {
up, down := 0.0, 0.0
if r > 0 {
up = whole.FloatAt((r-1)*cols + c)
}
if r < rows-1 {
down = whole.FloatAt((r+1)*cols + c)
}
left, right := 0.0, 0.0
if c > 0 {
left = whole.FloatAt(r*cols + c - 1)
}
if c < cols-1 {
right = whole.FloatAt(r*cols + c + 1)
}
serial[r*cols+c] = (up + left + whole.FloatAt(r*cols+c) + right + down) / 5
}
}
const halos = 1
err := Launch(6, func(w *World) error {
rlo, rhi, clo, chi := rankTile(whole, []int{2, 3}, w.Rank())
local := dealTile(t, whole, rlo, rhi, clo, chi)
upper, lower, err := w.ExchangeHalosOnGrid(local, halos, 0, []int{2, 3})
if err != nil {
return err
}
leftHalo, rightHalo, err := w.ExchangeHalosOnGrid(local, halos, 1, []int{2, 3})
if err != nil {
return err
}
tileRows, tileCols := rhi-rlo, chi-clo
for ri := range tileRows {
for cj := range tileCols {
gi, gj := rlo+ri, clo+cj
centre := local.FloatAt(ri*tileCols + cj)
up := 0.0
if ri == 0 {
if gi > 0 {
up = upper.FloatAt(cj) // the halo's last row adjoins the tile
}
} else {
up = local.FloatAt((ri-1)*tileCols + cj)
}
down := 0.0
if ri == tileRows-1 {
if gi < rows-1 {
down = lower.FloatAt(cj)
}
} else {
down = local.FloatAt((ri+1)*tileCols + cj)
}
left := 0.0
if cj == 0 {
if gj > 0 {
left = leftHalo.FloatAt(ri) // the halo's last column adjoins the tile
}
} else {
left = local.FloatAt(ri*tileCols + cj - 1)
}
right := 0.0
if cj == tileCols-1 {
if gj < cols-1 {
right = rightHalo.FloatAt(ri)
}
} else {
right = local.FloatAt(ri*tileCols + cj + 1)
}
got := (up + left + centre + right + down) / 5
if got != serial[gi*cols+gj] {
t.Fatalf("rank %d global (%d, %d): %v against the serial %v",
w.Rank(), gi, gj, got, serial[gi*cols+gj])
}
}
}
return w.Barrier()
})
if err != nil {
t.Fatal(err)
}
}
// TestExchangeHalosOnGridOverTCP runs the grid exchange over real
// connections, the axis 1 slabs included.
func TestExchangeHalosOnGridOverTCP(t *testing.T) {
whole := gridFixture(24, 18)
runTCPWorld(t, 6, Options{Timeout: 30 * time.Second}, func(w *World) error {
checkGridExchange(t, w, whole, 2, 1, []int{2, 3})
checkGridExchange(t, w, whole, 2, 0, []int{2, 3})
return nil
})
}
// TestExchangeHalosOnGridRefusesTheHostile: every invalid argument is
// a named error before any frame moves. Every rank makes the same
// invalid call, so an exchange that ever started would deadlock the
// world instead of answering, which is what pins the ordering too.
func TestExchangeHalosOnGridRefusesTheHostile(t *testing.T) {
local := mk(core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3))
cases := []struct {
name string
local *core.Array
halos, axis int
grid []int
want string
}{
{"negative halos", local, -1, 0, []int{2}, "negative halo width"},
{"axis past the array", local, 1, 2, []int{2, 2}, "outside the array"},
{"negative axis", local, 1, -1, []int{2, 2}, "outside the array"},
{"axis past the grid", local, 1, 1, []int{2}, "outside the grid"},
{"grid short of the world", local, 1, 0, []int{1, 3}, "against the world's"},
{"negative grid extent", local, 1, 0, []int{-2}, "negative extent"},
{"halos wider than the axis", local, 3, 0, []int{2}, "exceeds the piece's"},
}
for _, tc := range cases {
err := Launch(2, func(w *World) error {
_, _, err := w.ExchangeHalosOnGrid(tc.local, tc.halos, tc.axis, tc.grid)
if err == nil {
t.Fatalf("%s: the hostile call was accepted", tc.name)
}
if !strings.Contains(err.Error(), tc.want) {
t.Fatalf("%s: the error %q does not name the fault", tc.name, err)
}
return nil
})
if err != nil {
t.Fatalf("%s: %v", tc.name, err)
}
}
}
+180
View File
@@ -0,0 +1,180 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// globalRows builds the [rows, width] fixture the halo tests deal out.
func globalRows(rows, width int) *core.Array {
vals := make([]float64, rows*width)
for i := range vals {
vals[i] = float64(i*31%(rows*width)) * 0.25
}
a, err := core.FromFloats(vals, rows, width)
if err != nil {
panic(err)
}
return a
}
// dealRows2D cuts rows [lo, hi) off the global array, an empty array
// for an empty range.
func dealRows2D(tb testing.TB, whole *core.Array, lo, hi int) *core.Array {
tb.Helper()
rest := 1
for _, d := range whole.Shape()[1:] {
rest *= d
}
wire, err := encodePart(nil, whole, append([]int{hi - lo}, whole.Shape()[1:]...), lo*rest, (hi-lo)*rest)
if err != nil {
tb.Fatal(err)
}
a, err := decodeWire(wire)
if err != nil {
tb.Fatal(err)
}
return a
}
// checkHalos asserts one rank's halos against the whole's neighbouring
// rows, treating empty neighbours and empty pieces as no data.
func checkHalos(t *testing.T, w *World, whole *core.Array, rows, halos int, upper, lower *core.Array) {
t.Helper()
span := mustPartition(t, rows, w.Size(), w.Rank())
neighbourRows := func(rank int) int {
if rank < 0 || rank >= w.Size() {
return 0
}
return mustPartition(t, rows, w.Size(), rank).Len()
}
if span.Len() == 0 {
if upper != nil || lower != nil {
t.Fatalf("rank %d: an empty piece answered halos", w.Rank())
}
return
}
if upper != nil && upper.Len() == 0 {
upper = nil
}
if lower != nil && lower.Len() == 0 {
lower = nil
}
if w.Rank() == 0 || neighbourRows(w.Rank()-1) == 0 {
if upper != nil {
t.Fatalf("rank %d: an upper halo arrived where no rows live", w.Rank())
}
} else {
if upper == nil {
t.Fatalf("rank %d: the upper halo is nil with a living neighbour", w.Rank())
}
if want := dealRows2D(t, whole, span.Lo-halos, span.Lo); !sameBits(want, upper) {
t.Fatalf("rank %d: the upper halo differs from rows [%d, %d)", w.Rank(), span.Lo-halos, span.Lo)
}
}
if w.Rank()+1 == w.Size() || neighbourRows(w.Rank()+1) == 0 {
if lower != nil {
t.Fatalf("rank %d: a lower halo arrived where no rows live", w.Rank())
}
} else {
if lower == nil {
t.Fatalf("rank %d: the lower halo is nil with a living neighbour", w.Rank())
}
if want := dealRows2D(t, whole, span.Hi, span.Hi+halos); !sameBits(want, lower) {
t.Fatalf("rank %d: the lower halo differs from rows [%d, %d)", w.Rank(), span.Hi, span.Hi+halos)
}
}
}
// TestExchangeHalosMatchTheWhole: every rank's upper and lower halos
// are exactly the whole's neighbouring rows, and edges without a
// living neighbour answer nil.
func TestExchangeHalosMatchTheWhole(t *testing.T) {
for _, size := range []int{1, 2, 3, 5, 8} {
for _, halos := range []int{0, 1, 3} {
const rows, width = 40, 5
whole := globalRows(rows, width)
err := Launch(size, func(w *World) error {
span := mustPartition(t, rows, w.Size(), w.Rank())
local := dealRows2D(t, whole, span.Lo, span.Hi)
upper, lower, err := w.ExchangeHalos(local, halos)
if err != nil {
return err
}
checkHalos(t, w, whole, rows, halos, upper, lower)
return nil
})
if err != nil {
t.Fatalf("size %d halos %d: %v", size, halos, err)
}
}
}
}
// TestExchangeHalosStencilMatchesSerial runs a three-point smoothing
// step over a sharded series and compares it, element for element,
// with the same step computed on the whole array: the halos must make
// the distributed stencil see exactly the neighbours the serial one
// sees.
func TestExchangeHalosStencilMatchesSerial(t *testing.T) {
const n = 1000
whole := fixtureArray(n)
serial := make([]float64, n)
for i := 1; i < n-1; i++ {
serial[i] = (whole.FloatAt(i-1) + whole.FloatAt(i) + whole.FloatAt(i+1)) / 3
}
err := Launch(4, func(w *World) error {
span := mustPartition(t, n, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
upper, lower, err := w.ExchangeHalos(local, 1)
if err != nil {
return err
}
for i := range span.Len() {
g := span.Lo + i
if g == 0 || g == n-1 {
continue
}
var left, right float64
if i == 0 {
left = upper.FloatAt(0)
} else {
left = local.FloatAt(i - 1)
}
if i == span.Len()-1 {
right = lower.FloatAt(0)
} else {
right = local.FloatAt(i + 1)
}
if got := (left + local.FloatAt(i) + right) / 3; got != serial[g] {
t.Fatalf("rank %d global %d: %v against the serial %v", w.Rank(), g, got, serial[g])
}
}
return w.Barrier()
})
if err != nil {
t.Fatal(err)
}
}
// TestExchangeHalosOverTCP runs the neighbour exchange over real
// connections.
func TestExchangeHalosOverTCP(t *testing.T) {
const rows, width = 200, 3
whole := globalRows(rows, width)
runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error {
span := mustPartition(t, rows, w.Size(), w.Rank())
local := dealRows2D(t, whole, span.Lo, span.Hi)
upper, lower, err := w.ExchangeHalos(local, 2)
if err != nil {
return err
}
checkHalos(t, w, whole, rows, 2, upper, lower)
return nil
})
}
+79
View File
@@ -0,0 +1,79 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Span names one rank's contiguous piece of a global axis of data: the
// global length, the piece's first element and one past its last.
type Span struct {
// Global is the axis length every rank's piece is a cut of.
Global int
// Lo is the piece's first global element index; Hi is one past its
// last. The piece is Lo..Hi, possibly empty.
Lo, Hi int
}
// Len returns the number of elements the piece carries along the axis.
func (s Span) Len() int { return s.Hi - s.Lo }
// Partition cuts a global axis of globalN elements into the world's
// contiguous pieces: rank rank's piece is [Lo, Hi). Every piece starts
// and ends on a boundary of the canonical fold partition, which is what
// lets the shards' reductions compose into the single-array fold's
// exact bits, so a program that cuts its data any other way gives up
// that contract. With more ranks than blocks, the pieces beyond the blocks are empty,
// wherever the boundary falls, so rank 0 may hold nothing at all.
//
// A negative globalN, a size below 1 or a rank outside [0, size) is
// refused with an error naming the input. Valid arguments never fail.
func Partition(globalN, size, rank int) (Span, error) {
if globalN < 0 {
return Span{}, base.Errf("spmd: Partition of a negative global length %d", globalN)
}
if size < 1 || rank < 0 || rank >= size {
return Span{}, base.Errf("spmd: Partition of rank %d in a world of %d", rank, size)
}
parts := core.FoldParts(globalN)
return Span{
Global: globalN,
Lo: core.FoldBoundary(globalN, rank*parts/size),
Hi: core.FoldBoundary(globalN, (rank+1)*parts/size),
}, nil
}
// checkSpan verifies that a caller's span is the canonical partition's
// own cut for this rank and that the local slab leads with exactly the
// span's run of the axis. Any other cut is refused by name: the
// bit-identity contract lives on the canonical boundaries.
func (w *World) checkSpan(span Span, local *core.Array) error {
want, err := Partition(span.Global, w.size, w.rank)
if err != nil {
return err
}
if span != want {
return base.Errf("spmd: rank %d holds [%d, %d) of %d, but the canonical partition puts this rank on [%d, %d); cut the data with Partition",
w.rank, span.Lo, span.Hi, span.Global, want.Lo, want.Hi)
}
if local.NDim() == 0 || local.Shape()[0] != span.Len() {
lead := 0
if local.NDim() > 0 {
lead = local.Shape()[0]
}
return base.Errf("spmd: rank %d's slab leads with %d elements for a span of %d", w.rank, lead, span.Len())
}
return nil
}
// checkOneDimSpan is checkSpan for the vector reductions: the shards
// carry one-dimensional arrays, whose length is the span's run.
func (w *World) checkOneDimSpan(local *core.Array, span Span) error {
if local.NDim() != 1 {
return base.Errf("spmd: the sharded product, norm and dot carry 1-D arrays; got %d dimensions", local.NDim())
}
return w.checkSpan(span, local)
}
+44
View File
@@ -0,0 +1,44 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"strings"
"testing"
)
// mustPartition is Partition for the tests' fixed valid arguments: the
// piece, or a failure the tests can never meet. Benchmarks pass their
// own tb.
func mustPartition(tb testing.TB, globalN, size, rank int) Span {
tb.Helper()
span, err := Partition(globalN, size, rank)
if err != nil {
tb.Fatal(err)
}
return span
}
// TestPartitionRefusesInvalidArguments: the inputs Partition refuses
// come back as errors naming the input, never as panics.
func TestPartitionRefusesInvalidArguments(t *testing.T) {
for _, tc := range []struct {
name string
globalN, size, rank int
want string
}{
{"negative global length", -1, 4, 0, "Partition of a negative global length -1"},
{"empty world", 100, 0, 0, "Partition of rank 0 in a world of 0"},
{"negative rank", 100, 4, -1, "Partition of rank -1 in a world of 4"},
{"rank past the world", 100, 4, 4, "Partition of rank 4 in a world of 4"},
} {
span, err := Partition(tc.globalN, tc.size, tc.rank)
if err == nil {
t.Fatalf("%s: Partition answered %v", tc.name, span)
}
if !strings.Contains(err.Error(), tc.want) {
t.Fatalf("%s: the error does not name the input: %v", tc.name, err)
}
}
}
+201
View File
@@ -0,0 +1,201 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"fmt"
"math"
"net"
"os"
"os/exec"
"strconv"
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The loopback tests run the package as a cluster: N real processes on
// one machine, rank 0 inside the test process and the rest re-executed
// out of this same test binary. It is the same code a multi-machine
// run executes, which is what makes the multi-machine contract
// testable without one.
func TestMain(m *testing.M) {
if os.Getenv("TENSOR_SPMD_WORKER") != "" {
os.Exit(spmdWorker(os.Getenv("TENSOR_SPMD_ADDR")))
}
os.Exit(m.Run())
}
// spmdWorker is the program a re-executed worker process runs: join
// the world, reduce the shards, compare the bits, leave cleanly. The
// fixture is deterministic, so no data crosses with the address.
func spmdWorker(addr string) int {
rank, _ := strconv.Atoi(os.Getenv("TENSOR_SPMD_RANK"))
size, _ := strconv.Atoi(os.Getenv("TENSOR_SPMD_SIZE"))
w, err := Join(addr, Options{Timeout: 2 * time.Minute})
if err != nil {
fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err)
return 1
}
const gn = 131073
whole := fixtureArray(gn)
want := core.Sum(whole)
span, err := Partition(gn, size, w.Rank())
if err != nil {
fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err)
return 1
}
local, err := dealRows(whole, span)
if err != nil {
fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err)
return 1
}
got, err := w.AllReduceShards(local, span, Sum)
if err != nil {
fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err)
return 1
}
if scalarBits(got) != scalarBits(want) {
fmt.Fprintf(os.Stderr, "worker %d: sharded %s against single-array %s\n",
rank, scalarBits(got), scalarBits(want))
return 1
}
if err := w.Barrier(); err != nil {
fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err)
return 1
}
if err := w.Close(); err != nil {
fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err)
return 1
}
return 0
}
// dealRows cuts a shard's rows off a 1-D whole, bits exactly.
func dealRows(whole *core.Array, span Span) (*core.Array, error) {
wire, err := encodePart(nil, whole, []int{span.Len()}, span.Lo, span.Len())
if err != nil {
return nil, err
}
return decodeWire(wire)
}
func TestMultiprocessLoopback(t *testing.T) {
const size = 4
ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer ln.Close()
rankErr := make(chan error, 1)
go func() {
w, err := listen(ln, size, Options{Timeout: 2 * time.Minute})
if err != nil {
rankErr <- err
return
}
defer w.Close()
const gn = 131073
whole := fixtureArray(gn)
want := core.Sum(whole)
span, err := Partition(gn, size, w.Rank())
if err != nil {
rankErr <- err
return
}
local, err := dealRows(whole, span)
if err != nil {
rankErr <- err
return
}
got, err := w.AllReduceShards(local, span, Sum)
if err != nil {
rankErr <- err
return
}
if scalarBits(got) != scalarBits(want) {
rankErr <- fmt.Errorf("rank 0: sharded %s against single-array %s",
scalarBits(got), scalarBits(want))
return
}
if err := w.Barrier(); err != nil {
rankErr <- err
return
}
rankErr <- nil
}()
workers := make([]*exec.Cmd, size-1)
for r := 1; r < size; r++ {
cmd := exec.Command(os.Args[0], "-test.run=^$")
cmd.Env = append(os.Environ(),
"TENSOR_SPMD_WORKER=1",
"TENSOR_SPMD_ADDR="+ln.Addr().String(),
"TENSOR_SPMD_RANK="+strconv.Itoa(r),
"TENSOR_SPMD_SIZE="+strconv.Itoa(size))
workers[r-1] = cmd
}
for r, cmd := range workers {
if err := cmd.Start(); err != nil {
t.Fatalf("worker %d never started: %v", r+1, err)
}
}
for r, cmd := range workers {
if err := cmd.Wait(); err != nil {
t.Fatalf("worker %d failed: %v", r+1, err)
}
}
if err := <-rankErr; err != nil {
t.Fatalf("rank 0: %v", err)
}
}
// TestRepeatedRunsAnswerIdenticalBits is the arrival-order claim,
// exercised: one world runs the sharded reduction and the movement
// collectives interleaved many times, and six worlds run the whole
// battery again, so goroutine scheduling arrives at every order it
// can find and the bits may not move once.
func TestRepeatedRunsAnswerIdenticalBits(t *testing.T) {
const gn = 200001
whole := fixtureArray(gn)
want := core.Sum(whole)
wantMax, err := core.Max(whole)
if err != nil {
t.Fatal(err)
}
for attempt := range 6 {
err := Launch(5, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local, err := dealRows(whole, span)
if err != nil {
return err
}
for i := range 8 {
got, err := w.AllReduceShards(local, span, Sum)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("attempt %d run %d: sharded %s against single-array %s",
attempt, i, scalarBits(got), scalarBits(want))
}
gotMax, err := w.AllReduceShards(local, span, Max)
if err != nil {
return err
}
if math.Float64bits(gotMax.Float()) != math.Float64bits(wantMax.Float()) {
t.Fatalf("attempt %d run %d: sharded max moved", attempt, i)
}
if _, err := w.AllGather(local); err != nil {
return err
}
}
return nil
})
if err != nil {
t.Fatalf("attempt %d: %v", attempt, err)
}
}
}
+1291
View File
File diff suppressed because it is too large Load Diff
+886
View File
@@ -0,0 +1,886 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"encoding/binary"
"math"
"strings"
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The sharded reduction tests judge the whole contract in one claim:
// whatever the world's size, the sharded answer carries the
// single-array reduction's exact bits. The single-array answer is
// computed from the same fixture beside the test, so the comparison is
// the contract and not a self-check.
// sparseFixture spreads magnitudes and drops NaNs and an infinity at
// fixed spots, so the folds meet cancellation and the extremum rules
// meet their NaN edges.
func sparseFixture(n int) []float64 {
v := make([]float64, n)
for i := range v {
v[i] = float64((i*6559)%2001-1000) * math.Pow(10, float64(i%7)-3)
if i%401 == 3 {
v[i] = math.NaN()
}
if i%401 == 200 {
v[i] = math.Inf(-1)
}
}
return v
}
// fixtureDtypes builds the same data in every element type the shard
// reductions carry.
func fixtureDtypes(t *testing.T, gn int) map[core.Dtype]*core.Array {
t.Helper()
build := func(a *core.Array, err error) *core.Array {
if err != nil {
t.Fatal(err)
}
return a
}
vals := sparseFixture(gn)
ints := make([]int64, gn)
bools := make([]bool, gn)
for i := range ints {
ints[i] = int64((i*13)%97) - 48
bools[i] = i%3 == 0
}
halves := make([]uint16, gn)
f32s := make([]float32, gn)
for i := range halves {
halves[i] = core.HalfFromFloat64(vals[i])
f32s[i] = float32(vals[i])
}
complexes := make([]complex128, gn)
for i := range complexes {
complexes[i] = complex(vals[i], -vals[i]/2)
}
return map[core.Dtype]*core.Array{
core.Float: build(core.FromFloats(vals, gn)),
core.Float32: build(core.FromFloat32s(f32s, gn)),
core.Float16: build(core.HalvesFromArray(halves, gn)),
core.Complex: build(core.FromComplexes(complexes, gn)),
core.Int: build(core.FromInts(ints, gn)),
core.Int8: build(core.FromInt8s(narrow8s(ints), gn)),
core.Uint8: build(core.FromUint8s(narrowu8s(ints), gn)),
core.Int16: build(core.FromInt16s(narrow16s(ints), gn)),
core.Uint16: build(core.FromUint16s(narrowu16s(ints), gn)),
core.Int32: build(core.FromInt32s(narrow32s(ints), gn)),
core.Uint32: build(core.FromUint32s(narrowu32s(ints), gn)),
core.Bool: build(core.FromBools(bools, gn)),
}
}
func narrow8s(v []int64) []int8 {
out := make([]int8, len(v))
for i := range v {
out[i] = int8(v[i])
}
return out
}
func narrowu8s(v []int64) []uint8 {
out := make([]uint8, len(v))
for i := range v {
out[i] = uint8(v[i])
}
return out
}
func narrow16s(v []int64) []int16 {
out := make([]int16, len(v))
for i := range v {
out[i] = int16(v[i])
}
return out
}
func narrowu16s(v []int64) []uint16 {
out := make([]uint16, len(v))
for i := range v {
out[i] = uint16(v[i])
}
return out
}
func narrow32s(v []int64) []int32 {
out := make([]int32, len(v))
for i := range v {
out[i] = int32(v[i])
}
return out
}
func narrowu32s(v []int64) []uint32 {
out := make([]uint32, len(v))
for i := range v {
out[i] = uint32(v[i])
}
return out
}
// narrowSliceFor deals any dtype's rows off the whole array, keeping
// the bits exactly: what Scatter deals, built without the collectives
// under test.
func narrowSliceFor(t *testing.T, whole *core.Array, span Span) *core.Array {
t.Helper()
wire, err := encodePart(nil, whole, append([]int{span.Len()}, whole.Shape()[1:]...), span.Lo, span.Len())
if err != nil {
t.Fatal(err)
}
a, err := decodeWire(wire)
if err != nil {
t.Fatal(err)
}
return a
}
// scalarBits renders a scalar's exact bits, NaN payloads and signed
// zeros included, so the comparison is bit for bit.
func scalarBits(s core.Scalar) string {
switch {
case s.IsComplex():
c := s.Complex()
return "c" + fBits(real(c)) + "/" + fBits(imag(c))
case s.IsFloat():
return "f" + fBits(s.Float())
default:
return "i" + iFormat(s.Int())
}
}
func fBits(f float64) string { return iFormat(int64(math.Float64bits(f))) }
func iFormat(i int64) string {
if i < 0 {
return "-" + uFormat(-uint64(i))
}
return uFormat(uint64(i))
}
func uFormat(u uint64) string {
if u == 0 {
return "0"
}
var b []byte
for u > 0 {
b = append([]byte{byte('0' + u%10)}, b...)
u /= 10
}
return string(b)
}
// TestShardedSumIsTheSingleArraySum is the headline claim of the whole
// package: shard the data, reduce the shards, and the bits are the
// single array's bits, for every world size, every element type and
// every length that touches the partition's edges.
func TestShardedSumIsTheSingleArraySum(t *testing.T) {
for _, size := range []int{1, 2, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537, 131073, 327681} {
for dt, whole := range fixtureDtypes(t, gn) {
want := core.Sum(whole)
err := Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceShards(local, span, Sum)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("size %d gn %d %s: sharded %s against single-array %s",
w.Size(), gn, dt, scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d %s: %v", size, gn, dt, err)
}
}
}
}
}
// TestShardedExtremaAreTheSingleArrayExtrema pins Min and Max against
// the single-array walk, NaN rules and the all-NaN fallback included.
func TestShardedExtremaAreTheSingleArrayExtrema(t *testing.T) {
for _, size := range []int{1, 2, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537, 131073} {
for _, dt := range []core.Dtype{core.Float, core.Float32, core.Float16, core.Int, core.Int8} {
whole := fixtureDtypes(t, gn)[dt]
for _, op := range []Op{Min, Max} {
want, err := (func() (core.Scalar, error) {
if op == Min {
return core.Min(whole)
}
return core.Max(whole)
})()
if err != nil {
t.Fatal(err)
}
err = Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceShards(local, span, op)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("size %d gn %d %s %s: sharded %s against single-array %s",
w.Size(), gn, dt, op, scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d %s %s: %v", size, gn, dt, op, err)
}
}
}
}
}
// The all-NaN array: every block reports no candidate, so the answer
// is the last block's fallback, the single-array walk's own rule.
const gn = 131073 // two blocks
allNaN := make([]float64, gn)
for i := range allNaN {
allNaN[i] = math.NaN()
}
nans := mk(core.FromFloats(allNaN, gn))
for _, op := range []Op{Min, Max} {
want, err := (func() (core.Scalar, error) {
if op == Min {
return core.Min(nans)
}
return core.Max(nans)
})()
if err != nil {
t.Fatal(err)
}
err = Launch(3, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, nans, span)
got, err := w.AllReduceShards(local, span, op)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("all-NaN %s: sharded %s against single-array %s",
op, scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("all-NaN %s: %v", op, err)
}
}
}
// TestShardedBoolReductions pins Any and All against the single-array
// answers.
func TestShardedBoolReductions(t *testing.T) {
for _, size := range []int{1, 2, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537} {
whole := fixtureDtypes(t, gn)[core.Bool]
anyWant, _ := core.Any(whole)
allWant, _ := core.All(whole)
err := Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
anyGot, err := w.AllReduceShards(local, span, Any)
if err != nil {
return err
}
allGot, err := w.AllReduceShards(local, span, All)
if err != nil {
return err
}
if anyGot.Int() != bToF(anyWant) || allGot.Int() != bToF(allWant) {
t.Fatalf("size %d gn %d: any %d all %d against %v/%v",
w.Size(), gn, anyGot.Int(), allGot.Int(), anyWant, allWant)
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d: %v", size, gn, err)
}
}
}
}
// TestReduceShardsAnswersTheRootAlone pins the root-directed shape of
// the collective: the root holds the answer, nobody else holds
// anything.
func TestReduceShardsAnswersTheRootAlone(t *testing.T) {
const gn = 131074
whole := fixtureDtypes(t, gn)[core.Float]
want := core.Sum(whole)
err := Launch(3, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.ReduceShards(local, span, Sum, 2)
if err != nil {
return err
}
if w.Rank() == 2 {
if got == nil || scalarBits(*got) != scalarBits(want) {
t.Fatalf("the root's sharded sum against the single-array %s", scalarBits(want))
}
return nil
}
if got != nil {
t.Fatalf("rank %d received an answer it should not have", w.Rank())
}
return nil
})
if err != nil {
t.Fatal(err)
}
}
// TestShardsRefuseAForeignSpan: a span the partition did not cut is an
// error naming the canonical boundaries, never a number.
func TestShardsRefuseAForeignSpan(t *testing.T) {
const gn = 65537
whole := fixtureDtypes(t, gn)[core.Float]
err := Launch(3, func(w *World) error {
span := mustPartition(t, gn, 3, w.Rank())
local := narrowSliceFor(t, whole, mustPartition(t, gn, 3, w.Rank()))
if w.Rank() == 1 {
// A plausible but foreign cut: one element over.
span.Lo++
}
_, err := w.AllReduceShards(local, span, Sum)
return err
})
if err == nil || !strings.Contains(err.Error(), "canonical partition") {
t.Fatalf("a foreign span did not name the canonical boundaries: %v", err)
}
}
// TestShardedSumOverTCP runs the headline claim over real connections.
func TestShardedSumOverTCP(t *testing.T) {
const gn = 131073
whole := fixtureDtypes(t, gn)[core.Float]
want := core.Sum(whole)
runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceShards(local, span, Sum)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("rank %d: sharded %s against single-array %s", w.Rank(), scalarBits(got), scalarBits(want))
}
return nil
})
}
// TestAllReduceFoldsInRankOrder pins family A: elementwise across the
// ranks' same-shaped arrays, folded in rank index order, which the
// expected answer rebuilds with the core's own pairwise ops.
func TestAllReduceFoldsInRankOrder(t *testing.T) {
for _, size := range []int{1, 2, 3, 5} {
rankArrays := make([]*core.Array, size)
for r := range rankArrays {
vals := make([]float64, 100)
for i := range vals {
vals[i] = float64(r+1) * float64((i*17)%31-15) / 3.0
}
rankArrays[r] = mk(core.FromFloats(vals, 100))
}
want := rankArrays[0]
for _, a := range rankArrays[1:] {
next, err := core.Add(want, a)
if err != nil {
t.Fatal(err)
}
want = next
}
wantMin := rankArrays[0]
for _, a := range rankArrays[1:] {
next, err := core.Minimum(wantMin, a)
if err != nil {
t.Fatal(err)
}
wantMin = next
}
err := Launch(size, func(w *World) error {
got, err := w.AllReduce(rankArrays[w.Rank()], Sum)
if err != nil {
return err
}
if !sameBits(want, got) {
t.Fatalf("size %d: the AllReduce sum differs from the rank-order fold", size)
}
got, err = w.AllReduce(rankArrays[w.Rank()], Min)
if err != nil {
return err
}
if !sameBits(wantMin, got) {
t.Fatalf("size %d: the AllReduce min differs from the rank-order fold", size)
}
return nil
})
if err != nil {
t.Fatalf("size %d: %v", size, err)
}
}
}
// TestShardContributionRefusesTheHostile: the counts the wire names
// are bounded as unsigned values before anything is allocated, so a
// top-bit count is refused rather than turned into a negative length,
// and the candidate flags of an extremum contribution are read back as
// the 0/1 bytes the encoder writes, an honest run and an empty one
// included.
func TestShardContributionRefusesTheHostile(t *testing.T) {
head := make([]byte, 8+1+8*4)
head[8] = kindSumF64
binary.LittleEndian.PutUint64(head[9:], uint64(1)<<63)
if _, err := decodeBlockValues(head); err == nil {
t.Fatal("a top-bit float count decoded")
}
honest := binary.LittleEndian.AppendUint64(nil, 0)
honest = append(honest, kindSumF64)
honest = binary.LittleEndian.AppendUint64(honest, 1)
honest = binary.LittleEndian.AppendUint64(honest, 0)
honest = binary.LittleEndian.AppendUint64(honest, 0)
honest = binary.LittleEndian.AppendUint64(honest, 0)
honest = binary.LittleEndian.AppendUint64(honest, math.Float64bits(1))
bv, err := decodeBlockValues(honest)
if err != nil || len(bv.f) != 1 || bv.f[0] != 1 {
t.Fatalf("an honest contribution was refused: %v", err)
}
extremum := func(flags ...byte) []byte {
buf := binary.LittleEndian.AppendUint64(nil, 0)
buf = append(buf, kindExtF64)
buf = binary.LittleEndian.AppendUint64(buf, uint64(len(flags)))
buf = binary.LittleEndian.AppendUint64(buf, 0)
buf = binary.LittleEndian.AppendUint64(buf, 0)
buf = binary.LittleEndian.AppendUint64(buf, 0)
for range flags {
buf = binary.LittleEndian.AppendUint64(buf, math.Float64bits(1))
}
return append(buf, flags...)
}
if _, err := decodeBlockValues(extremum()); err != nil {
t.Fatalf("an empty flag run was refused: %v", err)
}
bv, err = decodeBlockValues(extremum(1, 0))
if err != nil {
t.Fatalf("honest candidate flags were refused: %v", err)
}
if !bv.oks[0] || bv.oks[1] {
t.Fatalf("the candidate flags came back as %v", bv.oks)
}
if _, err := decodeBlockValues(extremum(1, 2)); err == nil {
t.Fatal("a candidate flag byte of 2 decoded")
}
}
// TestShardedSumAtZeroLength: the degenerate global axis answers the
// single-array reduction's zero at any world size, no exchange needed.
func TestShardedSumAtZeroLength(t *testing.T) {
for _, size := range []int{1, 2, 3, 5} {
empty := mk(core.FromFloats(nil, 0))
want := core.Sum(empty)
err := Launch(size, func(w *World) error {
span := mustPartition(t, 0, w.Size(), w.Rank())
got, err := w.AllReduceShards(empty, span, Sum)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("size %d: sharded %s against single-array %s",
w.Size(), scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("size %d: %v", size, err)
}
}
}
// TestShardedExtremaReachTheRemotePieces pins the candidate flags'
// travel: the extremum living strictly on a non-root rank, without a
// tie to mask it, must win the combine whoever holds it. The saturating
// fixture above cannot tell this, because its extrema repeat in every
// window and the tie rule hands the answer to the root anyway.
func TestShardedExtremaReachTheRemotePieces(t *testing.T) {
// Ascending data: the maximum lives on the last rank alone, and the
// minimum on the root. The root holds its own block here.
const gn = 131073
asc := make([]float64, gn)
for i := range asc {
asc[i] = float64(i)
}
whole := mk(core.FromFloats(asc, gn))
for _, op := range []Op{Min, Max} {
want, err := (func() (core.Scalar, error) {
if op == Min {
return core.Min(whole)
}
return core.Max(whole)
})()
if err != nil {
t.Fatal(err)
}
err = Launch(2, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceShards(local, span, op)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("%s: sharded %s against single-array %s",
op, scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("%s: %v", op, err)
}
}
// A world whose root holds nothing at all (four ranks over three
// blocks) and whose maximum lives only in the middle block: with
// every root-side candidate absent, the combine must still take the
// middle block's value, never the last block's.
const gn2 = 131073
pyr := make([]float64, gn2)
for i := range pyr {
d := i - gn2/2
if d < 0 {
d = -d
}
pyr[i] = -float64(d)
}
peak := mk(core.FromFloats(pyr, gn2))
want, err := core.Max(peak)
if err != nil {
t.Fatal(err)
}
err = Launch(4, func(w *World) error {
span := mustPartition(t, gn2, w.Size(), w.Rank())
local := narrowSliceFor(t, peak, span)
got, err := w.AllReduceShards(local, span, Max)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("peak: sharded %s against single-array %s",
scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatal(err)
}
// A tie in value with different bits: the earlier block's signed zero
// must win, whichever sign each block holds.
const gn3 = 131073 // two blocks
for _, tc := range []struct {
first, second float64
}{
{0, math.Copysign(0, -1)}, // tie keeps the first block's +0
{math.Copysign(0, -1), 0}, // tie keeps the first block's -0
} {
vals := make([]float64, gn3)
for i := range vals {
vals[i] = tc.first
if i >= gn3/2 {
vals[i] = tc.second
}
}
zeros := mk(core.FromFloats(vals, gn3))
want, err := core.Max(zeros)
if err != nil {
t.Fatal(err)
}
err = Launch(2, func(w *World) error {
span := mustPartition(t, gn3, w.Size(), w.Rank())
local := narrowSliceFor(t, zeros, span)
got, err := w.AllReduceShards(local, span, Max)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("tie: sharded max %s against single-array %s",
scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatal(err)
}
}
}
// TestShardedExtremaOverTCP runs the extrema's candidate flags over
// real connections, the pattern of TestShardedSumOverTCP.
func TestShardedExtremaOverTCP(t *testing.T) {
const gn = 131073
pyr := make([]float64, gn)
for i := range pyr {
d := i - gn/2
if d < 0 {
d = -d
}
pyr[i] = -float64(d)
}
peak := mk(core.FromFloats(pyr, gn))
want, err := core.Max(peak)
if err != nil {
t.Fatal(err)
}
runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, peak, span)
got, err := w.AllReduceShards(local, span, Max)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("rank %d: sharded %s against single-array %s",
w.Rank(), scalarBits(got), scalarBits(want))
}
return nil
})
}
// TestShardedEmptyAxisMatchesTheDtypeRules pins the degenerate length's
// dtype rules against the non-empty path's: a Bool sum counts its trues
// and answers zero, and Any and All keep refusing the non-Bool dtypes.
func TestShardedEmptyAxisMatchesTheDtypeRules(t *testing.T) {
emptyBool := mk(core.FromBools(nil, 0))
if got := core.Sum(emptyBool); got.Int() != 0 {
t.Fatalf("the single-array Bool sum of an empty array: %v", got)
}
err := Launch(3, func(w *World) error {
span := mustPartition(t, 0, w.Size(), w.Rank())
got, err := w.AllReduceShards(emptyBool, span, Sum)
if err != nil {
return err
}
if got.Int() != 0 {
t.Fatalf("the empty Bool sum answered %v", got)
}
if any, err := w.AllReduceShards(emptyBool, span, Any); err != nil || any.Int() != 0 {
t.Fatalf("the empty Bool Any answered %v %v", any, err)
}
if all, err := w.AllReduceShards(emptyBool, span, All); err != nil || all.Int() != 1 {
t.Fatalf("the empty Bool All answered %v %v", all, err)
}
emptyFloat := mk(core.FromFloats(nil, 0))
if _, err := w.AllReduceShards(emptyFloat, span, Any); err == nil {
t.Fatal("the empty Any of a Float shard was accepted")
}
if _, err := w.AllReduceShards(emptyFloat, span, All); err == nil {
t.Fatal("the empty All of a Float shard was accepted")
}
return nil
})
if err != nil {
t.Fatal(err)
}
}
// prodFixture keeps every factor a relative hair away from one, so a
// long product stays finite and meaningful against the single-array
// answer.
func prodFixture(n int) []float64 {
v := make([]float64, n)
for i := range v {
v[i] = 1 + float64(int64(i%2001)-1000)/1e6
}
return v
}
// TestShardedProdIsTheSingleArrayProd: the sharded product carries the
// single-array product's exact bits, the dtype's own per-step rounding
// included.
func TestShardedProdIsTheSingleArrayProd(t *testing.T) {
for _, size := range []int{1, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537, 131073} {
build := func(t *testing.T) map[core.Dtype]*core.Array {
build1 := func(a *core.Array, err error) *core.Array {
if err != nil {
t.Fatal(err)
}
return a
}
vals := prodFixture(gn)
ints := make([]int64, gn)
for i := range ints {
ints[i] = int64(i%5) - 2
}
halves := make([]uint16, gn)
f32s := make([]float32, gn)
for i := range halves {
halves[i] = core.HalfFromFloat64(vals[i])
f32s[i] = float32(vals[i])
}
return map[core.Dtype]*core.Array{
core.Float: build1(core.FromFloats(vals, gn)),
core.Float32: build1(core.FromFloat32s(f32s, gn)),
core.Float16: build1(core.HalvesFromArray(halves, gn)),
core.Int: build1(core.FromInts(ints, gn)),
}
}
for dt, whole := range build(t) {
want, err := core.Prod(whole, 0, false)
if err != nil {
t.Fatal(err)
}
var wantBits string
switch dt {
case core.Int:
wantBits = scalarBits(core.IntScalar(want.RawInts()[0]))
case core.Float16:
wantBits = scalarBits(core.FloatScalar(core.HalfToFloat64(want.RawHalves()[0])))
case core.Float32:
wantBits = scalarBits(core.FloatScalar(float64(want.RawFloat32s()[0])))
default:
wantBits = scalarBits(core.FloatScalar(want.FloatAt(0)))
}
err = Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceShards(local, span, Prod)
if err != nil {
return err
}
if scalarBits(got) != wantBits {
t.Fatalf("size %d gn %d %s: sharded %s against single-array %s",
w.Size(), gn, dt, scalarBits(got), wantBits)
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d %s: %v", size, gn, dt, err)
}
}
}
}
}
// TestShardedNormIsTheSingleArrayNorm pins the power sums and their
// closing against the single-array norm, finite exponents only.
func TestShardedNormIsTheSingleArrayNorm(t *testing.T) {
for _, size := range []int{1, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537, 131073} {
for _, p := range []float64{1, 2, 3.5} {
whole := sparseFixtureOf(gn, core.Float)
want, err := core.Norm(whole, p, 0, false)
if err != nil {
t.Fatal(err)
}
err = Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceNormShards(local, span, p)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(core.FloatScalar(want.FloatAt(0))) {
t.Fatalf("size %d gn %d p %v: sharded %s against single-array %s",
w.Size(), gn, p, scalarBits(got), scalarBits(core.FloatScalar(want.FloatAt(0))))
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d p %v: %v", size, gn, p, err)
}
}
}
}
}
// sparseFixtureOf is sparseFixture landing in one dtype.
func sparseFixtureOf(gn int, dt core.Dtype) *core.Array {
vals := sparseFixture(gn)
var a *core.Array
var err error
switch dt {
case core.Float32:
f32s := make([]float32, gn)
for i := range f32s {
f32s[i] = float32(vals[i])
}
a, err = core.FromFloat32s(f32s, gn)
default:
a, err = core.FromFloats(vals, gn)
}
if err != nil {
panic(err)
}
return a
}
// TestShardedDotIsTheSingleArrayDot: two equally sharded arrays answer
// the single-array Dot's exact bits.
func TestShardedDotIsTheSingleArrayDot(t *testing.T) {
for _, size := range []int{1, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537, 131073} {
x := fixtureDtypes(t, gn)[core.Float]
yWhole := sparseFixtureOf(gn, core.Float)
want, err := core.Dot(x, yWhole)
if err != nil {
t.Fatal(err)
}
err = Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
lx := narrowSliceFor(t, x, span)
ly := narrowSliceFor(t, yWhole, span)
got, err := w.AllReduceDotShards(lx, ly, span)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("size %d gn %d: sharded %s against single-array %s",
w.Size(), gn, scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d: %v", size, gn, err)
}
}
}
}
// TestShardedVectorRefusals: the infinity norm belongs to Max, and the
// vector reductions carry 1-D arrays alone.
func TestShardedVectorRefusals(t *testing.T) {
err := Launch(2, func(w *World) error {
span := mustPartition(t, 100, w.Size(), w.Rank())
local := narrowSliceFor(t, sparseFixtureOf(100, core.Float), span)
if _, err := w.AllReduceNormShards(local, span, math.Inf(1)); err == nil {
t.Fatal("the infinity norm was accepted")
}
two, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
if err != nil {
return err
}
span2 := mustPartition(t, 2, w.Size(), w.Rank())
if _, err := w.AllReduceNormShards(two, span2, 2); err == nil {
t.Fatal("a two-dimensional shard was accepted")
}
if _, err := w.AllReduceShards(two, span2, Prod); err == nil {
t.Fatal("a two-dimensional prod shard was accepted")
}
return nil
})
if err != nil {
t.Fatal(err)
}
}
+237
View File
@@ -0,0 +1,237 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"bufio"
"bytes"
"encoding/binary"
"io"
"net"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/base"
)
// A frame is one message on a link: a 24-byte header, then the payload.
//
// offset 0: tag u8
// offset 1: version u8 (frameVersion, a guard against a confused peer)
// offset 2: reserved u16
// offset 4: from u32, little-endian, the sending rank
// offset 8: dest u32, little-endian, the receiving rank
// offset 12: length u64, little-endian, payload bytes that follow
//
// The header is the only place a length arrives from the wire, and no
// read ever allocates a payload longer than the world's message
// ceiling: a peer announcing more is an error before a byte of payload
// is read. From and dest name the endpoints because a networked world
// is a star over rank 0: a frame from rank 2 to rank 5 rides rank 2's
// connection in, and rank 5's connection out.
const (
frameHeaderLen = 24
frameVersion = 1
)
// The tags the collectives speak. The handshake speaks its own words on
// the fresh connection, before any frame.
const (
tagBarrier uint8 = 1
tagBarrierAck uint8 = 2
tagBroadcast uint8 = 3
tagScatterHead uint8 = 4
tagScatter uint8 = 5
tagGather uint8 = 6
tagShardsValues uint8 = 8
tagShardsWhole uint8 = 9
tagReduce uint8 = 10
tagHalo uint8 = 11
)
// message is one frame in the world's own terms.
type message struct {
tag uint8
from int
dest int
data []byte
// pooled marks a payload whose buffer the hub's pump took from the
// routed frame pool, so the drain, the frame's single consumer,
// returns it there after the write. Every other frame leaves it
// false and its buffer belongs to whoever holds the frame.
pooled bool
}
// peer is one rank's link. An in-process world joins the two ranks'
// channels directly: the sender writes into the receiver's inbox. A
// networked hub pumps each of its links with one reader and one writer;
// a networked rank that is not the hub drives its single connection
// itself and lets the hub route by dest.
type peer struct {
rank int
// inbox carries frames from this rank that this world reads. In
// process it is the direct channel from the peer; at the hub a
// reader feeds it; a non-hub rank leaves it nil and reads its one
// connection itself.
inbox chan message
// outbox carries frames to this rank's connection. In process it
// is the direct channel to the peer; at the hub a writer drains
// it; a non-hub rank leaves it nil and writes its one connection
// itself.
outbox chan message
// gone closes when the peer's world has left; only an in-process
// link has one, and only a sender waits on it.
gone chan struct{}
// conn is this world's end of the peer's networked link.
conn *tcpLink
// pending holds frames that arrived before the collective asked
// for their rank and tag, so no answer ever depends on the
// arrival order. Owned by the one goroutine that drives the
// world's collectives.
pending []message
}
// tcpLink is a framed TCP connection to one peer rank.
type tcpLink struct {
conn net.Conn
rd *bufio.Reader
wr *bufio.Writer
}
func newTCPLink(conn net.Conn) *tcpLink {
return &tcpLink{conn: conn, rd: bufio.NewReader(conn), wr: bufio.NewWriter(conn)}
}
func (l *tcpLink) readFrame(max int64, deadline time.Time) (message, error) {
return l.readFrameInto(max, deadline, nil)
}
// readFrameInto reads one frame as readFrame does. Take, when not
// nil, is asked with the parsed header and the payload length the
// header announced for the buffer the payload reads into; a nil
// answer allocates the payload as usual. A buffer take supplied
// leaves the read with the message's pooled mark on, which commits it
// to the single ownership chain the routed frame pool lives on: the
// pump that reads the frame hands it through one outbox channel to
// the one drain, which returns the buffer after the write.
func (l *tcpLink) readFrameInto(max int64, deadline time.Time, take func(message, int64) []byte) (message, error) {
if err := l.conn.SetReadDeadline(deadline); err != nil {
return message{}, err
}
var head [frameHeaderLen]byte
if _, err := io.ReadFull(l.rd, head[:]); err != nil {
return message{}, err
}
if head[1] != frameVersion {
return message{}, base.Errf("spmd: frame version %d from the link is not %d", head[1], frameVersion)
}
m := message{
tag: head[0],
from: int(binary.LittleEndian.Uint32(head[4:])),
dest: int(binary.LittleEndian.Uint32(head[8:])),
}
length := int64(binary.LittleEndian.Uint64(head[12:]))
if length < 0 || length > max {
return message{}, base.Errf("spmd: the link announces a %d byte payload beyond the %d byte ceiling", length, max)
}
if take != nil {
if buf := take(m, length); buf != nil {
m.data, m.pooled = buf, true
}
}
if m.data == nil {
m.data = make([]byte, length)
}
if _, err := io.ReadFull(l.rd, m.data); err != nil {
return message{}, err
}
return m, nil
}
func (l *tcpLink) writeFrame(m message, deadline time.Time) error {
if err := l.conn.SetWriteDeadline(deadline); err != nil {
return err
}
var head [frameHeaderLen]byte
head[0] = m.tag
head[1] = frameVersion
binary.LittleEndian.PutUint32(head[4:], uint32(m.from))
binary.LittleEndian.PutUint32(head[8:], uint32(m.dest))
binary.LittleEndian.PutUint64(head[12:], uint64(len(m.data)))
if _, err := l.wr.Write(head[:]); err != nil {
return err
}
if _, err := l.wr.Write(m.data); err != nil {
return err
}
return l.wr.Flush()
}
// The handshake words on a fresh TCP connection: the joining rank
// sends hello, the listening rank answers with the world's size and
// the joining rank's place in it.
const handshakeLen = 16
var handshakeMagic = [4]byte{'T', 'S', 'P', 'M'}
// sendHello is the joining side's word.
func sendHello(conn net.Conn) error {
var hello [handshakeLen]byte
copy(hello[0:4], handshakeMagic[:])
hello[4] = frameVersion
_, err := conn.Write(hello[:])
return err
}
// readHello is the listening side's read of it.
func readHello(conn net.Conn, deadline time.Time) error {
if err := conn.SetReadDeadline(deadline); err != nil {
return err
}
var hello [handshakeLen]byte
if _, err := io.ReadFull(conn, hello[:]); err != nil {
return err
}
if !bytes.Equal(hello[0:4], handshakeMagic[:]) {
return base.Errf("spmd: the joining connection did not say the spmd magic")
}
if hello[4] != frameVersion {
return base.Errf("spmd: joining protocol version %d is not %d", hello[4], frameVersion)
}
return nil
}
// sendWelcome is the listening side's answer: the world size and the
// rank the joining connection carries.
func sendWelcome(conn net.Conn, size, rank int) error {
var welcome [handshakeLen]byte
copy(welcome[0:4], handshakeMagic[:])
welcome[4] = frameVersion
binary.LittleEndian.PutUint32(welcome[8:], uint32(size))
binary.LittleEndian.PutUint32(welcome[12:], uint32(rank))
_, err := conn.Write(welcome[:])
return err
}
// readWelcome is the joining side's read of it.
func readWelcome(conn net.Conn, deadline time.Time) (size, rank int, err error) {
if err := conn.SetReadDeadline(deadline); err != nil {
return 0, 0, err
}
var welcome [handshakeLen]byte
if _, err := io.ReadFull(conn, welcome[:]); err != nil {
return 0, 0, err
}
if !bytes.Equal(welcome[0:4], handshakeMagic[:]) {
return 0, 0, base.Errf("spmd: the listener did not answer with the spmd magic")
}
if welcome[4] != frameVersion {
return 0, 0, base.Errf("spmd: listener protocol version %d is not %d", welcome[4], frameVersion)
}
size = int(binary.LittleEndian.Uint32(welcome[8:]))
rank = int(binary.LittleEndian.Uint32(welcome[12:]))
if size < 1 || rank < 0 || rank >= size {
return 0, 0, base.Errf("spmd: the listener answered with an impossible size %d and rank %d", size, rank)
}
return size, rank, nil
}
+315
View File
@@ -0,0 +1,315 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"encoding/binary"
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The wire form of an array is a self-describing frame payload: one
// dtype byte, one dimension-count byte, the extents as little-endian
// int64, then the elements as little-endian fixed-width raw bits in
// row-major order. Every multi-byte field is written little-endian by
// explicit conversion, never by copying the payload's memory, so the
// form is the same on every architecture the library builds for.
// Floating-point payloads keep their exact bit patterns, NaN payloads
// and signed zeros included.
// dtypeWidth is the wire width of one element per dtype.
func dtypeWidth(dt core.Dtype) int {
switch dt {
case core.Int8, core.Uint8, core.Bool:
return 1
case core.Int16, core.Uint16, core.Float16:
return 2
case core.Int32, core.Uint32, core.Float32:
return 4
case core.Int, core.Float:
return 8
case core.Complex:
return 16
}
return 0
}
// encodeArray appends the wire form of a to dst and returns the grown
// slice.
func encodeArray(dst []byte, a *core.Array) ([]byte, error) {
return encodePart(dst, a, a.Shape(), 0, a.Len())
}
// encodeHead appends just the head of a wire form: dtype, dimension
// count and extents. A rank that holds the head can name its piece of
// the data before any payload flows.
func encodeHead(dst []byte, dt core.Dtype, shape []int) ([]byte, error) {
if dtypeWidth(dt) == 0 {
return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt)
}
if len(shape) > 255 {
return nil, base.Errf("spmd: an array of %d dimensions exceeds the wire form's 255", len(shape))
}
dst = append(dst, byte(dt), byte(len(shape)))
for _, d := range shape {
if d < 0 {
return nil, base.Errf("spmd: a wire form cannot name a negative extent %d", d)
}
dst = binary.LittleEndian.AppendUint64(dst, uint64(d))
}
return dst, nil
}
// decodeHead reads a wire form's head: the dtype and the shape. It
// validates only the head's own length; the payload against the shape
// is decodeArray's check.
func decodeHead(wire []byte) (core.Dtype, []int, error) {
if len(wire) < 2 {
return 0, nil, base.Errf("spmd: a wire form needs at least two bytes, got %d", len(wire))
}
dt := core.Dtype(wire[0])
if dtypeWidth(dt) == 0 {
return 0, nil, base.Errf("spmd: dtype %s cannot travel the wire", dt)
}
ndim := int(wire[1])
if len(wire) < 2+8*ndim {
return 0, nil, base.Errf("spmd: a wire form for %d dimensions is short by %d bytes", ndim, 2+8*ndim-len(wire))
}
shape := make([]int, ndim)
for i := range shape {
v := binary.LittleEndian.Uint64(wire[2+8*i:])
if v > math.MaxInt {
return 0, nil, base.Errf("spmd: an extent of %d does not fit this machine's int", v)
}
shape[i] = int(v)
}
return dt, shape, nil
}
// decodeWire builds an array from its complete wire form, head and
// payload together. Every extent is bounded before anything is
// allocated: an overflowing shape or a payload that does not carry
// exactly the named elements is an error, never a partial answer.
func decodeWire(wire []byte) (*core.Array, error) {
dt, shape, err := decodeHead(wire)
if err != nil {
return nil, err
}
return decodeArray(dt, shape, wire[2+8*len(shape):])
}
// partHeadLen is the byte length of a one-dimensional part's wire
// head, the shape encodePart writes for `[]int{count}`: one dtype
// byte, one dimension-count byte, and the single uint64 extent.
const partHeadLen = 2 + 8
// encodePart appends the wire form of a contiguous element range of a,
// presented under the given shape: the head names the shape, the
// payload carries elements [first, first+count) of a. Only the array's
// own elements take part: a rebased view's payload may run past its
// element count, so every payload walk stops at Len.
func encodePart(dst []byte, a *core.Array, shape []int, first, count int) ([]byte, error) {
dt := a.Dtype()
width := dtypeWidth(dt)
if width == 0 {
return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt)
}
if len(shape) > 255 {
return nil, base.Errf("spmd: an array of %d dimensions exceeds the wire form's 255", len(shape))
}
if first < 0 || count < 0 || first+count > a.Len() {
return nil, base.Errf("spmd: element range [%d, %d) is outside the array's %d elements", first, first+count, a.Len())
}
for _, d := range shape {
if d < 0 {
return nil, base.Errf("spmd: a wire form cannot name a negative extent %d", d)
}
}
dst = append(dst, byte(dt), byte(len(shape)))
for _, d := range shape {
dst = binary.LittleEndian.AppendUint64(dst, uint64(d))
}
start := len(dst)
switch dt {
case core.Int:
for _, v := range a.RawInts()[first : first+count] {
dst = binary.LittleEndian.AppendUint64(dst, uint64(v))
}
case core.Float:
for _, v := range a.RawFloats()[first : first+count] {
dst = binary.LittleEndian.AppendUint64(dst, math.Float64bits(v))
}
case core.Float32:
for _, v := range a.RawFloat32s()[first : first+count] {
dst = binary.LittleEndian.AppendUint32(dst, math.Float32bits(v))
}
case core.Float16:
for _, v := range a.RawHalves()[first : first+count] {
dst = binary.LittleEndian.AppendUint16(dst, v)
}
case core.Complex:
for _, v := range a.RawComplexes()[first : first+count] {
dst = binary.LittleEndian.AppendUint64(dst, math.Float64bits(real(v)))
dst = binary.LittleEndian.AppendUint64(dst, math.Float64bits(imag(v)))
}
case core.Bool:
for _, v := range a.RawBools()[first : first+count] {
b := byte(0)
if v {
b = 1
}
dst = append(dst, b)
}
case core.Int8:
for _, v := range a.RawInt8s()[first : first+count] {
dst = append(dst, byte(v))
}
case core.Uint8:
for _, v := range a.RawUint8s()[first : first+count] {
dst = append(dst, v)
}
case core.Int16:
for _, v := range a.RawInt16s()[first : first+count] {
dst = binary.LittleEndian.AppendUint16(dst, uint16(v))
}
case core.Uint16:
for _, v := range a.RawUint16s()[first : first+count] {
dst = binary.LittleEndian.AppendUint16(dst, v)
}
case core.Int32:
for _, v := range a.RawInt32s()[first : first+count] {
dst = binary.LittleEndian.AppendUint32(dst, uint32(v))
}
case core.Uint32:
for _, v := range a.RawUint32s()[first : first+count] {
dst = binary.LittleEndian.AppendUint32(dst, v)
}
}
if got := len(dst) - start; got != count*width {
return nil, base.Errf("spmd: array of dtype %s encoded %d bytes for %d elements", dt, got, count)
}
return dst, nil
}
// decodeArray builds an array from the wire payload of a dtype and
// shape the caller has read. The payload must carry exactly the
// elements the shape names at the dtype's width; the caller has
// already bounded payload by the world's message ceiling.
func decodeArray(dt core.Dtype, shape []int, payload []byte) (*core.Array, error) {
width := dtypeWidth(dt)
if width == 0 {
return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt)
}
n, ok := elementCount(shape)
if !ok {
return nil, base.Errf("spmd: shape %v overflows the element count", shape)
}
if int64(len(payload)) != int64(n)*int64(width) {
return nil, base.Errf("spmd: wire payload of %d bytes does not carry %d elements of %s", len(payload), n, dt)
}
switch dt {
case core.Int:
vals := make([]int64, n)
for i := range vals {
vals[i] = int64(binary.LittleEndian.Uint64(payload[i*8:]))
}
return core.FromInts(vals, shape...)
case core.Float:
vals := make([]float64, n)
for i := range vals {
vals[i] = math.Float64frombits(binary.LittleEndian.Uint64(payload[i*8:]))
}
return core.FromFloats(vals, shape...)
case core.Float32:
vals := make([]float32, n)
for i := range vals {
vals[i] = math.Float32frombits(binary.LittleEndian.Uint32(payload[i*4:]))
}
return core.FromFloat32s(vals, shape...)
case core.Float16:
vals := make([]uint16, n)
for i := range vals {
vals[i] = binary.LittleEndian.Uint16(payload[i*2:])
}
return core.HalvesFromArray(vals, shape...)
case core.Complex:
vals := make([]complex128, n)
for i := range vals {
re := math.Float64frombits(binary.LittleEndian.Uint64(payload[i*16:]))
im := math.Float64frombits(binary.LittleEndian.Uint64(payload[i*16+8:]))
vals[i] = complex(re, im)
}
return core.FromComplexes(vals, shape...)
case core.Bool:
vals := make([]bool, n)
for i := range vals {
switch payload[i] {
case 0:
case 1:
vals[i] = true
default:
return nil, base.Errf("spmd: bool wire byte %d at index %d is neither 0 nor 1", payload[i], i)
}
}
return core.FromBools(vals, shape...)
case core.Int8:
vals := make([]int8, n)
for i := range vals {
vals[i] = int8(payload[i])
}
return core.FromInt8s(vals, shape...)
case core.Uint8:
vals := make([]uint8, n)
for i := range vals {
vals[i] = payload[i]
}
return core.FromUint8s(vals, shape...)
case core.Int16:
vals := make([]int16, n)
for i := range vals {
vals[i] = int16(binary.LittleEndian.Uint16(payload[i*2:]))
}
return core.FromInt16s(vals, shape...)
case core.Uint16:
vals := make([]uint16, n)
for i := range vals {
vals[i] = binary.LittleEndian.Uint16(payload[i*2:])
}
return core.FromUint16s(vals, shape...)
case core.Int32:
vals := make([]int32, n)
for i := range vals {
vals[i] = int32(binary.LittleEndian.Uint32(payload[i*4:]))
}
return core.FromInt32s(vals, shape...)
case core.Uint32:
vals := make([]uint32, n)
for i := range vals {
vals[i] = binary.LittleEndian.Uint32(payload[i*4:])
}
return core.FromUint32s(vals, shape...)
}
return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt)
}
// elementCount is the product of the shape, reported as false when any
// extent is negative or the product overflows an int.
func elementCount(shape []int) (int, bool) {
n := 1
for _, d := range shape {
if d < 0 {
return 0, false
}
if d == 0 {
return 0, true
}
if n > math.MaxInt/d {
return 0, false
}
n *= d
}
return n, true
}
+175
View File
@@ -0,0 +1,175 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"bytes"
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The wire form is judged by one rule: the array that comes out carries
// the exact bits of the array that went in, shape and dtype included.
func TestWireRoundTrip(t *testing.T) {
nanLow := math.Float64frombits(0x7ff8000000000001) // NaN, low payload bit set
nanHigh := math.Float64frombits(0xfff8deadbeef0000)
negZero := math.Copysign(0, -1)
cases := []*core.Array{
mk(core.FromFloats([]float64{1, -2.5, 0, negZero, math.Inf(1), math.Inf(-1), nanLow, nanHigh, math.MaxFloat64, math.SmallestNonzeroFloat64}, 10)),
mk(core.FromFloat32s([]float32{1.5, -0.25, float32(negZero), float32(math.Inf(1)), float32(nanLow)}, 5)),
mk(core.FromComplexes([]complex128{1 + 2i, complex(negZero, nanHigh), complex(0, negZero)}, 3)),
mk(core.FromInts([]int64{math.MaxInt64, math.MinInt64, -1, 0, 42}, 5)),
mk(core.FromBools([]bool{true, false, true, true}, 4)),
mk(core.FromInt8s([]int8{math.MinInt8, math.MaxInt8, -1, 0}, 4)),
mk(core.FromUint8s([]uint8{0, 255, 128}, 3)),
mk(core.FromInt16s([]int16{math.MinInt16, math.MaxInt16, -1}, 3)),
mk(core.FromUint16s([]uint16{0, 65535, 32768}, 3)),
mk(core.FromInt32s([]int32{math.MinInt32, math.MaxInt32, -1}, 3)),
mk(core.FromUint32s([]uint32{0, 4294967295, 2147483648}, 3)),
// Float16 raw halves: denormals, infinities, NaN payloads, the
// whole span the payload can hold.
mk(core.HalvesFromArray([]uint16{0x0001, 0x03ff, 0x7bff, 0x7c00, 0xfc00, 0x7e00, 0x7eaa}, 7)),
// Empty and multi-dimensional shapes.
mk(core.FromFloats(nil, 0)),
mk(core.FromInts([]int64{1, 2, 3, 4, 5, 6}, 2, 3)),
mk(core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12}, 2, 1, 2, 3)),
}
for i, a := range cases {
wire, err := encodeArray(nil, a)
if err != nil {
t.Fatalf("case %d (%s): %v", i, a.Dtype(), err)
}
b, err := decodeWire(wire)
if err != nil {
t.Fatalf("case %d (%s): %v", i, a.Dtype(), err)
}
if b.Dtype() != a.Dtype() {
t.Fatalf("case %d: dtype %s came back as %s", i, a.Dtype(), b.Dtype())
}
if !sameShape(a.Shape(), b.Shape()) {
t.Fatalf("case %d: shape %v came back as %v", i, a.Shape(), b.Shape())
}
if !sameBits(a, b) {
t.Fatalf("case %d (%s): the round trip moved bits", i, a.Dtype())
}
}
}
// TestWireRefusesTheHostile is the hostile-reader rule: a form that
// does not carry exactly what its head names is an error, never a short
// read and never a partial answer.
func TestWireRefusesTheHostile(t *testing.T) {
a := mk(core.FromFloats([]float64{1, 2, 3, 4}, 2, 2))
wire, err := encodeArray(nil, a)
if err != nil {
t.Fatal(err)
}
payload := wire[2+8*2:]
if _, err := decodeWire(wire[:len(wire)-1]); err == nil {
t.Fatal("a short payload decoded")
}
if _, err := decodeWire(append(bytes.Clone(wire), 0)); err == nil {
t.Fatal("a long payload decoded")
}
if _, err := decodeWire(wire[:1]); err == nil {
t.Fatal("a headless form decoded")
}
if _, err := decodeWire(wire[:6]); err == nil {
t.Fatal("a truncated head decoded")
}
big := []int{math.MaxInt32, math.MaxInt32}
if _, err := decodeArray(core.Float, big, nil); err == nil {
t.Fatal("an overflowing shape decoded")
}
if _, err := decodeArray(core.Float, []int{2, -2}, payload); err == nil {
t.Fatal("a negative extent decoded")
}
badBool := mk(core.FromBools([]bool{true, false}, 2))
bw, err := encodeArray(nil, badBool)
if err != nil {
t.Fatal(err)
}
bp := bytes.Clone(bw)
bp[2+8] = 2
if _, err := decodeWire(bp); err == nil {
t.Fatal("a bool byte of 2 decoded")
}
}
// mk builds a fixture or panics; the fixtures are package-level
// literals, so a panic lands in the test that declared them.
func mk(a *core.Array, err error) *core.Array {
if err != nil {
panic(err)
}
return a
}
// sameBits compares two arrays' raw payloads bit for bit, dtype by
// dtype. NaN payloads and signed zeros are part of the contract.
func sameBits(a, b *core.Array) bool {
if a.Len() != b.Len() || a.Dtype() != b.Dtype() {
return false
}
switch a.Dtype() {
case core.Int:
return slicesEqual(a.RawInts()[:a.Len()], b.RawInts()[:b.Len()])
case core.Float:
x, y := a.RawFloats()[:a.Len()], b.RawFloats()[:b.Len()]
for i := range x {
if math.Float64bits(x[i]) != math.Float64bits(y[i]) {
return false
}
}
return true
case core.Float32:
x, y := a.RawFloat32s()[:a.Len()], b.RawFloat32s()[:b.Len()]
for i := range x {
if math.Float32bits(x[i]) != math.Float32bits(y[i]) {
return false
}
}
return true
case core.Float16:
return slicesEqual(a.RawHalves()[:a.Len()], b.RawHalves()[:b.Len()])
case core.Complex:
x, y := a.RawComplexes()[:a.Len()], b.RawComplexes()[:b.Len()]
for i := range x {
if math.Float64bits(real(x[i])) != math.Float64bits(real(y[i])) ||
math.Float64bits(imag(x[i])) != math.Float64bits(imag(y[i])) {
return false
}
}
return true
case core.Bool:
return slicesEqual(a.RawBools()[:a.Len()], b.RawBools()[:b.Len()])
case core.Int8:
return slicesEqual(a.RawInt8s()[:a.Len()], b.RawInt8s()[:b.Len()])
case core.Uint8:
return slicesEqual(a.RawUint8s()[:a.Len()], b.RawUint8s()[:b.Len()])
case core.Int16:
return slicesEqual(a.RawInt16s()[:a.Len()], b.RawInt16s()[:b.Len()])
case core.Uint16:
return slicesEqual(a.RawUint16s()[:a.Len()], b.RawUint16s()[:b.Len()])
case core.Int32:
return slicesEqual(a.RawInt32s()[:a.Len()], b.RawInt32s()[:b.Len()])
case core.Uint32:
return slicesEqual(a.RawUint32s()[:a.Len()], b.RawUint32s()[:b.Len()])
}
return false
}
func slicesEqual[T comparable](a, b []T) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
+576
View File
@@ -0,0 +1,576 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"errors"
"io"
"net"
"sync"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/base"
)
// Default bounds of a networked world: one collective's wait, one
// frame's payload, and the number of frames a link may queue. All
// three exist so that a stuck or hostile peer is an error, never a
// hang and never an allocation.
const (
defaultTimeout = 10 * time.Minute
defaultMaxMessage = int64(16) << 30
linkQueue = 64
// pendingCap bounds how many unmatched frames one recvFrom may
// hold: a peer flooding frames nobody asked for is an error,
// never an unbounded allocator.
pendingCap = 4096
)
// Options bounds a networked world. An in-process world from Launch
// takes none of them: its links are channels, and progress is the
// program's own business, as it is in MPI.
type Options struct {
// Timeout bounds one collective's wait on the network: dialing,
// the handshake and every send and receive carry it as a
// deadline, refreshed each time a frame moves. Zero means the
// default of ten minutes; a negative value means no deadline at
// all.
Timeout time.Duration
// MaxMessage is the largest frame payload the world accepts, in
// bytes. A peer announcing more is refused before any allocation.
// Zero means the default of 16 GiB.
MaxMessage int64
}
func (o Options) timeout() time.Duration {
switch {
case o.Timeout > 0:
return o.Timeout
case o.Timeout < 0:
return 0 // no deadline
default:
return defaultTimeout
}
}
func (o Options) maxMessage() int64 {
if o.MaxMessage > 0 {
return o.MaxMessage
}
return defaultMaxMessage
}
// World is one rank's end of an SPMD world: its place in it, the links
// to the other ranks and the collectives. A World is driven by one
// goroutine: like an MPI rank, it never runs two collectives at once.
type World struct {
rank int
size int
timeout time.Duration
maxMessage int64
networked bool
peers []*peer // peers[r] is the link to rank r; nil for this rank
closers []io.Closer
// The hub's readers and writers. The writers are waited on before
// an orderly close, so the last collective's frames are on the
// wire before the connections go away; the readers are waited on
// after, once those connections have broken their blocking reads.
drainWg sync.WaitGroup
pumpWg sync.WaitGroup
done chan struct{} // closed when the world failed or left
closeDone sync.Once
failErr error
}
// Rank returns this rank's index, from 0 to Size-1.
func (w *World) Rank() int { return w.rank }
// Size returns the number of ranks in the world.
func (w *World) Size() int { return w.size }
// Close ends a networked world: whatever the collectives queued is
// written to the wire, then the connections close and the other ranks
// see the departure as their next receive failing. On an in-process
// world it is a no-op, because Launch tears the world down.
func (w *World) Close() error {
w.leave()
w.drainWg.Wait()
var first error
for _, c := range w.closers {
if err := c.Close(); err != nil && first == nil {
first = err
}
}
w.pumpWg.Wait()
return first
}
// leave closes done without recording a failure: the ordinary exit of
// the rank's program.
func (w *World) leave() {
w.closeDone.Do(func() { close(w.done) })
}
// fail records the world's first failure, closes done so that every
// wait on this world wakes, and closes the networked links. It returns
// the failure an outside caller should see.
func (w *World) fail(err error) error {
w.closeDone.Do(func() {
w.failErr = err
close(w.done)
for _, c := range w.closers {
c.Close()
}
})
return w.status()
}
// status is the entry check every public operation makes: a world that
// failed, or a rank whose program already returned, answers with an
// error and never with data.
func (w *World) status() error {
select {
case <-w.done:
if w.failErr != nil {
return base.Errf("spmd: rank %d world is in a failed state: %v", w.rank, w.failErr)
}
return base.Errf("spmd: rank %d world has left", w.rank)
default:
return nil
}
}
// deadline is the absolute time one collective may run to on a
// networked world, refreshed each time a frame moves; the zero time
// means no deadline, which is what an in-process world always gets.
func (w *World) deadline() time.Time {
if !w.networked || w.timeout <= 0 {
return time.Time{}
}
return time.Now().Add(w.timeout)
}
// depart closes the peer's gone channel: the peer closed its side of
// the link, so no further frame will ever come from it.
func (p *peer) depart() {
select {
case <-p.gone:
default:
close(p.gone)
}
}
// hasLeft reports whether the peer closed its side of the link.
func (p *peer) hasLeft() bool {
if p.gone == nil {
return false
}
select {
case <-p.gone:
return true
default:
return false
}
}
// sendTo carries one frame to rank r. On an in-process world it walks
// the direct channel; on a networked world it walks this rank's own
// link to the hub, because rank 0 routes every frame by its
// destination.
func (w *World) sendTo(r int, tag uint8, payload []byte) error {
if err := w.status(); err != nil {
return err
}
m := message{tag: tag, from: w.rank, dest: r, data: payload}
p := w.peers[r]
if p != nil && p.outbox != nil {
select {
case p.outbox <- m:
// On a buffered networked link a queued frame is not yet
// a delivered frame, so a peer that left between the two
// dropped it. An in-process handoff is the delivery
// itself: the receiver taking the frame and then leaving
// is its own healthy business.
if w.networked && p.hasLeft() {
return w.fail(base.Errf("spmd: rank %d sending to rank %d: rank %d left the world", w.rank, r, r))
}
return nil
case <-p.gone:
return w.fail(base.Errf("spmd: rank %d sending to rank %d: rank %d left the world", w.rank, r, r))
case <-w.done:
return w.status()
}
}
// A networked rank that is not the hub owns one connection, and
// every word it sends rides it; the hub reads the destination.
if err := w.peers[0].conn.writeFrame(m, w.deadline()); err != nil {
return w.fail(base.Errf("spmd: rank %d sending to rank %d: %v", w.rank, r, err))
}
return nil
}
// recvFrom returns the payload of the next frame rank r sent under the
// wanted tag, holding frames that arrived earlier for other ranks or
// tags until the collective asks for them: no answer ever depends on
// the arrival order. Any failure fails the world.
func (w *World) recvFrom(r int, tag uint8) ([]byte, error) {
if err := w.status(); err != nil {
return nil, err
}
p := w.peers[r]
if w.networked && w.rank != 0 {
// One stream carries every rank's words to a rank that is not
// the hub, so one pending stash serves them all.
p = w.peers[0]
}
for i, m := range p.pending {
if m.from == r && m.tag == tag {
p.pending = append(p.pending[:i], p.pending[i+1:]...)
return m.data, nil
}
}
for {
var m message
var err error
switch {
case p.inbox != nil:
select {
case m = <-p.inbox:
case <-p.gone:
// The peer left: whatever it queued before leaving is
// still in the buffer and still counts; only an empty
// buffer means the frames will never come.
for {
select {
case m = <-p.inbox:
if m.from == r && m.tag == tag {
return m.data, nil
}
p.pending = append(p.pending, m)
continue
default:
}
return nil, w.fail(base.Errf("spmd: rank %d receiving from rank %d: rank %d left the world", w.rank, r, r))
}
case <-w.done:
return nil, w.status()
}
default:
// This rank drives its single connection; the hub has
// already routed whatever was not for it.
m, err = p.conn.readFrame(w.maxMessage, w.deadline())
if err != nil {
return nil, w.fail(base.Errf("spmd: rank %d receiving from rank %d: %v", w.rank, r, w.readErr(err)))
}
if m.dest != w.rank {
return nil, w.fail(base.Errf("spmd: rank %d got a frame addressed to rank %d", w.rank, m.dest))
}
}
if m.from == r && m.tag == tag {
return m.data, nil
}
p.pending = append(p.pending, m)
if len(p.pending) >= pendingCap {
return nil, w.fail(base.Errf("spmd: rank %d holds %d unmatched frames against rank %d", w.rank, len(p.pending), r))
}
}
}
// readErr names a networked read failure for what it is: a peer that
// closed or dropped its connection.
func (w *World) readErr(err error) error {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
return errors.New("the peer closed its connection")
}
return err
}
// pump reads one networked hub link for the world's lifetime: frames
// for rank 0 join the link's inbox, frames for anybody else join that
// rank's outbox unchanged. The hub is a post office, never an
// interpreter: what a frame carries is the collectives' business. A
// frame bound for another rank reads into a buffer from the routed
// frame pool, whose ownership rides the message through the outbox
// channel to the link's one drain. A peer that closes its connection
// has left the world, which is its own business too; only a broken or
// unreadable link fails this world.
func (w *World) pump(p *peer) {
for {
m, err := p.conn.readFrameInto(w.maxMessage, w.deadline(), takeFrameBuffer)
if err != nil {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) || errors.Is(err, net.ErrClosed) {
p.depart()
return
}
w.fail(base.Errf("spmd: rank 0 reading from rank %d: %v", p.rank, w.readErr(err)))
return
}
if m.from != p.rank || m.dest < 0 || m.dest >= w.size {
w.fail(base.Errf("spmd: rank 0 got a mislabelled frame from rank %d", p.rank))
return
}
var out chan message
if m.dest == 0 {
out = p.inbox
} else {
out = w.peers[m.dest].outbox
}
select {
case out <- m:
case <-w.done:
return
}
}
}
// drain writes one networked hub link for the world's lifetime: it is
// the only goroutine that ever writes the connection, so the frames
// the collective logic and the routed traffic send share one ordered
// stream without a lock. It is also the single consumer of the
// routed frames the pump queues, which makes it the one place a
// pooled payload is returned: writeFrame is the payload's last
// reader, so the buffer goes back the moment the write returns,
// whatever the answer was. When the world ends it writes out whatever
// the last collectives queued before it leaves, so an orderly close
// never drops a frame that was sent. A write against a link whose peer
// already left is not failed, because the departure was the peer's own
// clean act; any other write error is, named at once rather than left
// to surface later as a hang.
func (w *World) drain(p *peer) {
write := func(m message) bool {
err := p.conn.writeFrame(m, w.deadline())
if m.pooled {
routedFrames.retire(m.data)
}
if err == nil {
return true
}
if p.hasLeft() || w.ended() {
return false
}
w.fail(base.Errf("spmd: rank 0 writing to rank %d: %v", p.rank, err))
return false
}
for {
select {
case m := <-p.outbox:
if !write(m) {
return
}
case <-w.done:
for {
select {
case m := <-p.outbox:
if !write(m) {
return
}
default:
return
}
}
}
}
}
// ended reports whether the world's done channel has closed.
func (w *World) ended() bool {
select {
case <-w.done:
return true
default:
return false
}
}
// startPumps launches the hub's readers and writers, one of each per
// link. They live as long as the world does.
func (w *World) startPumps() {
for r := 1; r < w.size; r++ {
p := w.peers[r]
w.pumpWg.Go(func() { w.pump(p) })
w.drainWg.Go(func() { w.drain(p) })
}
}
// Launch runs the same function on size ranks of one process, one
// goroutine per rank, over in-process links: the same collectives, the
// same answers and the same rules as a networked world, which makes it
// the development and test surface of the package. The errors of the
// ranks that failed come back joined in rank order, so the report is
// deterministic and the rank that caused the trouble is in it; a rank
// whose function panics fails its world, and the panic comes back as
// that rank's error.
func Launch(size int, fn func(w *World) error) error {
if size < 1 {
return base.Errf("spmd: a world needs at least one rank, got %d", size)
}
worlds := make([]*World, size)
for r := range worlds {
worlds[r] = &World{
rank: r,
size: size,
done: make(chan struct{}),
peers: make([]*peer, size),
}
}
// One channel per direction of every pair: the sender's outbox is
// the receiver's inbox, so a frame crosses without a middleman and
// a receive wakes the moment its rank's world leaves.
for r := range worlds {
for q := r + 1; q < size; q++ {
rToQ := make(chan message, linkQueue)
qToR := make(chan message, linkQueue)
worlds[r].peers[q] = &peer{rank: q, inbox: qToR, outbox: rToQ, gone: worlds[q].done}
worlds[q].peers[r] = &peer{rank: r, inbox: rToQ, outbox: qToR, gone: worlds[r].done}
}
}
errs := make([]error, size)
var wg sync.WaitGroup
for r := range worlds {
wg.Go(func() {
defer func() {
if p := recover(); p != nil {
errs[r] = base.Errf("spmd: rank %d panicked: %v", r, p)
worlds[r].fail(errs[r])
}
worlds[r].leave()
}()
errs[r] = fn(worlds[r])
})
}
wg.Wait()
return errors.Join(errs...)
}
// Listen builds the rank 0 end of a networked world: it listens on the
// address until every other rank has joined, assigning ranks in dial
// order. Connections that do not say the spmd handshake are closed
// and skipped, so they cannot take a rank's place.
func Listen(addr string, size int, opts Options) (*World, error) {
ln, err := net.Listen("tcp", addr)
if err != nil {
return nil, base.Errf("spmd: listening on %s: %v", addr, err)
}
return listen(ln, size, opts)
}
// listen assembles the rank 0 world on a ready listener; the tests use
// it to hand over a listener whose address they already know.
func listen(ln net.Listener, size int, opts Options) (*World, error) {
if size < 1 {
ln.Close()
return nil, base.Errf("spmd: a world needs at least one rank, got %d", size)
}
timeout := opts.timeout()
w := &World{
rank: 0,
size: size,
timeout: timeout,
maxMessage: opts.maxMessage(),
networked: true,
done: make(chan struct{}),
peers: make([]*peer, size),
}
w.closers = append(w.closers, ln)
if tcp, ok := ln.(*net.TCPListener); ok && timeout > 0 {
tcp.SetDeadline(time.Now().Add(timeout))
}
deadline := w.deadline()
for joined := 1; joined < size; joined++ {
conn, err := ln.Accept()
if err != nil {
return nil, w.fail(base.Errf("spmd: rank 0 accepting rank %d on %s: %v", joined, ln.Addr(), err))
}
if err := readHello(conn, deadline); err != nil {
// A connection that does not say the handshake is not one
// of ours; it takes no rank's place.
conn.Close()
joined--
continue
}
if err := sendWelcome(conn, size, joined); err != nil {
conn.Close()
return nil, w.fail(base.Errf("spmd: rank 0 welcoming rank %d: %v", joined, err))
}
w.peers[joined] = &peer{
rank: joined,
conn: newTCPLink(conn),
inbox: make(chan message, linkQueue),
outbox: make(chan message, linkQueue),
gone: make(chan struct{}),
}
w.closers = append(w.closers, conn)
}
// Every rank has its link; nobody else joins this world. The
// listener was closers[0], and its job is done.
ln.Close()
w.closers = w.closers[1:]
w.startPumps()
return w, nil
}
// Join builds the other ranks' end of a networked world: it dials the
// listening rank 0, which answers with the world's size and this
// connection's rank. The world is a star over rank 0's listener, so
// one address is the whole world's knowledge; rank 0 routes every
// frame to its destination.
func Join(addr string, opts Options) (*World, error) {
timeout := opts.timeout()
d := net.Dialer{Timeout: timeout}
conn, err := d.Dial("tcp", addr)
if err != nil {
return nil, base.Errf("spmd: rank dialling %s: %v", addr, err)
}
w := &World{
timeout: timeout,
maxMessage: opts.maxMessage(),
networked: true,
done: make(chan struct{}),
closers: []io.Closer{conn},
}
if err := sendHello(conn); err != nil {
return nil, w.fail(base.Errf("spmd: rank saying hello to %s: %v", addr, err))
}
size, rank, err := readWelcome(conn, w.deadline())
if err != nil {
return nil, w.fail(base.Errf("spmd: rank reading %s's welcome: %v", addr, err))
}
if rank == 0 {
return nil, w.fail(base.Errf("spmd: the listener at %s answered a joining connection with rank 0", addr))
}
w.rank, w.size = rank, size
w.peers = make([]*peer, size)
w.peers[0] = &peer{rank: 0, conn: newTCPLink(conn)}
return w, nil
}
// Barrier blocks until every rank of the world has reached it. It
// carries no data and no arithmetic, so there is nothing in it to be
// anything but deterministic.
func (w *World) Barrier() error {
if err := w.status(); err != nil {
return err
}
if w.rank == 0 {
for r := 1; r < w.size; r++ {
if _, err := w.recvFrom(r, tagBarrier); err != nil {
return err
}
}
for r := 1; r < w.size; r++ {
if err := w.sendTo(r, tagBarrierAck, nil); err != nil {
return err
}
}
return nil
}
if err := w.sendTo(0, tagBarrier, nil); err != nil {
return err
}
_, err := w.recvFrom(0, tagBarrierAck)
return err
}
+416
View File
@@ -0,0 +1,416 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"encoding/binary"
"errors"
"io"
"net"
"strings"
"sync"
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/base"
)
// The world tests cover both transports with the same battery: an
// in-process world over channels and a loopback TCP world over real
// connections, because the contract says they are the same machine.
func TestLaunchOneRank(t *testing.T) {
err := Launch(1, func(w *World) error {
if w.Rank() != 0 || w.Size() != 1 {
t.Fatalf("rank %d of %d", w.Rank(), w.Size())
}
return w.Barrier()
})
if err != nil {
t.Fatal(err)
}
}
func TestLaunchBarrier(t *testing.T) {
for _, size := range []int{2, 3, 5, 8} {
t.Run("", func(t *testing.T) {
reached := make([]int, size)
err := Launch(size, func(w *World) error {
if err := w.Barrier(); err != nil {
return err
}
reached[w.Rank()] = 1
return w.Barrier()
})
if err != nil {
t.Fatal(err)
}
for r, got := range reached {
if got != 1 {
t.Fatalf("rank %d never reported", r)
}
}
})
}
}
// TestLaunchFailsTogether is the fail-fast rule: the rank that sees the
// error fails its world, every other rank's next collective answers
// with an error, and Launch returns the first rank's error in rank
// order.
func TestLaunchFailsTogether(t *testing.T) {
boom := errors.New("boom")
saw := make([]error, 3)
err := Launch(3, func(w *World) error {
if w.Rank() == 1 {
return boom
}
saw[w.Rank()] = w.Barrier()
return saw[w.Rank()]
})
if err == nil {
t.Fatal("a world with a failing rank returned nil")
}
for r, got := range saw {
if r == 1 {
continue
}
if got == nil {
t.Fatalf("rank %d's barrier survived a failed world", r)
}
}
}
func TestLaunchPanicIsAnError(t *testing.T) {
err := Launch(2, func(w *World) error {
if w.Rank() == 1 {
panic("rank one fell over")
}
return w.Barrier()
})
if err == nil || !strings.Contains(err.Error(), "rank 1 panicked") {
t.Fatalf("a panicking rank came back as %v", err)
}
}
// runTCPWorld assembles a loopback TCP world of size ranks, running fn
// on every rank in its own goroutine, and fails the test if any rank
// errors. The listener the tests hand over lets rank 0 know its
// address before it starts. A final barrier runs after fn on every
// rank, the orderly end of the program: no rank tears its world down
// while another still expects words from it.
func runTCPWorld(t *testing.T, size int, opts Options, fn func(w *World) error) {
t.Helper()
ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer ln.Close()
var wg sync.WaitGroup
errs := make([]error, size)
wg.Go(func() {
w, err := listen(ln, size, opts)
if err != nil {
errs[0] = err
return
}
defer w.Close()
if err := fn(w); err != nil {
errs[0] = err
return
}
errs[0] = w.Barrier()
})
for r := 1; r < size; r++ {
wg.Go(func() {
w, err := Join(ln.Addr().String(), opts)
if err != nil {
errs[r] = err
return
}
defer w.Close()
if err := fn(w); err != nil {
errs[r] = err
return
}
errs[r] = w.Barrier()
})
}
wg.Wait()
for r, err := range errs {
if err != nil {
t.Fatalf("rank %d: %v", r, err)
}
}
}
func TestTCPBarrier(t *testing.T) {
runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error {
if err := w.Barrier(); err != nil {
return err
}
return w.Barrier()
})
}
// TestTCPRankOrderIsDialOrder pins the one place ranks come from: the
// order the peers dial in, which the collectives' results never depend
// on.
func TestTCPRankOrderIsDialOrder(t *testing.T) {
runTCPWorld(t, 3, Options{Timeout: 30 * time.Second}, func(w *World) error {
return w.Barrier()
})
}
// TestTCPStrayConnectionTakesNoRank dials the listener with a
// connection that says nothing the handshake would recognise; the world
// must still assemble on the real ranks.
func TestTCPStrayConnectionTakesNoRank(t *testing.T) {
ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer ln.Close()
// The stray speaks first, then goes silent.
stray, err := net.Dial("tcp", ln.Addr().String())
if err != nil {
t.Fatal(err)
}
defer stray.Close()
if _, err := stray.Write([]byte("not the magic at all, but long enough to fill the read")); err != nil {
t.Fatal(err)
}
done := make(chan struct{})
var worldErr error
go func() {
defer close(done)
w, err := listen(ln, 2, Options{Timeout: 30 * time.Second})
if err != nil {
worldErr = err
return
}
w.Close()
}()
// The real rank joins behind the stray.
joined := make(chan error, 1)
go func() {
w, err := Join(ln.Addr().String(), Options{Timeout: 30 * time.Second})
if err != nil {
joined <- err
return
}
w.Close()
joined <- nil
}()
select {
case err := <-joined:
if err != nil {
t.Fatalf("the real rank did not join behind the stray: %v", err)
}
case <-time.After(30 * time.Second):
t.Fatal("the world never assembled")
}
<-done
if worldErr != nil {
t.Fatalf("rank 0: %v", worldErr)
}
}
// TestTCPDeadlineFailsTheCollective is the stuck-peer rule: a rank
// whose peer stops answering is errored out by its own deadline, never
// left hanging. The peer's world carries the short timeout from birth,
// so nothing changes under a running pump.
func TestTCPDeadlineFailsTheCollective(t *testing.T) {
ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer ln.Close()
hubDone := make(chan *World, 1)
go func() {
// The hub assembles and stays silent: it never answers.
w, err := listen(ln, 2, Options{Timeout: 30 * time.Second})
if err != nil {
w = nil
}
hubDone <- w
}()
peer, err := Join(ln.Addr().String(), Options{Timeout: 80 * time.Millisecond})
if err != nil {
t.Fatal(err)
}
defer peer.Close()
hub := <-hubDone
if hub == nil {
t.Fatal("the hub did not assemble")
}
defer hub.Close()
if err := peer.Barrier(); err == nil {
t.Fatal("a barrier against a silent hub succeeded")
}
}
// TestTCPMaxMessageRefused: a peer announcing a payload beyond the
// ceiling is an error before any allocation.
func TestTCPMaxMessageRefused(t *testing.T) {
ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer ln.Close()
const ceiling = int64(1 << 20)
wCh := make(chan *World, 1)
go func() {
w, err := listen(ln, 2, Options{Timeout: 30 * time.Second, MaxMessage: ceiling})
if err != nil {
wCh <- nil
return
}
wCh <- w
}()
// The fake peer does the handshake by hand, then announces an
// oversized frame.
peerDone := make(chan error, 1)
go func() {
conn, err := net.Dial("tcp", ln.Addr().String())
if err != nil {
peerDone <- err
return
}
defer conn.Close()
if err := sendHello(conn); err != nil {
peerDone <- err
return
}
if _, _, err := readWelcome(conn, time.Now().Add(30*time.Second)); err != nil {
peerDone <- err
return
}
link := newTCPLink(conn)
var head [frameHeaderLen]byte
head[0] = tagBarrier
head[1] = frameVersion
binary.LittleEndian.PutUint32(head[4:], 1) // from
binary.LittleEndian.PutUint32(head[8:], 0) // dest
binary.LittleEndian.PutUint64(head[12:], uint64(ceiling+1))
if _, err := link.wr.Write(head[:]); err != nil {
peerDone <- err
return
}
peerDone <- link.wr.Flush()
}()
w0 := <-wCh
if w0 == nil {
t.Fatal("rank 0 did not assemble")
}
defer w0.Close()
if err := <-peerDone; err != nil {
t.Fatalf("the fake peer: %v", err)
}
if _, err := w0.recvFrom(1, tagBarrier); err == nil {
t.Fatal("an oversized frame was received")
}
}
// TestHandshake walks the joining words over an in-memory connection:
// the right magic passes both ways, a wrong magic is refused, and an
// impossible rank answer is refused.
func TestHandshake(t *testing.T) {
c, s := net.Pipe()
defer c.Close()
defer s.Close()
deadline := time.Now().Add(5 * time.Second)
go sendHello(c)
if err := readHello(s, deadline); err != nil {
t.Fatalf("the right magic was refused: %v", err)
}
go sendWelcome(c, 4, 2)
size, rank, err := readWelcome(s, deadline)
if err != nil {
t.Fatalf("the welcome did not read: %v", err)
}
if size != 4 || rank != 2 {
t.Fatalf("welcome answered size %d rank %d", size, rank)
}
var bad [handshakeLen]byte
copy(bad[0:4], []byte("XXXX"))
bad[4] = frameVersion
go c.Write(bad[:])
if err := readHello(s, deadline); err == nil {
t.Fatal("a wrong magic was accepted")
}
go sendWelcome(c, 4, 4)
if _, _, err := readWelcome(s, deadline); err == nil {
t.Fatal("a rank beyond the world's size was accepted")
}
}
// The compile-time guards on the error surface: every failure this
// package reports keeps the library's prefix.
func TestErrorPrefix(t *testing.T) {
err := base.Errf("spmd: test")
if err == nil || !strings.HasPrefix(err.Error(), "tensor: ") {
t.Fatalf("the package error lost its prefix: %v", err)
}
var _ io.Closer = (*World)(nil)
}
// TestTCPNegativeLengthRefused: a frame length with its top bit set
// turns negative through the signed conversion; the receiver refuses
// it instead of allocating from it, which once panicked the hub.
func TestTCPNegativeLengthRefused(t *testing.T) {
ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer ln.Close()
wCh := make(chan *World, 1)
go func() {
w, err := listen(ln, 2, Options{Timeout: 30 * time.Second})
if err != nil {
w = nil
}
wCh <- w
}()
peerDone := make(chan error, 1)
go func() {
conn, err := net.Dial("tcp", ln.Addr().String())
if err != nil {
peerDone <- err
return
}
defer conn.Close()
if err := sendHello(conn); err != nil {
peerDone <- err
return
}
if _, _, err := readWelcome(conn, time.Now().Add(30*time.Second)); err != nil {
peerDone <- err
return
}
link := newTCPLink(conn)
var head [frameHeaderLen]byte
head[0] = tagBarrier
head[1] = frameVersion
binary.LittleEndian.PutUint32(head[4:], 1)
binary.LittleEndian.PutUint32(head[8:], 0)
binary.LittleEndian.PutUint64(head[12:], uint64(1)<<63)
if _, err := link.wr.Write(head[:]); err != nil {
peerDone <- err
return
}
peerDone <- link.wr.Flush()
}()
w0 := <-wCh
if w0 == nil {
t.Fatal("rank 0 did not assemble")
}
defer w0.Close()
if err := <-peerDone; err != nil {
t.Fatalf("the fake peer: %v", err)
}
if _, err := w0.recvFrom(1, tagBarrier); err == nil {
t.Fatal("a negative length was received")
}
}