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