Files
tensor/spmd/halo.go
T

210 lines
7.0 KiB
Go
Raw 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 (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// ExchangeHalos serves the domain-decomposed simulations: a rank whose
// piece is cut along the first axis hands its edge rows to its two
// neighbours and receives theirs. upper carries the rows that precede
// the piece in the global array (the lower edge of rank-1), lower the
// rows that follow it (the upper edge of rank+1); a rank at the
// world's edge receives nil for the side where no neighbour lives,
// and a zero halo width answers the same nil on both sides: nil is
// the answer every zero-length slab takes. The exchange moves bits
// and moves nothing else, so there is
// no order for it to get wrong: the runs in both directions are fixed
// by the rank indices, never by the arrival order. It is
// ExchangeHalosOnGrid on the one-dimensional grid of the whole world,
// cut along axis 0.
func (w *World) ExchangeHalos(local *core.Array, halos int) (*core.Array, *core.Array, error) {
return w.ExchangeHalosOnGrid(local, halos, 0, []int{w.size})
}
// ExchangeHalosOnGrid is the halo exchange of a domain-decomposed
// simulation on a process grid. grid lays the world's ranks out
// row-major, grid[a] naming how many positions the grid runs along
// axis a, and its product must be the world's size; rank r sits at
// coordinate (r/prod(grid[axis+1:]))%grid[axis] along axis, so its
// two neighbours there sit one grid step away, at rank-dist and
// rank+dist with dist = prod(grid[axis+1:]), and a rank at the grid's
// edge has no neighbour on that side and receives nil for it. Each
// rank hands the slab of width halos at its piece's edge along axis
// to the neighbour it borders there and receives theirs: upper
// carries the slab that precedes the piece along the axis, lower the
// slab that follows it, both keeping every other dimension of the
// piece's shape. A zero halo width answers nil on both sides, an
// empty piece joins with empty wires and answers no halos, and a
// neighbour with no piece to offer answers nil too, so the protocol
// stays symmetric. The exchange moves bits and moves nothing else, so
// there is no order
// for it to get wrong: the runs in both directions are fixed by the
// rank indices, never by the arrival order.
func (w *World) ExchangeHalosOnGrid(local *core.Array, halos, axis int, grid []int) (*core.Array, *core.Array, error) {
if err := w.status(); err != nil {
return nil, nil, err
}
if halos < 0 {
return nil, nil, w.fail(base.Errf("spmd: a negative halo width %d", halos))
}
if local.NDim() == 0 {
return nil, nil, w.fail(base.Errf("spmd: the halo exchange needs a dimension to cut the edges from"))
}
if axis < 0 || axis >= local.NDim() {
return nil, nil, w.fail(base.Errf("spmd: halo axis %d is outside the array's %d dimensions", axis, local.NDim()))
}
if axis >= len(grid) {
return nil, nil, w.fail(base.Errf("spmd: halo axis %d is outside the grid's %d axes", axis, len(grid)))
}
extent := 1
for _, d := range grid {
if d < 0 {
return nil, nil, w.fail(base.Errf("spmd: a grid cannot name a negative extent %d", d))
}
if extent > 0 && d > math.MaxInt/extent {
return nil, nil, w.fail(base.Errf("spmd: a grid of %s lays out more ranks than the world can hold", base.ShapeText(grid)))
}
extent *= d
}
if extent != w.size {
return nil, nil, w.fail(base.Errf("spmd: a grid of %s lays out %d ranks against the world's %d", base.ShapeText(grid), extent, w.size))
}
// The rank's place in the row-major grid: the coordinate along the
// cut axis, and the rank distance of one grid step along it.
dist := 1
for _, d := range grid[axis+1:] {
dist *= d
}
coord := (w.rank / dist) % grid[axis]
hasLower, hasUpper := coord > 0, coord+1 < grid[axis]
shape := local.Shape()
if local.Len() == 0 {
// An empty piece owns no slab, so it carries no edges: it
// still joins the exchange with empty wires, so the
// neighbours' protocol stays symmetric, and answers no halos.
emptyShape := make([]int, len(shape))
copy(emptyShape, shape)
emptyShape[axis] = 0
empty, err := encodeHead(nil, local.Dtype(), emptyShape)
if err != nil {
return nil, nil, w.fail(err)
}
if hasUpper {
if err := w.sendTo(w.rank+dist, tagHalo, empty); err != nil {
return nil, nil, err
}
}
if hasLower {
if err := w.sendTo(w.rank-dist, tagHalo, empty); err != nil {
return nil, nil, err
}
}
if hasUpper {
if _, err := w.recvFrom(w.rank+dist, tagHalo); err != nil {
return nil, nil, err
}
}
if hasLower {
if _, err := w.recvFrom(w.rank-dist, tagHalo); err != nil {
return nil, nil, err
}
}
return nil, nil, nil
}
span := shape[axis]
if halos > span {
return nil, nil, w.fail(base.Errf("spmd: a halo width of %d rows exceeds the piece's %d", halos, span))
}
lead, rest := 1, 1
for _, d := range shape[:axis] {
lead *= d
}
for _, d := range shape[axis+1:] {
rest *= d
}
edgeShape := make([]int, len(shape))
copy(edgeShape, shape)
edgeShape[axis] = halos
// The edge slab is one run of halos*rest elements per leading
// position, contiguous in the row-major layout; the wire carries
// the runs joined under the slab's own shape, whose extent along
// the axis is the halo width and whose other extents are the
// piece's own.
edge := func(high bool) ([]byte, error) {
wire, err := encodeHead(nil, local.Dtype(), edgeShape)
if err != nil {
return nil, err
}
count := halos * rest
for p := range lead {
first := p*span*rest + (span-halos)*rest
if !high {
first = p * span * rest
}
part, err := encodePart(nil, local, []int{count}, first, count)
if err != nil {
return nil, err
}
wire = append(wire, part[partHeadLen:]...)
}
return wire, nil
}
// Two phases: the lower halos travel to the upper neighbour
// first, the upper halos to the lower one second. The sends ride
// the links' buffers and the networked links drain through the
// hub, so no rank waits on a rank that is waiting on it.
if hasUpper {
highEdge, err := edge(true)
if err != nil {
return nil, nil, w.fail(err)
}
if err := w.sendTo(w.rank+dist, tagHalo, highEdge); err != nil {
return nil, nil, err
}
}
var lower *core.Array
if hasUpper {
data, err := w.recvFrom(w.rank+dist, tagHalo)
if err != nil {
return nil, nil, err
}
lower, err = decodeWire(data)
if err != nil {
return nil, nil, w.fail(err)
}
if lower.Len() == 0 {
lower = nil // the neighbour owns no slab
}
}
if hasLower {
lowEdge, err := edge(false)
if err != nil {
return nil, nil, w.fail(err)
}
if err := w.sendTo(w.rank-dist, tagHalo, lowEdge); err != nil {
return nil, nil, err
}
}
var upper *core.Array
if hasLower {
data, err := w.recvFrom(w.rank-dist, tagHalo)
if err != nil {
return nil, nil, err
}
upper, err = decodeWire(data)
if err != nil {
return nil, nil, w.fail(err)
}
if upper.Len() == 0 {
upper = nil // the neighbour owns no slab
}
}
return upper, lower, nil
}