feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
+209
@@ -0,0 +1,209 @@
|
||||
// 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
|
||||
}
|
||||
Reference in New Issue
Block a user