1292 lines
39 KiB
Go
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")
|
||
|
|
}
|
||
|
|
}
|