// Copyright (c) 2026 Petr BalvĂ­n (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) } } }