// Copyright (c) 2026 Petr BalvĂ­n (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) } } }