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