// 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" ) // 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 }