feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
+248
@@ -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()
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user