feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
+248
@@ -0,0 +1,248 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package spmd
|
||||
|
||||
import (
|
||||
"math"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// The arg reductions answer where an extremum lives in the global
|
||||
// array: the same global index the single-array ArgMax and ArgMin
|
||||
// answer, ties resolved by the earliest index, NaN elements skipped
|
||||
// as missing. The shards compare values exactly, so the answer is the
|
||||
// single-array answer's at any world size.
|
||||
|
||||
// AllReduceArgShards answers the global index of the extremum (Max or
|
||||
// Min) of one global array whose canonical pieces the ranks hold,
|
||||
// every rank receiving the same index. Ties keep the earliest global
|
||||
// index, the single-array walk's own rule; NaN elements never win,
|
||||
// and an array with no candidate at all is an error, as the
|
||||
// single-array reduction is.
|
||||
func (w *World) AllReduceArgShards(local *core.Array, span Span, op Op) (int, error) {
|
||||
ans, err := w.reduceArgShards(local, span, op, 0)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
s, err := w.broadcastScalar(ans, 0)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return int(s.Int()), nil
|
||||
}
|
||||
|
||||
// ReduceArgShards is AllReduceArgShards with the answer on the root
|
||||
// alone; every other rank receives nil.
|
||||
func (w *World) ReduceArgShards(local *core.Array, span Span, op Op, root int) (*int, error) {
|
||||
if err := checkRoot(root, w.size); err != nil {
|
||||
return nil, w.fail(err)
|
||||
}
|
||||
ans, err := w.reduceArgShards(local, span, op, root)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if w.rank != root {
|
||||
return nil, nil
|
||||
}
|
||||
i := int(ans.Int())
|
||||
return &i, nil
|
||||
}
|
||||
|
||||
// argCandidate is one rank's extremum candidate: the extremum's value
|
||||
// and its global index. A slab with no candidate (all NaN, or no rows)
|
||||
// carries ok=false.
|
||||
type argCandidate struct {
|
||||
ok bool
|
||||
value float64
|
||||
iv int64 // the exact value for the integer dtypes
|
||||
idx int
|
||||
isInt bool
|
||||
}
|
||||
|
||||
// reduceArgShards runs the arg reduction; the answer exists at the
|
||||
// root.
|
||||
func (w *World) reduceArgShards(local *core.Array, span Span, op Op, root int) (core.Scalar, error) {
|
||||
if err := w.status(); err != nil {
|
||||
return core.Scalar{}, err
|
||||
}
|
||||
if op != Min && op != Max {
|
||||
return core.Scalar{}, w.fail(base.Errf("spmd: arg reductions take Min or Max, got %s", op))
|
||||
}
|
||||
if err := w.checkOneDimSpan(local, span); err != nil {
|
||||
return core.Scalar{}, w.fail(err)
|
||||
}
|
||||
if local.Dtype() == core.Bool || local.Dtype() == core.Complex {
|
||||
return core.Scalar{}, w.fail(base.Errf("spmd: %s of dtype %s has no ordering", op, local.Dtype()))
|
||||
}
|
||||
if span.Global == 0 {
|
||||
return core.Scalar{}, w.fail(base.Errf("spmd: %s of an empty array has no answer", op))
|
||||
}
|
||||
cand := argCandidateFor(local, span, op)
|
||||
if w.rank != root {
|
||||
if err := w.sendTo(root, tagShardsValues, encodeBlockValues(blockValues{
|
||||
first: cand.idx, kind: kindArg, i: []int64{cand.iv, int64(cand.idx)}, f: []float64{cand.value}, oks: []bool{cand.ok, cand.isInt},
|
||||
})); err != nil {
|
||||
return core.Scalar{}, err
|
||||
}
|
||||
return core.Scalar{}, nil
|
||||
}
|
||||
all := make([]argCandidate, w.size)
|
||||
all[w.rank] = cand
|
||||
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 != kindArg || len(bv.i) != 2 || len(bv.f) != 1 || len(bv.oks) != 2 {
|
||||
return core.Scalar{}, w.fail(base.Errf("spmd: arg shards disagree on the payload shape"))
|
||||
}
|
||||
all[r] = argCandidate{ok: bv.oks[0], value: bv.f[0], iv: bv.i[0], idx: int(bv.i[1]), isInt: bv.oks[1]}
|
||||
}
|
||||
greater := op == Max
|
||||
best := -1
|
||||
for r, c := range all {
|
||||
if !c.ok {
|
||||
continue
|
||||
}
|
||||
if best < 0 {
|
||||
best = r
|
||||
continue
|
||||
}
|
||||
b := all[best]
|
||||
better := false
|
||||
if c.isInt == b.isInt {
|
||||
if c.isInt {
|
||||
better = (greater && c.iv > b.iv) || (!greater && c.iv < b.iv)
|
||||
} else {
|
||||
better = (greater && c.value > b.value) || (!greater && c.value < b.value)
|
||||
}
|
||||
} else {
|
||||
return core.Scalar{}, w.fail(base.Errf("spmd: arg shards disagree on the payload kind"))
|
||||
}
|
||||
if better {
|
||||
best = r
|
||||
}
|
||||
// A tie keeps the earlier rank's candidate, whose global index
|
||||
// is the earlier one: the pieces are contiguous and ascend.
|
||||
}
|
||||
if best < 0 {
|
||||
return core.Scalar{}, w.fail(base.Errf("spmd: %s shards have no candidate", op))
|
||||
}
|
||||
winner := all[best]
|
||||
return core.IntScalar(int64(winner.idx)), nil
|
||||
}
|
||||
|
||||
// argCandidateFor finds the slab's own first extremum, the serial
|
||||
// walk's rule: the first strictly better element, NaNs skipped. The
|
||||
// integer dtypes compare exactly in int64, never through float64,
|
||||
// the reason the core keeps its own int64 walk.
|
||||
func argCandidateFor(local *core.Array, span Span, op Op) argCandidate {
|
||||
greater := op == Max
|
||||
best := -1
|
||||
var bv float64
|
||||
var biv int64
|
||||
intExact := false
|
||||
walkInt := func(v int64, i int) {
|
||||
if best < 0 || (greater && v > biv) || (!greater && v < biv) {
|
||||
best, biv, intExact = i, v, true
|
||||
}
|
||||
}
|
||||
walkFloat := func(v float64, i int) {
|
||||
if math.IsNaN(v) {
|
||||
return
|
||||
}
|
||||
if best < 0 || (greater && v > bv) || (!greater && v < bv) {
|
||||
best, bv, intExact = i, v, false
|
||||
}
|
||||
}
|
||||
for i := range local.Len() {
|
||||
switch local.Dtype() {
|
||||
case core.Float:
|
||||
walkFloat(local.FloatAt(i), i)
|
||||
case core.Float32:
|
||||
walkFloat(float64(local.RawFloat32s()[i]), i)
|
||||
case core.Float16:
|
||||
walkFloat(core.HalfToFloat64(local.RawHalves()[i]), i)
|
||||
case core.Int:
|
||||
walkInt(local.RawInts()[i], i)
|
||||
case core.Int8:
|
||||
walkInt(int64(local.RawInt8s()[i]), i)
|
||||
case core.Uint8:
|
||||
walkInt(int64(local.RawUint8s()[i]), i)
|
||||
case core.Int16:
|
||||
walkInt(int64(local.RawInt16s()[i]), i)
|
||||
case core.Uint16:
|
||||
walkInt(int64(local.RawUint16s()[i]), i)
|
||||
case core.Int32:
|
||||
walkInt(int64(local.RawInt32s()[i]), i)
|
||||
case core.Uint32:
|
||||
walkInt(int64(local.RawUint32s()[i]), i)
|
||||
}
|
||||
}
|
||||
if best < 0 {
|
||||
return argCandidate{}
|
||||
}
|
||||
if intExact {
|
||||
return argCandidate{ok: true, iv: biv, idx: span.Lo + best, isInt: true}
|
||||
}
|
||||
return argCandidate{ok: true, value: bv, idx: span.Lo + best}
|
||||
}
|
||||
|
||||
// AllReduceArgSortShards answers the global permutation that sorts the
|
||||
// whole array ascending: an Int array of global indices, the
|
||||
// single-array ArgSort's own answer with its own tie and NaN
|
||||
// placement, on every rank. The shards' values gather in global order
|
||||
// and the core sorts them, so the permutation is the single-array
|
||||
// one by construction.
|
||||
func (w *World) AllReduceArgSortShards(local *core.Array, span Span) (*core.Array, error) {
|
||||
ans, err := w.reduceArgSortShards(local, span, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return w.Broadcast(ans, 0)
|
||||
}
|
||||
|
||||
// ReduceArgSortShards is AllReduceArgSortShards with the answer on the
|
||||
// root alone; every other rank receives nil.
|
||||
func (w *World) ReduceArgSortShards(local *core.Array, span Span, root int) (*core.Array, error) {
|
||||
if err := checkRoot(root, w.size); err != nil {
|
||||
return nil, w.fail(err)
|
||||
}
|
||||
return w.reduceArgSortShards(local, span, root)
|
||||
}
|
||||
|
||||
func (w *World) reduceArgSortShards(local *core.Array, span Span, root int) (*core.Array, error) {
|
||||
if err := w.status(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := w.checkOneDimSpan(local, span); err != nil {
|
||||
return nil, w.fail(err)
|
||||
}
|
||||
switch local.Dtype() {
|
||||
case core.Int, core.Float32, core.Float16, core.Float:
|
||||
default:
|
||||
return nil, w.fail(base.Errf("spmd: ArgSort of dtype %s is not supported; convert with Astype", local.Dtype()))
|
||||
}
|
||||
gathered, err := w.gather(local, root)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if w.rank != root {
|
||||
return nil, nil
|
||||
}
|
||||
return core.ArgSort(gathered)
|
||||
}
|
||||
|
||||
// nanValue is the quiet NaN the no-candidate tests build their
|
||||
// fixtures from.
|
||||
var nanValue = math.NaN()
|
||||
Reference in New Issue
Block a user