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