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