Files

212 lines
5.7 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
})
})
}
}