Files
tensor/spmd/arg.go
T

249 lines
7.6 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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()