feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,180 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package spmd
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// globalRows builds the [rows, width] fixture the halo tests deal out.
|
||||
func globalRows(rows, width int) *core.Array {
|
||||
vals := make([]float64, rows*width)
|
||||
for i := range vals {
|
||||
vals[i] = float64(i*31%(rows*width)) * 0.25
|
||||
}
|
||||
a, err := core.FromFloats(vals, rows, width)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// dealRows2D cuts rows [lo, hi) off the global array, an empty array
|
||||
// for an empty range.
|
||||
func dealRows2D(tb testing.TB, whole *core.Array, lo, hi int) *core.Array {
|
||||
tb.Helper()
|
||||
rest := 1
|
||||
for _, d := range whole.Shape()[1:] {
|
||||
rest *= d
|
||||
}
|
||||
wire, err := encodePart(nil, whole, append([]int{hi - lo}, whole.Shape()[1:]...), lo*rest, (hi-lo)*rest)
|
||||
if err != nil {
|
||||
tb.Fatal(err)
|
||||
}
|
||||
a, err := decodeWire(wire)
|
||||
if err != nil {
|
||||
tb.Fatal(err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// checkHalos asserts one rank's halos against the whole's neighbouring
|
||||
// rows, treating empty neighbours and empty pieces as no data.
|
||||
func checkHalos(t *testing.T, w *World, whole *core.Array, rows, halos int, upper, lower *core.Array) {
|
||||
t.Helper()
|
||||
span := mustPartition(t, rows, w.Size(), w.Rank())
|
||||
neighbourRows := func(rank int) int {
|
||||
if rank < 0 || rank >= w.Size() {
|
||||
return 0
|
||||
}
|
||||
return mustPartition(t, rows, w.Size(), rank).Len()
|
||||
}
|
||||
if span.Len() == 0 {
|
||||
if upper != nil || lower != nil {
|
||||
t.Fatalf("rank %d: an empty piece answered halos", w.Rank())
|
||||
}
|
||||
return
|
||||
}
|
||||
if upper != nil && upper.Len() == 0 {
|
||||
upper = nil
|
||||
}
|
||||
if lower != nil && lower.Len() == 0 {
|
||||
lower = nil
|
||||
}
|
||||
if w.Rank() == 0 || neighbourRows(w.Rank()-1) == 0 {
|
||||
if upper != nil {
|
||||
t.Fatalf("rank %d: an upper halo arrived where no rows live", w.Rank())
|
||||
}
|
||||
} else {
|
||||
if upper == nil {
|
||||
t.Fatalf("rank %d: the upper halo is nil with a living neighbour", w.Rank())
|
||||
}
|
||||
if want := dealRows2D(t, whole, span.Lo-halos, span.Lo); !sameBits(want, upper) {
|
||||
t.Fatalf("rank %d: the upper halo differs from rows [%d, %d)", w.Rank(), span.Lo-halos, span.Lo)
|
||||
}
|
||||
}
|
||||
if w.Rank()+1 == w.Size() || neighbourRows(w.Rank()+1) == 0 {
|
||||
if lower != nil {
|
||||
t.Fatalf("rank %d: a lower halo arrived where no rows live", w.Rank())
|
||||
}
|
||||
} else {
|
||||
if lower == nil {
|
||||
t.Fatalf("rank %d: the lower halo is nil with a living neighbour", w.Rank())
|
||||
}
|
||||
if want := dealRows2D(t, whole, span.Hi, span.Hi+halos); !sameBits(want, lower) {
|
||||
t.Fatalf("rank %d: the lower halo differs from rows [%d, %d)", w.Rank(), span.Hi, span.Hi+halos)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestExchangeHalosMatchTheWhole: every rank's upper and lower halos
|
||||
// are exactly the whole's neighbouring rows, and edges without a
|
||||
// living neighbour answer nil.
|
||||
func TestExchangeHalosMatchTheWhole(t *testing.T) {
|
||||
for _, size := range []int{1, 2, 3, 5, 8} {
|
||||
for _, halos := range []int{0, 1, 3} {
|
||||
const rows, width = 40, 5
|
||||
whole := globalRows(rows, width)
|
||||
err := Launch(size, func(w *World) error {
|
||||
span := mustPartition(t, rows, w.Size(), w.Rank())
|
||||
local := dealRows2D(t, whole, span.Lo, span.Hi)
|
||||
upper, lower, err := w.ExchangeHalos(local, halos)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
checkHalos(t, w, whole, rows, halos, upper, lower)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("size %d halos %d: %v", size, halos, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestExchangeHalosStencilMatchesSerial runs a three-point smoothing
|
||||
// step over a sharded series and compares it, element for element,
|
||||
// with the same step computed on the whole array: the halos must make
|
||||
// the distributed stencil see exactly the neighbours the serial one
|
||||
// sees.
|
||||
func TestExchangeHalosStencilMatchesSerial(t *testing.T) {
|
||||
const n = 1000
|
||||
whole := fixtureArray(n)
|
||||
serial := make([]float64, n)
|
||||
for i := 1; i < n-1; i++ {
|
||||
serial[i] = (whole.FloatAt(i-1) + whole.FloatAt(i) + whole.FloatAt(i+1)) / 3
|
||||
}
|
||||
err := Launch(4, func(w *World) error {
|
||||
span := mustPartition(t, n, w.Size(), w.Rank())
|
||||
local := narrowSliceFor(t, whole, span)
|
||||
upper, lower, err := w.ExchangeHalos(local, 1)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for i := range span.Len() {
|
||||
g := span.Lo + i
|
||||
if g == 0 || g == n-1 {
|
||||
continue
|
||||
}
|
||||
var left, right float64
|
||||
if i == 0 {
|
||||
left = upper.FloatAt(0)
|
||||
} else {
|
||||
left = local.FloatAt(i - 1)
|
||||
}
|
||||
if i == span.Len()-1 {
|
||||
right = lower.FloatAt(0)
|
||||
} else {
|
||||
right = local.FloatAt(i + 1)
|
||||
}
|
||||
if got := (left + local.FloatAt(i) + right) / 3; got != serial[g] {
|
||||
t.Fatalf("rank %d global %d: %v against the serial %v", w.Rank(), g, got, serial[g])
|
||||
}
|
||||
}
|
||||
return w.Barrier()
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestExchangeHalosOverTCP runs the neighbour exchange over real
|
||||
// connections.
|
||||
func TestExchangeHalosOverTCP(t *testing.T) {
|
||||
const rows, width = 200, 3
|
||||
whole := globalRows(rows, width)
|
||||
runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error {
|
||||
span := mustPartition(t, rows, w.Size(), w.Rank())
|
||||
local := dealRows2D(t, whole, span.Lo, span.Hi)
|
||||
upper, lower, err := w.ExchangeHalos(local, 2)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
checkHalos(t, w, whole, rows, 2, upper, lower)
|
||||
return nil
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user