Files

373 lines
12 KiB
Go
Raw Permalink 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 (
"slices"
"strings"
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The grid halo tests pin ExchangeHalosOnGrid against the whole array:
// the rank's row-major place in the grid, its neighbours along the cut
// axis and the slabs that travel between them, with the serial stencil
// as the exacter judge of the geometry.
// gridFixture builds the [rows, cols] fixture the grid tests deal out.
func gridFixture(rows, cols int) *core.Array {
vals := make([]float64, rows*cols)
for i := range vals {
vals[i] = float64(i*37%(rows*cols)) * 0.5
}
return mk(core.FromFloats(vals, rows, cols))
}
// tileSpan cuts tile idx's run of an extent dealt into parts tiles,
// the deal that leaves the remainders on the last tiles.
func tileSpan(extent, parts, idx int) (int, int) {
return idx * extent / parts, (idx + 1) * extent / parts
}
// gridCoords maps a rank into a row-major grid, coordinate a being
// (rank/prod(grid[a+1:]))%grid[a].
func gridCoords(rank int, grid []int) []int {
coords := make([]int, len(grid))
rest := 1
for a := len(grid) - 1; a >= 0; a-- {
coords[a] = (rank / rest) % grid[a]
rest *= grid[a]
}
return coords
}
// gridRank is gridCoords' inverse: the rank a row-major grid puts on
// the named coordinates.
func gridRank(coords []int, grid []int) int {
rank := 0
for a := range grid {
rank = rank*grid[a] + coords[a]
}
return rank
}
// dealTile cuts rows [rlo, rhi) by cols [clo, chi) off a 2-D whole,
// an empty array for an empty range.
func dealTile(t *testing.T, whole *core.Array, rlo, rhi, clo, chi int) *core.Array {
t.Helper()
width := whole.Shape()[1]
wire, err := encodeHead(nil, whole.Dtype(), []int{rhi - rlo, chi - clo})
if err != nil {
t.Fatal(err)
}
for r := rlo; r < rhi; r++ {
part, err := encodePart(nil, whole, []int{chi - clo}, r*width+clo, chi-clo)
if err != nil {
t.Fatal(err)
}
wire = append(wire, part[2+8:]...)
}
a, err := decodeWire(wire)
if err != nil {
t.Fatal(err)
}
return a
}
// checkGridExchange runs one rank's grid exchange and judges its halos
// against the neighbouring tiles' own edge slabs, worked out from the
// row-major grid arithmetic beside the exchange: a halo keeps the
// piece's other dimensions, so the expected slab is the neighbour's
// tile narrowed to the halo width along the cut axis. Empty pieces
// answer nothing, and empty neighbours and grid edges hand nil back.
func checkGridExchange(t *testing.T, w *World, whole *core.Array, halos, axis int, grid []int) {
t.Helper()
two := grid
if len(two) == 1 {
two = []int{two[0], 1}
}
coords := gridCoords(w.Rank(), two)
rlo, rhi, clo, chi := rankTile(whole, two, w.Rank())
local := dealTile(t, whole, rlo, rhi, clo, chi)
upper, lower, err := w.ExchangeHalosOnGrid(local, halos, axis, grid)
if err != nil {
t.Fatalf("rank %d: %v", w.Rank(), err)
}
if local.Len() == 0 {
if upper != nil || lower != nil {
t.Fatalf("rank %d: an empty piece answered halos", w.Rank())
}
return
}
sides := []struct {
name string
got *core.Array
off int
}{
{"upper", upper, -1},
{"lower", lower, 1},
}
for _, side := range sides {
got := side.got
if got != nil && got.Len() == 0 {
got = nil // a zero-width slab answers nothing
}
c := coords[axis] + side.off
if c < 0 || c >= two[axis] {
if got != nil {
t.Fatalf("rank %d: a %s halo arrived where no neighbour lives", w.Rank(), side.name)
}
continue
}
sideCoords := slices.Clone(coords)
sideCoords[axis] = c
nrlo, nrhi, nclo, nchi := rankTile(whole, two, gridRank(sideCoords, two))
if nrhi == nrlo || nchi == nclo {
// The neighbour owns no elements, so it carried no slab.
if got != nil {
t.Fatalf("rank %d: a %s halo arrived from an empty neighbour", w.Rank(), side.name)
}
continue
}
// The neighbour's slab: its tile's edge along the cut axis,
// kept whole on every other dimension.
cutLo, cutHi := nclo, nchi
if axis == 0 {
cutLo, cutHi = nrlo, nrhi
}
sLo, sHi := cutLo, cutLo+halos
if side.off < 0 {
sLo, sHi = cutHi-halos, cutHi
}
var want *core.Array
if axis == 0 {
want = dealTile(t, whole, sLo, sHi, nclo, nchi)
} else {
want = dealTile(t, whole, nrlo, nrhi, sLo, sHi)
}
if want.Len() == 0 {
if got != nil {
t.Fatalf("rank %d: a %s halo arrived for a zero-width slab", w.Rank(), side.name)
}
continue
}
if got == nil {
t.Fatalf("rank %d: the %s halo is nil with a living neighbour", w.Rank(), side.name)
}
if !slices.Equal(got.Shape(), want.Shape()) || !sameBits(want, got) {
t.Fatalf("rank %d: the %s halo differs from the neighbouring tile's slab along axis %d",
w.Rank(), side.name, axis)
}
}
}
// rankTile is dealTile's own cut for one rank of a 2-D grid.
func rankTile(whole *core.Array, grid []int, rank int) (rlo, rhi, clo, chi int) {
coords := gridCoords(rank, grid)
rlo, rhi = tileSpan(whole.Shape()[0], grid[0], coords[0])
clo, chi = tileSpan(whole.Shape()[1], grid[1], coords[1])
return rlo, rhi, clo, chi
}
// TestExchangeHalosOnGridMatchesTheWhole walks the whole matrix the
// API claims: world sizes 1 to 8, halo widths 0, 1 and 3, every
// two-factor grid of each size beside the degenerate one-dimensional
// one, cut along every axis the grid names. The expected slabs come
// from the rank's grid coordinates worked out beside the exchange, so
// a wrong row-major mapping, a wrong neighbour distance or a wrong
// slab cannot pass: on the grid [2, 3] the axis 1 neighbours are
// rank+-1 and the axis 0 neighbours rank+-3, and the checks know it
// independently.
func TestExchangeHalosOnGridMatchesTheWhole(t *testing.T) {
// The extents divide out to tiles of at least three elements on
// every factor grid up to eight positions, so even the widest
// pinned halo fits the piece it travels from.
const rows, cols = 24, 24
whole := gridFixture(rows, cols)
for _, size := range []int{1, 2, 3, 5, 8} {
grids := [][]int{{size}}
for a := 1; a <= size; a++ {
if size%a == 0 {
grids = append(grids, []int{a, size / a})
}
}
for _, grid := range grids {
for _, halos := range []int{0, 1, 3} {
for axis := range len(grid) {
err := Launch(size, func(w *World) error {
checkGridExchange(t, w, whole, halos, axis, grid)
return nil
})
if err != nil {
t.Fatalf("size %d grid %v halos %d axis %d: %v", size, grid, halos, axis, err)
}
}
}
}
}
}
// TestExchangeHalosOnGridEmptyPieces: the pieces the tiling leaves
// empty join the exchange symmetrically, answering no halos, and the
// neighbours see nil from the empty side, in both grid orientations.
func TestExchangeHalosOnGridEmptyPieces(t *testing.T) {
whole := gridFixture(1, 2)
for _, grid := range [][]int{{2, 2}, {3, 2}, {2, 3}} {
for axis := range 2 {
err := Launch(grid[0]*grid[1], func(w *World) error {
checkGridExchange(t, w, whole, 1, axis, grid)
return nil
})
if err != nil {
t.Fatalf("grid %v axis %d: %v", grid, axis, err)
}
}
}
}
// TestExchangeHalosOnGridStencilMatchesSerial runs the five-point
// stencil over a field tiled 2 by 3, exchanging halos along both grid
// axes, and compares every tile element for element against the same
// step computed on the whole field. The comparison is the matrix
// claim made flesh: on the row-major grid [2, 3] the axis 1
// neighbours are rank+-1, the axis 0 neighbours rank+-3, and the
// distributed stencil may only agree bit for bit when the slabs
// deliver exactly the neighbours the serial walk sees.
func TestExchangeHalosOnGridStencilMatchesSerial(t *testing.T) {
const rows, cols = 6, 12
whole := gridFixture(rows, cols)
serial := make([]float64, rows*cols)
for r := range rows {
for c := range cols {
up, down := 0.0, 0.0
if r > 0 {
up = whole.FloatAt((r-1)*cols + c)
}
if r < rows-1 {
down = whole.FloatAt((r+1)*cols + c)
}
left, right := 0.0, 0.0
if c > 0 {
left = whole.FloatAt(r*cols + c - 1)
}
if c < cols-1 {
right = whole.FloatAt(r*cols + c + 1)
}
serial[r*cols+c] = (up + left + whole.FloatAt(r*cols+c) + right + down) / 5
}
}
const halos = 1
err := Launch(6, func(w *World) error {
rlo, rhi, clo, chi := rankTile(whole, []int{2, 3}, w.Rank())
local := dealTile(t, whole, rlo, rhi, clo, chi)
upper, lower, err := w.ExchangeHalosOnGrid(local, halos, 0, []int{2, 3})
if err != nil {
return err
}
leftHalo, rightHalo, err := w.ExchangeHalosOnGrid(local, halos, 1, []int{2, 3})
if err != nil {
return err
}
tileRows, tileCols := rhi-rlo, chi-clo
for ri := range tileRows {
for cj := range tileCols {
gi, gj := rlo+ri, clo+cj
centre := local.FloatAt(ri*tileCols + cj)
up := 0.0
if ri == 0 {
if gi > 0 {
up = upper.FloatAt(cj) // the halo's last row adjoins the tile
}
} else {
up = local.FloatAt((ri-1)*tileCols + cj)
}
down := 0.0
if ri == tileRows-1 {
if gi < rows-1 {
down = lower.FloatAt(cj)
}
} else {
down = local.FloatAt((ri+1)*tileCols + cj)
}
left := 0.0
if cj == 0 {
if gj > 0 {
left = leftHalo.FloatAt(ri) // the halo's last column adjoins the tile
}
} else {
left = local.FloatAt(ri*tileCols + cj - 1)
}
right := 0.0
if cj == tileCols-1 {
if gj < cols-1 {
right = rightHalo.FloatAt(ri)
}
} else {
right = local.FloatAt(ri*tileCols + cj + 1)
}
got := (up + left + centre + right + down) / 5
if got != serial[gi*cols+gj] {
t.Fatalf("rank %d global (%d, %d): %v against the serial %v",
w.Rank(), gi, gj, got, serial[gi*cols+gj])
}
}
}
return w.Barrier()
})
if err != nil {
t.Fatal(err)
}
}
// TestExchangeHalosOnGridOverTCP runs the grid exchange over real
// connections, the axis 1 slabs included.
func TestExchangeHalosOnGridOverTCP(t *testing.T) {
whole := gridFixture(24, 18)
runTCPWorld(t, 6, Options{Timeout: 30 * time.Second}, func(w *World) error {
checkGridExchange(t, w, whole, 2, 1, []int{2, 3})
checkGridExchange(t, w, whole, 2, 0, []int{2, 3})
return nil
})
}
// TestExchangeHalosOnGridRefusesTheHostile: every invalid argument is
// a named error before any frame moves. Every rank makes the same
// invalid call, so an exchange that ever started would deadlock the
// world instead of answering, which is what pins the ordering too.
func TestExchangeHalosOnGridRefusesTheHostile(t *testing.T) {
local := mk(core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3))
cases := []struct {
name string
local *core.Array
halos, axis int
grid []int
want string
}{
{"negative halos", local, -1, 0, []int{2}, "negative halo width"},
{"axis past the array", local, 1, 2, []int{2, 2}, "outside the array"},
{"negative axis", local, 1, -1, []int{2, 2}, "outside the array"},
{"axis past the grid", local, 1, 1, []int{2}, "outside the grid"},
{"grid short of the world", local, 1, 0, []int{1, 3}, "against the world's"},
{"negative grid extent", local, 1, 0, []int{-2}, "negative extent"},
{"halos wider than the axis", local, 3, 0, []int{2}, "exceeds the piece's"},
}
for _, tc := range cases {
err := Launch(2, func(w *World) error {
_, _, err := w.ExchangeHalosOnGrid(tc.local, tc.halos, tc.axis, tc.grid)
if err == nil {
t.Fatalf("%s: the hostile call was accepted", tc.name)
}
if !strings.Contains(err.Error(), tc.want) {
t.Fatalf("%s: the error %q does not name the fault", tc.name, err)
}
return nil
})
if err != nil {
t.Fatalf("%s: %v", tc.name, err)
}
}
}