210 lines
7.0 KiB
Go
210 lines
7.0 KiB
Go
// 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
|
|
}
|