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