202 lines
5.2 KiB
Go
202 lines
5.2 KiB
Go
// 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)
|
|
}
|
|
}
|
|
}
|