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