feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user