Files
tensor/spmd/reduce.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

1292 lines
39 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"encoding/binary"
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Op names a reduction's arithmetic.
type Op uint8
const (
// Sum adds: shard partials through the balanced tree, rank arrays
// in rank order.
Sum Op = iota
// Min takes the smallest; NaN never wins.
Min
// Max takes the largest; NaN never wins.
Max
// Any reports whether any element is true; Bool arrays only.
Any
// All reports whether every element is true; Bool arrays only.
All
// Prod multiplies: shard partials through the balanced product
// tree, rank arrays in rank order.
Prod
)
func (o Op) String() string {
switch o {
case Sum:
return "sum"
case Min:
return "min"
case Max:
return "max"
case Any:
return "any"
case All:
return "all"
case Prod:
return "prod"
}
return "?"
}
// The payload kinds a shard's contribution can travel as. The kind is
// part of the wire, so the root combines exactly what the ranks
// folded, never a guess from the dtype.
const (
kindNone byte = 0 // a rank whose piece is empty
kindSumF64 byte = 1 // one float64 fold per block
kindSumC128 byte = 2 // one complex fold per block
kindSumI64 byte = 3 // one exact int64 sum per block
kindExtF64 byte = 4 // one float64 extremum and candidate flag per block
kindExtI64 byte = 5 // one exact int64 extremum per block
kindFlag byte = 6 // one 0/1 flag for the whole piece
kindProdF64 byte = 7 // one float64 product per block
kindProdF32 byte = 8 // one native float32 product per block
kindProdHalf byte = 9 // one exact half product per block, in float64
kindProdI64 byte = 10 // one exact int64 product per block
kindArg byte = 11 // one extremum candidate: value, index and flags
)
// intPayload is the integer element type a shard folds exactly.
type intPayload interface {
int64 | int8 | uint8 | int16 | uint16 | int32 | uint32
}
// AllReduceShards answers the reduction of one global array whose
// canonical pieces the ranks hold, every rank receiving the answer.
// Each rank folds its own blocks with the single block fold, the block
// values gather at rank 0, and core's own tree and extremum rules
// combine them: the answer is the single-array reduction's exact bits,
// whatever the world's size, whatever the order the frames arrive in.
// Any and All answer the integer-class scalar, 1 or 0.
func (w *World) AllReduceShards(local *core.Array, span Span, op Op) (core.Scalar, error) {
ans, err := w.reduceShards(local, span, op, 0)
if err != nil {
return core.Scalar{}, err
}
return w.broadcastScalar(ans, 0)
}
// ReduceShards is AllReduceShards with the answer on the root alone;
// every other rank receives nil.
func (w *World) ReduceShards(local *core.Array, span Span, op Op, root int) (*core.Scalar, error) {
if err := checkRoot(root, w.size); err != nil {
return nil, w.fail(err)
}
ans, err := w.reduceShards(local, span, op, root)
if err != nil {
return nil, err
}
if w.rank != root {
return nil, nil
}
return &ans, nil
}
// blockValues is one rank's contribution to a sharded reduction: the
// index of its first block, how many blocks it covers, the payload
// kind, and the values.
type blockValues struct {
first int
kind byte
f []float64
f32 []float32
c []complex128
i []int64
oks []bool
}
// blocks reports how many partition blocks the piece covers.
func (bv blockValues) blocks() int {
switch bv.kind {
case kindSumF64, kindSumI64:
return len(bv.f) + len(bv.i)
case kindSumC128:
return len(bv.c)
case kindExtF64:
return len(bv.oks)
case kindExtI64:
return len(bv.i)
case kindProdF64:
return len(bv.f)
case kindProdF32:
return len(bv.f32)
case kindProdHalf, kindProdI64:
return len(bv.f) + len(bv.i)
}
return 0
}
// spanBlocks names the partition blocks the span covers: [first, last).
// A span that the partition gave out always sits on block boundaries.
func spanBlocks(span Span) (first, last, parts int) {
parts = core.FoldParts(span.Global)
first, last = -1, -1
for c := 0; c <= parts; c++ {
b := core.FoldBoundary(span.Global, c)
if b == span.Lo && first < 0 {
first = c
}
if b == span.Hi {
last = c
}
}
return first, last, parts
}
// reduceShards runs the sharded reduction; the answer exists at the
// root.
func (w *World) reduceShards(local *core.Array, span Span, op Op, root int) (core.Scalar, error) {
if err := w.status(); err != nil {
return core.Scalar{}, err
}
if err := w.checkSpan(span, local); err != nil {
return core.Scalar{}, w.fail(err)
}
if span.Global == 0 {
// An empty global axis answers without an exchange: the
// single-array reduction answers the same constants, so the
// contract holds at the degenerate length too. The dtype
// validates first, exactly as a non-empty array would.
switch op {
case Prod:
if dt := local.Dtype(); dt == core.Bool || dt == core.Int8 || dt == core.Uint8 || dt == core.Int16 || dt == core.Uint16 || dt == core.Int32 || dt == core.Uint32 || dt == core.Complex {
return core.Scalar{}, w.fail(base.Errf("spmd: prod of dtype %s is not supported; convert with Astype", dt))
}
case Min, Max:
if dt := local.Dtype(); dt == core.Bool || dt == core.Complex {
return core.Scalar{}, w.fail(base.Errf("spmd: %s of dtype %s has no ordering", op, dt))
}
case Any, All:
if dt := local.Dtype(); dt != core.Bool {
return core.Scalar{}, w.fail(base.Errf("spmd: %s needs Bool elements, got %s", op, dt))
}
}
switch op {
case Sum:
switch local.Dtype() {
case core.Complex:
return core.ComplexScalar(0), nil
case core.Int, core.Int8, core.Uint8, core.Int16, core.Uint16,
core.Int32, core.Uint32, core.Bool:
return core.IntScalar(0), nil
default:
return core.FloatScalar(0), nil
}
case Any:
return core.IntScalar(0), nil
case All:
return core.IntScalar(1), nil
case Prod:
// The empty product is the multiplicative identity, the
// answer the single-array product keeps.
if local.Dtype() == core.Int {
return core.IntScalar(1), nil
}
return core.FloatScalar(1), nil
default:
return core.Scalar{}, w.fail(base.Errf("spmd: %s of an empty array has no answer", op))
}
}
bv, err := w.foldBlocks(local, span, op)
if err != nil {
return core.Scalar{}, w.fail(err)
}
if w.rank != root {
if err := w.sendTo(root, tagShardsValues, encodeBlockValues(bv)); err != nil {
return core.Scalar{}, err
}
return core.Scalar{}, nil
}
all := make([]blockValues, w.size)
all[w.rank] = bv
for r := range w.size {
if r == root {
continue
}
data, err := w.recvFrom(r, tagShardsValues)
if err != nil {
return core.Scalar{}, err
}
if all[r], err = decodeBlockValues(data); err != nil {
return core.Scalar{}, w.fail(err)
}
}
return w.combineShards(all, span, op)
}
// foldBlocks computes the rank's own blocks' values with the same
// per-block arithmetic the single-array fold uses, dtype by dtype.
func (w *World) foldBlocks(local *core.Array, span Span, op Op) (blockValues, error) {
first, last, _ := spanBlocks(span)
if first < 0 || last < 0 || first > last {
return blockValues{}, base.Errf("spmd: rank %d's span [%d, %d) is not on the partition's block boundaries",
w.rank, span.Lo, span.Hi)
}
bv := blockValues{first: first}
if first == last {
return bv, nil
}
n := span.Global
count := last - first
blockF64 := func(c int) []float64 {
return sliceBlock(local.RawFloats()[:local.Len()], n, first, c)
}
switch op {
case Sum:
switch local.Dtype() {
case core.Float:
bv.kind = kindSumF64
bv.f = make([]float64, count)
for c := first; c < last; c++ {
bv.f[c-first] = core.FoldRange(blockF64(c))
}
case core.Float32:
bv.kind = kindSumF64
bv.f = make([]float64, count)
for c := first; c < last; c++ {
bv.f[c-first] = core.FoldRangeF32(sliceBlock(local.RawFloat32s()[:local.Len()], n, first, c))
}
case core.Float16:
bv.kind = kindSumF64
bv.f = make([]float64, count)
for c := first; c < last; c++ {
bv.f[c-first] = core.FoldRangeF16(sliceBlock(local.RawHalves()[:local.Len()], n, first, c))
}
case core.Complex:
bv.kind = kindSumC128
bv.c = make([]complex128, count)
for c := first; c < last; c++ {
bv.c[c-first] = core.FoldRangeC128(sliceBlock(local.RawComplexes()[:local.Len()], n, first, c))
}
case core.Int:
bv.kind = kindSumI64
bv.i = exactSums(local.RawInts()[:local.Len()], n, first, last)
case core.Int8:
bv.kind = kindSumI64
bv.i = exactSums(local.RawInt8s()[:local.Len()], n, first, last)
case core.Uint8:
bv.kind = kindSumI64
bv.i = exactSums(local.RawUint8s()[:local.Len()], n, first, last)
case core.Int16:
bv.kind = kindSumI64
bv.i = exactSums(local.RawInt16s()[:local.Len()], n, first, last)
case core.Uint16:
bv.kind = kindSumI64
bv.i = exactSums(local.RawUint16s()[:local.Len()], n, first, last)
case core.Int32:
bv.kind = kindSumI64
bv.i = exactSums(local.RawInt32s()[:local.Len()], n, first, last)
case core.Uint32:
bv.kind = kindSumI64
bv.i = exactSums(local.RawUint32s()[:local.Len()], n, first, last)
case core.Bool:
bv.kind = kindSumI64
src := local.RawBools()[:local.Len()]
bv.i = make([]int64, count)
for c := first; c < last; c++ {
var s int64
for _, v := range sliceBlock(src, n, first, c) {
if v {
s++
}
}
bv.i[c-first] = s
}
default:
return blockValues{}, base.Errf("spmd: sum of dtype %s is not a reduction", local.Dtype())
}
case Min, Max:
greater := op == Max
switch local.Dtype() {
case core.Float:
bv.kind = kindExtF64
bv.f = make([]float64, count)
bv.oks = make([]bool, count)
for c := first; c < last; c++ {
bv.f[c-first], bv.oks[c-first] = core.ExtremeRange(blockF64(c), greater)
}
case core.Float32:
bv.kind = kindExtF64
bv.f = make([]float64, count)
bv.oks = make([]bool, count)
for c := first; c < last; c++ {
bv.f[c-first], bv.oks[c-first] = core.ExtremeRangeF32(sliceBlock(local.RawFloat32s()[:local.Len()], n, first, c), greater)
}
case core.Float16:
bv.kind = kindExtF64
bv.f = make([]float64, count)
bv.oks = make([]bool, count)
for c := first; c < last; c++ {
bv.f[c-first], bv.oks[c-first] = core.ExtremeRangeF16(sliceBlock(local.RawHalves()[:local.Len()], n, first, c), greater)
}
case core.Int:
bv.kind = kindExtI64
bv.i = exactExtremes(local.RawInts()[:local.Len()], n, first, last, greater)
case core.Int8:
bv.kind = kindExtI64
bv.i = exactExtremes(local.RawInt8s()[:local.Len()], n, first, last, greater)
case core.Uint8:
bv.kind = kindExtI64
bv.i = exactExtremes(local.RawUint8s()[:local.Len()], n, first, last, greater)
case core.Int16:
bv.kind = kindExtI64
bv.i = exactExtremes(local.RawInt16s()[:local.Len()], n, first, last, greater)
case core.Uint16:
bv.kind = kindExtI64
bv.i = exactExtremes(local.RawUint16s()[:local.Len()], n, first, last, greater)
case core.Int32:
bv.kind = kindExtI64
bv.i = exactExtremes(local.RawInt32s()[:local.Len()], n, first, last, greater)
case core.Uint32:
bv.kind = kindExtI64
bv.i = exactExtremes(local.RawUint32s()[:local.Len()], n, first, last, greater)
default:
return blockValues{}, base.Errf("spmd: %s of dtype %s has no ordering", op, local.Dtype())
}
case Prod:
switch local.Dtype() {
case core.Float:
bv.kind = kindProdF64
bv.f = make([]float64, count)
for c := first; c < last; c++ {
bv.f[c-first] = core.FoldProd(sliceBlock(local.RawFloats()[:local.Len()], n, first, c))
}
case core.Float32:
bv.kind = kindProdF32
bv.f32 = make([]float32, count)
for c := first; c < last; c++ {
bv.f32[c-first] = core.FoldProdF32(sliceBlock(local.RawFloat32s()[:local.Len()], n, first, c))
}
case core.Float16:
bv.kind = kindProdHalf
bv.f = make([]float64, count)
for c := first; c < last; c++ {
bv.f[c-first] = core.FoldProdF16(sliceBlock(local.RawHalves()[:local.Len()], n, first, c))
}
case core.Int:
bv.kind = kindProdI64
bv.i = make([]int64, count)
for c := first; c < last; c++ {
bv.i[c-first] = core.FoldProdI64(sliceBlock(local.RawInts()[:local.Len()], n, first, c))
}
default:
return blockValues{}, base.Errf("spmd: prod of dtype %s is not supported; convert with Astype", local.Dtype())
}
case Any, All:
if local.Dtype() != core.Bool {
return blockValues{}, base.Errf("spmd: %s needs Bool elements, got %s", op, local.Dtype())
}
bv.kind = kindFlag
flag := op == All
for _, v := range local.RawBools()[:local.Len()] {
if op == Any && v {
flag = true
break
}
if op == All && !v {
flag = false
break
}
}
bv.i = []int64{bToF(flag)}
default:
return blockValues{}, base.Errf("spmd: unknown reduction op %s", op)
}
return bv, nil
}
// sliceBlock is block c of the canonical partition of n elements, seen
// inside a slab whose first block is first.
func sliceBlock[T any](src []T, n, first, c int) []T {
baseIdx := core.FoldBoundary(n, first)
return src[core.FoldBoundary(n, c)-baseIdx : core.FoldBoundary(n, c+1)-baseIdx]
}
// exactSums folds each of the piece's blocks into an exact int64 sum.
func exactSums[T intPayload](src []T, n, first, last int) []int64 {
out := make([]int64, last-first)
for c := first; c < last; c++ {
var s int64
for _, v := range sliceBlock(src, n, first, c) {
s += int64(v)
}
out[c-first] = s
}
return out
}
// exactExtremes folds each of the piece's blocks into an exact int64
// extremum: the values compare exactly, so the combine needs no flags.
func exactExtremes[T intPayload](src []T, n, first, last int, greater bool) []int64 {
out := make([]int64, last-first)
for c := first; c < last; c++ {
blk := sliceBlock(src, n, first, c)
best := int64(blk[0])
for _, v := range blk[1:] {
if w := int64(v); (greater && w > best) || (!greater && w < best) {
best = w
}
}
out[c-first] = best
}
return out
}
// payloadKind names the kind every non-empty piece carries; the pieces
// agree because the program runs one op on one dtype, and the root, whose
// own piece may be empty, reads it from the first rank that folded
// anything.
func payloadKind(all []blockValues) (byte, error) {
kind := kindNone
for _, bv := range all {
if bv.kind == kindNone {
continue
}
if kind == kindNone {
kind = bv.kind
} else if bv.kind != kind {
return 0, base.Errf("spmd: the shards disagree on the payload kind: %d against %d", bv.kind, kind)
}
}
return kind, nil
}
// combineShards lays the ranks' pieces over the whole partition and
// combines them with the same functions the single-array reduction
// combines its chunks with.
func (w *World) combineShards(all []blockValues, span Span, op Op) (core.Scalar, error) {
first, last, parts := spanBlocks(span)
if first < 0 || last < 0 {
return core.Scalar{}, base.Errf("spmd: the root's span is not on the partition's boundaries")
}
kind, kindErr := payloadKind(all)
if kindErr != nil {
return core.Scalar{}, kindErr
}
if op != Any && op != All {
covered := 0
for _, bv := range all {
covered += bv.blocks()
}
if covered != parts {
return core.Scalar{}, base.Errf("spmd: the shards' %d blocks do not cover the partition's %d", covered, parts)
}
// The root joins in rank order, which is the partition's
// order only when every piece starts at its own rank's first
// block; a piece naming any other start would scramble the
// tree's addends into a wrong number returned as success.
for r, bv := range all {
if bv.kind == kindNone || bv.kind == kindFlag {
continue
}
if bv.first != r*parts/w.size {
return core.Scalar{}, base.Errf("spmd: rank %d's piece starts at block %d, the partition puts it at %d", r, bv.first, r*parts/w.size)
}
}
}
switch op {
case Sum:
switch kind {
case kindSumF64:
vals := make([]float64, 0, parts)
for _, bv := range all {
vals = append(vals, bv.f...)
}
return core.FloatScalar(core.TreeSum(vals)), nil
case kindSumC128:
vals := make([]complex128, 0, parts)
for _, bv := range all {
vals = append(vals, bv.c...)
}
return core.ComplexScalar(core.TreeSum(vals)), nil
case kindSumI64:
vals := make([]int64, 0, parts)
for _, bv := range all {
vals = append(vals, bv.i...)
}
return core.IntScalar(core.TreeSum(vals)), nil
default:
return core.Scalar{}, base.Errf("spmd: sum shards disagree on the payload kind")
}
case Min, Max:
greater := op == Max
switch kind {
case kindExtF64:
vals := make([]float64, 0, parts)
oks := make([]bool, 0, parts)
for _, bv := range all {
vals = append(vals, bv.f...)
oks = append(oks, bv.oks...)
}
return core.FloatScalar(core.CombineExtrema(vals, oks, greater)), nil
case kindExtI64:
var best int64
have := false
for _, bv := range all {
for _, v := range bv.i {
if !have || (greater && v > best) || (!greater && v < best) {
best, have = v, true
}
}
}
return core.IntScalar(best), nil
default:
return core.Scalar{}, base.Errf("spmd: %s shards disagree on the payload kind", op)
}
case Prod:
switch kind {
case kindProdF64:
vals := make([]float64, 0, parts)
for _, bv := range all {
vals = append(vals, bv.f...)
}
return core.FloatScalar(core.TreeProd(vals)), nil
case kindProdF32:
vals := make([]float32, 0, parts)
for _, bv := range all {
vals = append(vals, bv.f32...)
}
return core.FloatScalar(float64(core.TreeProd(vals))), nil
case kindProdHalf:
vals := make([]float64, 0, parts)
for _, bv := range all {
vals = append(vals, bv.f...)
}
return core.FloatScalar(core.TreeProdHalf(vals)), nil
case kindProdI64:
vals := make([]int64, 0, parts)
for _, bv := range all {
vals = append(vals, bv.i...)
}
return core.IntScalar(core.TreeProd(vals)), nil
default:
return core.Scalar{}, base.Errf("spmd: prod shards disagree on the payload kind")
}
case Any, All:
flag := op == All
for _, bv := range all {
if bv.kind != kindFlag || len(bv.i) == 0 {
continue
}
v := bv.i[0] != 0
if op == Any && v {
flag = true
}
if op == All && !v {
flag = false
}
}
return core.IntScalar(bToF(flag)), nil
}
return core.Scalar{}, base.Errf("spmd: unknown reduction op %s", op)
}
// bToF carries a boolean in the int64 slot the integer-class scalar
// answers through.
func bToF(b bool) int64 {
if b {
return 1
}
return 0
}
// The shard contribution's wire form: first block index, payload kind,
// block count, then the values. Every field is checked against every
// other on the way back in.
func encodeBlockValues(bv blockValues) []byte {
buf := binary.LittleEndian.AppendUint64(nil, uint64(bv.first))
buf = append(buf, bv.kind)
buf = binary.LittleEndian.AppendUint64(buf, uint64(len(bv.f)))
buf = binary.LittleEndian.AppendUint64(buf, uint64(len(bv.f32)))
buf = binary.LittleEndian.AppendUint64(buf, uint64(len(bv.c)))
buf = binary.LittleEndian.AppendUint64(buf, uint64(len(bv.i)))
for _, v := range bv.f {
buf = binary.LittleEndian.AppendUint64(buf, math.Float64bits(v))
}
for _, v := range bv.c {
buf = binary.LittleEndian.AppendUint64(buf, math.Float64bits(real(v)))
buf = binary.LittleEndian.AppendUint64(buf, math.Float64bits(imag(v)))
}
for _, v := range bv.i {
buf = binary.LittleEndian.AppendUint64(buf, uint64(v))
}
for _, v := range bv.f32 {
buf = binary.LittleEndian.AppendUint32(buf, math.Float32bits(v))
}
for _, ok := range bv.oks {
b := byte(0)
if ok {
b = 1
}
buf = append(buf, b)
}
return buf
}
func decodeBlockValues(data []byte) (blockValues, error) {
const head = 8 + 1 + 8*4
if len(data) < head {
return blockValues{}, base.Errf("spmd: a shard contribution of %d bytes is shorter than its head", len(data))
}
// Every count is bounded as its unsigned wire value, before any
// signed conversion: a top-bit count would come back negative and
// sneak past an upper-bound check.
firstU := binary.LittleEndian.Uint64(data)
if firstU > math.MaxInt {
return blockValues{}, base.Errf("spmd: a shard contribution names a first block of %d", firstU)
}
bv := blockValues{first: int(firstU)}
bv.kind = data[8]
nf, n32, nc, ni, err := decodeCounts(data[9:])
if err != nil {
return blockValues{}, err
}
data = data[head:]
want := (nf+ni)*8 + n32*4 + nc*16
if bv.kind == kindExtF64 {
want += nf
}
if bv.kind == kindArg {
want += 2
}
if len(data) != want {
return blockValues{}, base.Errf("spmd: a shard contribution carries %d value bytes for %d floats, %d float32s, %d complexes and %d integers",
len(data), nf, n32, nc, ni)
}
bv.f = make([]float64, nf)
for i := range bv.f {
bv.f[i] = math.Float64frombits(binary.LittleEndian.Uint64(data[i*8:]))
}
data = data[nf*8:]
bv.c = make([]complex128, nc)
for i := range bv.c {
re := math.Float64frombits(binary.LittleEndian.Uint64(data[i*16:]))
im := math.Float64frombits(binary.LittleEndian.Uint64(data[i*16+8:]))
bv.c[i] = complex(re, im)
}
data = data[nc*16:]
bv.i = make([]int64, ni)
for i := range bv.i {
bv.i[i] = int64(binary.LittleEndian.Uint64(data[i*8:]))
}
data = data[ni*8:]
bv.f32 = make([]float32, n32)
for i := range bv.f32 {
bv.f32[i] = math.Float32frombits(binary.LittleEndian.Uint32(data[i*4:]))
}
data = data[n32*4:]
if bv.kind == kindExtF64 {
bv.oks = make([]bool, nf)
for i := range bv.oks {
switch data[i] {
case 0:
case 1:
bv.oks[i] = true
default:
return blockValues{}, base.Errf("spmd: a shard's candidate flag %d is neither 0 nor 1", data[i])
}
}
}
if bv.kind == kindArg {
bv.oks = make([]bool, 2)
for i := range bv.oks {
switch data[i] {
case 0:
case 1:
bv.oks[i] = true
default:
return blockValues{}, base.Errf("spmd: a shard's candidate flag %d is neither 0 nor 1", data[i])
}
}
}
return bv, nil
}
// decodeCounts reads the four block counts as unsigned wire values
// and bounds each before any signed conversion: a top-bit count would
// come back negative and sneak past an upper-bound check.
func decodeCounts(data []byte) (nf, n32, nc, ni int, err error) {
raw := [4]uint64{
binary.LittleEndian.Uint64(data[0:8]),
binary.LittleEndian.Uint64(data[8:16]),
binary.LittleEndian.Uint64(data[16:24]),
binary.LittleEndian.Uint64(data[24:32]),
}
for k, u := range raw {
if u > maxBlocks {
return 0, 0, 0, 0, base.Errf("spmd: a shard contribution names an impossible count %d in slot %d", u, k)
}
}
return int(raw[0]), int(raw[1]), int(raw[2]), int(raw[3]), nil
}
// maxBlocks bounds the counts a contribution may name before anything
// is allocated: the partition itself never exceeds this.
const maxBlocks = 1 << 12
// The scalar answer travels as one fixed frame: a kind byte, then the
// int64, float64 and complex slots, each little-endian. Only the slot
// the kind names is meaningful; every slot always travels, so the form
// has no variable length to get wrong.
const (
scalarKindInt byte = 1
scalarKindFloat byte = 2
scalarKindComplex byte = 3
)
const scalarFrameLen = 1 + 8 + 8 + 16
func encodeScalar(s core.Scalar) []byte {
buf := make([]byte, scalarFrameLen)
switch {
case s.IsComplex():
buf[0] = scalarKindComplex
binary.LittleEndian.PutUint64(buf[17:], math.Float64bits(real(s.Complex())))
binary.LittleEndian.PutUint64(buf[25:], math.Float64bits(imag(s.Complex())))
case s.IsFloat():
buf[0] = scalarKindFloat
binary.LittleEndian.PutUint64(buf[9:], math.Float64bits(s.Float()))
default:
buf[0] = scalarKindInt
binary.LittleEndian.PutUint64(buf[1:], uint64(s.Int()))
}
return buf
}
func sameShape(a, b []int) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
func decodeScalar(data []byte) (core.Scalar, error) {
if len(data) != scalarFrameLen {
return core.Scalar{}, base.Errf("spmd: a scalar answer of %d bytes is not the fixed %d", len(data), scalarFrameLen)
}
switch data[0] {
case scalarKindInt:
return core.IntScalar(int64(binary.LittleEndian.Uint64(data[1:]))), nil
case scalarKindFloat:
return core.FloatScalar(math.Float64frombits(binary.LittleEndian.Uint64(data[9:]))), nil
case scalarKindComplex:
re := math.Float64frombits(binary.LittleEndian.Uint64(data[17:]))
im := math.Float64frombits(binary.LittleEndian.Uint64(data[25:]))
return core.ComplexScalar(complex(re, im)), nil
}
return core.Scalar{}, base.Errf("spmd: a scalar answer of kind %d is none of the library's", data[0])
}
// broadcastScalar carries the root's answer to every rank.
func (w *World) broadcastScalar(s core.Scalar, root int) (core.Scalar, error) {
if w.size == 1 {
return s, nil
}
if w.rank == root {
wire := encodeScalar(s)
for r := range w.size {
if r == root {
continue
}
if err := w.sendTo(r, tagShardsWhole, wire); err != nil {
return core.Scalar{}, err
}
}
return s, nil
}
data, err := w.recvFrom(root, tagShardsWhole)
if err != nil {
return core.Scalar{}, err
}
got, err := decodeScalar(data)
if err != nil {
return core.Scalar{}, w.fail(err)
}
return got, nil
}
// Reduce folds the ranks' same-shaped arrays together and answers on
// the root alone: element e of the answer is the fold of the ranks'
// elements e in rank index order, left to right. The fold's order is
// the program's own rank order, never the arrival order, so the same
// program and data answer the same bits on one machine and across a
// cluster. Every other rank receives nil.
func (w *World) Reduce(a *core.Array, op Op, 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 err := checkArrayOp(a, op); err != nil {
return nil, w.fail(err)
}
if w.size == 1 {
return a, nil
}
k := w.size
n := a.Len()
bounds := func(c int) (int, int) { return c * n / k, (c + 1) * n / k }
// Every rank hands its chunk c to the rank that owns c.
for c := range k {
if c == w.rank {
continue
}
lo, hi := bounds(c)
wire, err := encodePart(nil, a, []int{hi - lo}, lo, hi-lo)
if err != nil {
return nil, w.fail(err)
}
if err := w.sendTo(c, tagReduce, wire); err != nil {
return nil, err
}
}
// The fold runs as the chunks land: the rank's own chunk seeds the
// accumulator, rank 0's chunk replaces it when it arrives, and
// every later chunk joins the running accumulator at its rank's
// turn. The chain of operands is chunk 0 left to chunk k-1 right,
// exactly the rank order it always was, so the bits cannot move,
// while folding one chunk overlaps the transfer of the next and
// no rank piles the whole exchange up before it computes.
lo, hi := bounds(w.rank)
ownWire, err := encodePart(nil, a, []int{hi - lo}, lo, hi-lo)
if err != nil {
return nil, w.fail(err)
}
own, err := decodeWire(ownWire)
if err != nil {
return nil, w.fail(err)
}
acc := own
for i := range k {
chunk := own
if i != w.rank {
data, rerr := w.recvFrom(i, tagReduce)
if rerr != nil {
return nil, rerr
}
chunk, rerr = decodeWire(data)
if rerr != nil {
return nil, w.fail(rerr)
}
if chunk.Dtype() != own.Dtype() {
return nil, w.fail(base.Errf("spmd: rank %d's %s disagrees with rank %d's %s",
i, chunk.Dtype(), w.rank, own.Dtype()))
}
}
if i == 0 {
acc = chunk // the chain starts at rank 0's chunk
continue
}
var ferr error
switch op {
case Sum:
acc, ferr = core.Add(acc, chunk)
case Min:
acc, ferr = core.Minimum(acc, chunk)
case Max:
acc, ferr = core.Maximum(acc, chunk)
case Any:
acc, ferr = core.Or(acc, chunk)
case All:
acc, ferr = core.And(acc, chunk)
case Prod:
acc, ferr = core.Mul(acc, chunk)
}
if ferr != nil {
return nil, w.fail(ferr)
}
}
// The folded chunk travels back to the root, which joins the
// world's chunks in rank order and reshapes to the input's shape.
flat, err := w.gather(acc, root)
if err != nil {
return nil, err
}
if w.rank != root {
return nil, nil
}
return core.Reshape(flat, a.Shape()...)
}
// AllReduce is Reduce with the answer on every rank: one world, one
// answer, identical bits everywhere.
func (w *World) AllReduce(a *core.Array, op Op) (*core.Array, error) {
whole, err := w.Reduce(a, op, 0)
if err != nil {
return nil, err
}
return w.Broadcast(whole, 0)
}
// checkArrayOp refuses the combinations that carry no arithmetic
// before any frame moves: Bool has no sum and no ordering, complex has
// no ordering, and Any with All are Bool's alone.
func checkArrayOp(a *core.Array, op Op) error {
dt := a.Dtype()
switch op {
case Sum:
if dt == core.Bool {
return base.Errf("spmd: sum of Bool arrays has no arithmetic; Any and All are Bool's reductions")
}
case Prod:
if dt == core.Bool {
return base.Errf("spmd: prod of Bool arrays has no arithmetic; Any and All are Bool's reductions")
}
case Min, Max:
if dt == core.Bool || dt == core.Complex {
return base.Errf("spmd: %s of dtype %s has no ordering", op, dt)
}
case Any, All:
if dt != core.Bool {
return base.Errf("spmd: %s needs Bool elements, got %s", op, dt)
}
default:
return base.Errf("spmd: unknown reduction op %s", op)
}
return nil
}
// The sharded product, norm and dot extend the reduction families to
// the vector operations the iterative solvers live on. All three carry
// one-dimensional arrays: the canonical partition cuts the axis, the
// shards fold their own blocks with the single block kernels, and the
// combine is the same tree the single-array answers use.
// AllReduceNormShards answers the Lp norm of one global array whose
// canonical pieces the ranks hold, every rank receiving the answer.
// The power sums fold the canonical blocks with the norm fold's own
// per-element arithmetic and combine through the same balanced tree,
// so the answer is the single-array norm's exact bits at any world
// size. p must be finite and positive; the infinity norm is a maximum
// and belongs to Max, not here.
func (w *World) AllReduceNormShards(local *core.Array, span Span, p float64) (core.Scalar, error) {
ans, err := w.reduceNormShards(local, span, p, 0)
if err != nil {
return core.Scalar{}, err
}
return w.broadcastScalar(ans, 0)
}
// ReduceNormShards is AllReduceNormShards with the answer on the root
// alone; every other rank receives nil.
func (w *World) ReduceNormShards(local *core.Array, span Span, p float64, root int) (*core.Scalar, error) {
if err := checkRoot(root, w.size); err != nil {
return nil, w.fail(err)
}
ans, err := w.reduceNormShards(local, span, p, root)
if err != nil {
return nil, err
}
if w.rank != root {
return nil, nil
}
return &ans, nil
}
func (w *World) reduceNormShards(local *core.Array, span Span, p float64, root int) (core.Scalar, error) {
if err := w.status(); err != nil {
return core.Scalar{}, err
}
if math.IsNaN(p) || p <= 0 || math.IsInf(p, 1) {
return core.Scalar{}, w.fail(base.Errf("spmd: Norm shards need a finite positive p, got %v", p))
}
if err := w.checkOneDimSpan(local, span); err != nil {
return core.Scalar{}, w.fail(err)
}
if dt := local.Dtype(); dt != core.Float && dt != core.Float32 && dt != core.Float16 && dt != core.Int {
return core.Scalar{}, w.fail(base.Errf("spmd: norm of dtype %s is not supported; convert with Astype", dt))
}
if span.Global == 0 {
return core.FloatScalar(0), nil
}
sums, err := w.foldNormBlocks(local, span, p)
if err != nil {
return core.Scalar{}, err
}
if w.rank != root {
if err := w.sendTo(root, tagShardsValues, encodeBlockValues(blockValues{first: blockFirst(span), kind: kindSumF64, f: sums})); err != nil {
return core.Scalar{}, err
}
return core.Scalar{}, nil
}
all := make([][]float64, w.size)
all[w.rank] = sums
for r := range w.size {
if r == root {
continue
}
data, err := w.recvFrom(r, tagShardsValues)
if err != nil {
return core.Scalar{}, err
}
bv, err := decodeBlockValues(data)
if err != nil {
return core.Scalar{}, w.fail(err)
}
if bv.kind != kindSumF64 {
return core.Scalar{}, w.fail(base.Errf("spmd: norm shards disagree on the payload kind"))
}
all[r] = bv.f
}
total := make([]float64, 0, core.FoldParts(span.Global))
for _, part := range all {
total = append(total, part...)
}
return core.FloatScalar(core.NormRoot(core.TreeSum(total), p)), nil
}
// foldNormBlocks computes the rank's own blocks' power sums with the
// norm fold's per-element arithmetic, dtype by dtype.
func (w *World) foldNormBlocks(local *core.Array, span Span, p float64) ([]float64, error) {
first, last, _ := spanBlocks(span)
if first < 0 || last < 0 || first > last {
return nil, base.Errf("spmd: rank %d's span [%d, %d) is not on the partition's block boundaries",
w.rank, span.Lo, span.Hi)
}
n := span.Global
count := last - first
sums := make([]float64, count)
for c := first; c < last; c++ {
switch local.Dtype() {
case core.Float:
sums[c-first] = core.FoldNormPower(sliceBlock(local.RawFloats()[:local.Len()], n, first, c), p)
case core.Float32:
sums[c-first] = core.FoldNormPowerF32(sliceBlock(local.RawFloat32s()[:local.Len()], n, first, c), p)
case core.Float16:
sums[c-first] = core.FoldNormPowerF16(sliceBlock(local.RawHalves()[:local.Len()], n, first, c), p)
case core.Int:
sums[c-first] = core.FoldNormPowerI64(sliceBlock(local.RawInts()[:local.Len()], n, first, c), p)
default:
return nil, base.Errf("spmd: norm of dtype %s is not supported; convert with Astype", local.Dtype())
}
}
return sums, nil
}
// blockFirst names the partition block a span starts on.
func blockFirst(span Span) int {
first, _, _ := spanBlocks(span)
return first
}
// AllReduceDotShards answers the dot product of two global arrays the
// ranks hold as the same canonical pieces, every rank receiving the
// answer. The shards fold their own blocks with the single block dot
// kernel and the block values combine through the same balanced tree,
// so the answer is the single-array Dot's exact bits at any world
// size. The two arrays must carry the same dtype.
func (w *World) AllReduceDotShards(x, y *core.Array, span Span) (core.Scalar, error) {
ans, err := w.reduceDotShards(x, y, span, 0)
if err != nil {
return core.Scalar{}, err
}
return w.broadcastScalar(ans, 0)
}
// ReduceDotShards is AllReduceDotShards with the answer on the root
// alone; every other rank receives nil.
func (w *World) ReduceDotShards(x, y *core.Array, span Span, root int) (*core.Scalar, error) {
if err := checkRoot(root, w.size); err != nil {
return nil, w.fail(err)
}
ans, err := w.reduceDotShards(x, y, span, root)
if err != nil {
return nil, err
}
if w.rank != root {
return nil, nil
}
return &ans, nil
}
func (w *World) reduceDotShards(x, y *core.Array, span Span, root int) (core.Scalar, error) {
if err := w.status(); err != nil {
return core.Scalar{}, err
}
if err := w.checkOneDimSpan(x, span); err != nil {
return core.Scalar{}, w.fail(err)
}
if y.NDim() != 1 || y.Len() != span.Len() {
return core.Scalar{}, w.fail(base.Errf("spmd: the second shard leads with %d elements against the span's %d", y.Len(), span.Len()))
}
if x.Dtype() != y.Dtype() {
return core.Scalar{}, w.fail(base.Errf("spmd: Dot shards carry %s against %s; convert with Astype", x.Dtype(), y.Dtype()))
}
switch x.Dtype() {
case core.Int:
case core.Float, core.Float32, core.Float16, core.Complex:
default:
return core.Scalar{}, w.fail(base.Errf("spmd: dot of dtype %s is not supported; convert with Astype", x.Dtype()))
}
if span.Global == 0 {
if x.Dtype() == core.Int {
return core.IntScalar(0), nil
}
if x.Dtype() == core.Complex {
return core.ComplexScalar(0), nil
}
return core.FloatScalar(0), nil
}
bv, err := w.foldDotBlocks(x, y, span)
if err != nil {
return core.Scalar{}, w.fail(err)
}
if w.rank != root {
if err := w.sendTo(root, tagShardsValues, encodeBlockValues(bv)); err != nil {
return core.Scalar{}, err
}
return core.Scalar{}, nil
}
all := make([]blockValues, w.size)
all[w.rank] = bv
for r := range w.size {
if r == root {
continue
}
data, err := w.recvFrom(r, tagShardsValues)
if err != nil {
return core.Scalar{}, err
}
if all[r], err = decodeBlockValues(data); err != nil {
return core.Scalar{}, w.fail(err)
}
}
return w.combineDot(all, span)
}
// foldDotBlocks computes the rank's own blocks' dot values with the
// single block dot kernel, dtype by dtype.
func (w *World) foldDotBlocks(x, y *core.Array, span Span) (blockValues, error) {
first, last, _ := spanBlocks(span)
if first < 0 || last < 0 || first > last {
return blockValues{}, base.Errf("spmd: rank %d's span [%d, %d) is not on the partition's block boundaries",
w.rank, span.Lo, span.Hi)
}
n := span.Global
count := last - first
bv := blockValues{first: first}
if count == 0 {
return bv, nil
}
switch x.Dtype() {
case core.Float:
bv.kind = kindSumF64
bv.f = make([]float64, count)
for c := first; c < last; c++ {
bv.f[c-first] = core.FoldDot(
sliceBlock(x.RawFloats()[:x.Len()], n, first, c),
sliceBlock(y.RawFloats()[:y.Len()], n, first, c))
}
case core.Float32:
bv.kind = kindSumF64
bv.f = make([]float64, count)
for c := first; c < last; c++ {
bv.f[c-first] = core.FoldDotF32(
sliceBlock(x.RawFloat32s()[:x.Len()], n, first, c),
sliceBlock(y.RawFloat32s()[:y.Len()], n, first, c))
}
case core.Float16:
bv.kind = kindSumF64
bv.f = make([]float64, count)
for c := first; c < last; c++ {
bv.f[c-first] = core.FoldDotF16(
sliceBlock(x.RawHalves()[:x.Len()], n, first, c),
sliceBlock(y.RawHalves()[:y.Len()], n, first, c))
}
case core.Complex:
bv.kind = kindSumC128
bv.c = make([]complex128, count)
for c := first; c < last; c++ {
bv.c[c-first] = core.FoldDotC128(
sliceBlock(x.RawComplexes()[:x.Len()], n, first, c),
sliceBlock(y.RawComplexes()[:y.Len()], n, first, c))
}
case core.Int:
bv.kind = kindSumI64
bv.i = make([]int64, count)
for c := first; c < last; c++ {
bv.i[c-first] = core.FoldDotI64(
sliceBlock(x.RawInts()[:x.Len()], n, first, c),
sliceBlock(y.RawInts()[:y.Len()], n, first, c))
}
default:
return blockValues{}, base.Errf("spmd: dot of dtype %s is not supported; convert with Astype", x.Dtype())
}
return bv, nil
}
// combineDot folds the shards' dot values through the balanced tree.
func (w *World) combineDot(all []blockValues, span Span) (core.Scalar, error) {
first, _, parts := spanBlocks(span)
if first < 0 {
return core.Scalar{}, base.Errf("spmd: the root's span is not on the partition's boundaries")
}
covered := 0
for r, bv := range all {
covered += bv.blocks()
if bv.kind == kindNone {
continue
}
if bv.first != r*parts/w.size {
return core.Scalar{}, base.Errf("spmd: rank %d's piece starts at block %d, the partition puts it at %d", r, bv.first, r*parts/w.size)
}
}
if covered != parts {
return core.Scalar{}, base.Errf("spmd: the shards' %d blocks do not cover the partition's %d", covered, parts)
}
kind, kindErr := payloadKind(all)
if kindErr != nil {
return core.Scalar{}, kindErr
}
switch kind {
case kindSumI64:
vals := make([]int64, 0, parts)
for _, bv := range all {
vals = append(vals, bv.i...)
}
return core.IntScalar(core.TreeSum(vals)), nil
case kindSumC128:
vals := make([]complex128, 0, parts)
for _, bv := range all {
vals = append(vals, bv.c...)
}
return core.ComplexScalar(core.TreeSum(vals)), nil
case kindSumF64:
vals := make([]float64, 0, parts)
for _, bv := range all {
vals = append(vals, bv.f...)
}
return core.FloatScalar(core.TreeSum(vals)), nil
default:
return core.Scalar{}, base.Errf("spmd: dot shards disagree on the payload kind")
}
}