181 lines
5.1 KiB
Go
181 lines
5.1 KiB
Go
// 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
|
||
|
|
})
|
||
|
|
}
|