// Copyright (c) 2026 Petr BalvĂ­n (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") } }