212 lines
5.7 KiB
Go
212 lines
5.7 KiB
Go
// 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
|
|
})
|
|
})
|
|
}
|
|
}
|