286 lines
7.9 KiB
Go
286 lines
7.9 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"
|
|
)
|
|
|
|
// The movement collectives carry bits and carry nothing else: no
|
|
// arithmetic happens in transit, so their answers are deterministic by
|
|
// construction. The wire is the only form that travels, and every
|
|
// collective's join order is a function of the data, never of the
|
|
// order frames happen to arrive in.
|
|
|
|
// Broadcast delivers root's array to every rank of the world. The root
|
|
// gets the array it passed; every other rank gets an equal copy, bits
|
|
// included. A world of one rank hands the array back.
|
|
func (w *World) Broadcast(a *core.Array, root int) (*core.Array, error) {
|
|
if err := w.status(); err != nil {
|
|
return nil, err
|
|
}
|
|
if err := checkRoot(root, w.size); err != nil {
|
|
return nil, w.fail(err)
|
|
}
|
|
if w.size == 1 {
|
|
return a, nil
|
|
}
|
|
if w.rank == root {
|
|
wire, err := encodeArray(nil, a)
|
|
if err != nil {
|
|
return nil, w.fail(err)
|
|
}
|
|
for r := range w.size {
|
|
if r == root {
|
|
continue
|
|
}
|
|
if err := w.sendTo(r, tagBroadcast, wire); err != nil {
|
|
return nil, err
|
|
}
|
|
}
|
|
return a, nil
|
|
}
|
|
data, err := w.recvFrom(root, tagBroadcast)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
got, err := decodeWire(data)
|
|
if err != nil {
|
|
return nil, w.fail(err)
|
|
}
|
|
return got, nil
|
|
}
|
|
|
|
// Scatter deals root's array out along its first dimension: rank r
|
|
// receives exactly the canonical partition's piece of the axis, so the
|
|
// world's data lands on the boundaries the shard reductions compose
|
|
// on. The global shape travels first, so every rank can name its own
|
|
// piece before any payload moves; the span that comes back names it.
|
|
func (w *World) Scatter(global *core.Array, root int) (*core.Array, Span, error) {
|
|
if err := w.status(); err != nil {
|
|
return nil, Span{}, err
|
|
}
|
|
if err := checkRoot(root, w.size); err != nil {
|
|
return nil, Span{}, w.fail(err)
|
|
}
|
|
if global.NDim() == 0 {
|
|
return nil, Span{}, w.fail(base.Errf("spmd: Scatter needs a dimension to deal along"))
|
|
}
|
|
var shape []int
|
|
if w.rank == root {
|
|
shape = global.Shape()
|
|
head, err := encodeHead(nil, global.Dtype(), shape)
|
|
if err != nil {
|
|
return nil, Span{}, w.fail(err)
|
|
}
|
|
for r := range w.size {
|
|
if r == root {
|
|
continue
|
|
}
|
|
if err := w.sendTo(r, tagScatterHead, head); err != nil {
|
|
return nil, Span{}, err
|
|
}
|
|
}
|
|
} else {
|
|
data, err := w.recvFrom(root, tagScatterHead)
|
|
if err != nil {
|
|
return nil, Span{}, err
|
|
}
|
|
_, shape, err = decodeHead(data)
|
|
if err != nil {
|
|
return nil, Span{}, w.fail(err)
|
|
}
|
|
}
|
|
if len(shape) == 0 {
|
|
return nil, Span{}, w.fail(base.Errf("spmd: the dealt head names no dimension to deal along"))
|
|
}
|
|
gn := shape[0]
|
|
rest, ok := elementCount(shape[1:])
|
|
if !ok {
|
|
return nil, Span{}, w.fail(base.Errf("spmd: Scatter's shape %v overflows the element count", shape))
|
|
}
|
|
full, ok := elementCount(shape)
|
|
if !ok {
|
|
return nil, Span{}, w.fail(base.Errf("spmd: Scatter's shape %v overflows the element count", shape))
|
|
}
|
|
if w.rank == root && full != global.Len() {
|
|
return nil, Span{}, w.fail(base.Errf("spmd: Scatter's shape %s names %d elements against the array's %d",
|
|
base.ShapeText(shape), full, global.Len()))
|
|
}
|
|
span, err := Partition(gn, w.size, w.rank)
|
|
if err != nil {
|
|
return nil, Span{}, w.fail(err)
|
|
}
|
|
slabShape := make([]int, 0, len(shape))
|
|
slabShape = append(slabShape, span.Len())
|
|
slabShape = append(slabShape, shape[1:]...)
|
|
local, err := func() (*core.Array, error) {
|
|
if w.rank == root {
|
|
wire, err := encodePart(nil, global, slabShape, span.Lo*rest, span.Len()*rest)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
for r := range w.size {
|
|
if r == root {
|
|
continue
|
|
}
|
|
s, err := Partition(gn, w.size, r)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
wire, err := encodePart(nil, global,
|
|
append([]int{s.Len()}, shape[1:]...), s.Lo*rest, s.Len()*rest)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if err := w.sendTo(r, tagScatter, wire); err != nil {
|
|
return nil, err
|
|
}
|
|
}
|
|
return decodeWire(wire)
|
|
}
|
|
wire, err := w.recvFrom(root, tagScatter)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return decodeWire(wire)
|
|
}()
|
|
if err != nil {
|
|
return nil, Span{}, w.fail(err)
|
|
}
|
|
return local, span, nil
|
|
}
|
|
|
|
// Gather raises the dealt pieces back into the global array on root,
|
|
// joining them in rank order, which is the canonical partition's order.
|
|
// Ranks other than root answer with nil: they hold their piece, the
|
|
// root holds the whole.
|
|
func (w *World) Gather(local *core.Array, root int) (*core.Array, error) {
|
|
global, err := w.gather(local, root)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if w.rank != root {
|
|
return nil, nil
|
|
}
|
|
return global, nil
|
|
}
|
|
|
|
// gather is Gather's body; AllGather shares it with a fixed root.
|
|
func (w *World) gather(local *core.Array, root int) (*core.Array, error) {
|
|
if err := w.status(); err != nil {
|
|
return nil, err
|
|
}
|
|
if err := checkRoot(root, w.size); err != nil {
|
|
return nil, w.fail(err)
|
|
}
|
|
if local.NDim() == 0 {
|
|
return nil, w.fail(base.Errf("spmd: Gather needs a dimension to raise along"))
|
|
}
|
|
wire, err := encodeArray(nil, local)
|
|
if err != nil {
|
|
return nil, w.fail(err)
|
|
}
|
|
// A world of one rank is already the whole.
|
|
if w.size == 1 {
|
|
return local, nil
|
|
}
|
|
// The pieces travel to the root; the root takes them in rank
|
|
// order, which is the partition's order, whatever order they
|
|
// arrive in. Its own piece sits at its own rank's place in the
|
|
// join.
|
|
if w.rank != root {
|
|
if err := w.sendTo(root, tagGather, wire); err != nil {
|
|
return nil, err
|
|
}
|
|
return nil, nil
|
|
}
|
|
pieces := make([][]byte, 0, w.size)
|
|
for r := range w.size {
|
|
if r == root {
|
|
pieces = append(pieces, wire)
|
|
continue
|
|
}
|
|
data, err := w.recvFrom(r, tagGather)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
pieces = append(pieces, data)
|
|
}
|
|
return w.joinPieces(pieces)
|
|
}
|
|
|
|
// joinPieces builds the global array from the pieces' wire forms: the
|
|
// heads must agree in dtype and in every dimension but the first, and
|
|
// the payloads concatenate in the order given.
|
|
func (w *World) joinPieces(pieces [][]byte) (*core.Array, error) {
|
|
var dt core.Dtype
|
|
var trailing []int
|
|
total := 0
|
|
var payload []byte
|
|
for i, piece := range pieces {
|
|
pdt, shape, err := decodeHead(piece)
|
|
if err != nil {
|
|
return nil, w.fail(err)
|
|
}
|
|
if len(shape) == 0 {
|
|
return nil, w.fail(base.Errf("spmd: piece %d has no dimension to raise along", i))
|
|
}
|
|
if i == 0 {
|
|
dt = pdt
|
|
trailing = shape[1:]
|
|
} else {
|
|
if pdt != dt {
|
|
return nil, w.fail(base.Errf("spmd: piece %d is %s against %s", i, pdt, dt))
|
|
}
|
|
if len(shape) != len(trailing)+1 {
|
|
return nil, w.fail(base.Errf("spmd: piece %d has %d dimensions against %d", i, len(shape), len(trailing)+1))
|
|
}
|
|
for d := range trailing {
|
|
if shape[d+1] != trailing[d] {
|
|
return nil, w.fail(base.Errf("spmd: piece %d disagrees in dimension %d: %d against %d",
|
|
i, d+1, shape[d+1], trailing[d]))
|
|
}
|
|
}
|
|
}
|
|
if shape[0] > math.MaxInt-total {
|
|
return nil, w.fail(base.Errf("spmd: joining the pieces overflows the first dimension at piece %d", i))
|
|
}
|
|
total += shape[0]
|
|
payload = append(payload, piece[2+8*len(shape):]...)
|
|
}
|
|
globalShape := make([]int, 0, len(trailing)+1)
|
|
globalShape = append(globalShape, total)
|
|
globalShape = append(globalShape, trailing...)
|
|
head, err := encodeHead(nil, dt, globalShape)
|
|
if err != nil {
|
|
return nil, w.fail(err)
|
|
}
|
|
got, err := decodeWire(append(head, payload...))
|
|
if err != nil {
|
|
return nil, w.fail(err)
|
|
}
|
|
return got, nil
|
|
}
|
|
|
|
// AllGather raises every dealt piece into the whole on every rank: one
|
|
// world, one array, identical bits everywhere.
|
|
func (w *World) AllGather(local *core.Array) (*core.Array, error) {
|
|
whole, err := w.gather(local, 0)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return w.Broadcast(whole, 0)
|
|
}
|
|
|
|
func checkRoot(root, size int) error {
|
|
if root < 0 || root >= size {
|
|
return base.Errf("spmd: root %d is outside the world of %d ranks", root, size)
|
|
}
|
|
return nil
|
|
}
|