Files
tensor/spmd/halo_test.go
T

181 lines
5.1 KiB
Go
Raw 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 (
"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
})
}