Files
tensor/spmd/proc_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

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)
}
}
}