373 lines
12 KiB
Go
373 lines
12 KiB
Go
// 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)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|