Files

286 lines
7.9 KiB
Go
Raw Permalink 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"
)
// 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
}