// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package spmd import ( "testing" "time" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // The arg reductions answer the single-array walk's global index, ties // by the earliest index, NaNs skipped, and the sharded sort answers // the single-array permutation outright. // TestShardedArgMatchesSingleArray pins ArgMax and ArgMin against the // core's own walk across dtypes, NaNs included, with the tie at the // earliest index. func TestShardedArgMatchesSingleArray(t *testing.T) { for _, size := range []int{1, 3, 5, 8} { for _, gn := range []int{1, 100, 65537} { for dt, whole := range fixtureDtypes(t, gn) { switch dt { case core.Float, core.Float32, core.Float16, core.Int, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32: default: continue } for _, op := range []Op{Max, Min} { var want int var err error if op == Max { want, err = core.ArgMax(whole) } else { want, err = core.ArgMin(whole) } if err != nil { t.Fatal(err) } err = Launch(size, func(w *World) error { span := mustPartition(t, gn, w.Size(), w.Rank()) local := narrowSliceFor(t, whole, span) got, err := w.AllReduceArgShards(local, span, op) if err != nil { return err } if got != want { t.Fatalf("size %d gn %d %s %s: sharded index %d against single-array %d", w.Size(), gn, dt, op, got, want) } return nil }) if err != nil { t.Fatalf("size %d gn %d %s %s: %v", size, gn, dt, op, err) } } } } } } // TestShardedArgNoCandidate: an all-NaN array is an error on every // rank, the single-array walk's own refusal. func TestShardedArgNoCandidate(t *testing.T) { nans := make([]float64, 100) for i := range nans { nans[i] = math2NaN() } whole, err := core.FromFloats(nans, 100) if err != nil { t.Fatal(err) } err = Launch(3, func(w *World) error { span := mustPartition(t, 100, w.Size(), w.Rank()) local := narrowSliceFor(t, whole, span) _, err := w.AllReduceArgShards(local, span, Max) return err }) if err == nil { t.Fatal("an all-NaN array answered an index") } } func math2NaN() float64 { return nanValue } // TestShardedArgSortMatchesSingleArray: the sharded permutation is // the single-array ArgSort's own, values, ties, NaN placement and all. func TestShardedArgSortMatchesSingleArray(t *testing.T) { for _, size := range []int{1, 3, 5, 8} { for _, gn := range []int{1, 100, 65537} { for _, dt := range []core.Dtype{core.Float, core.Float32, core.Float16, core.Int} { whole := fixtureDtypes(t, gn)[dt] want, err := core.ArgSort(whole) if err != nil { t.Fatal(err) } err = Launch(size, func(w *World) error { span := mustPartition(t, gn, w.Size(), w.Rank()) local := narrowSliceFor(t, whole, span) got, err := w.AllReduceArgSortShards(local, span) if err != nil { return err } if got.Len() != want.Len() { t.Fatalf("size %d gn %d %s: permutation of %d against %d", w.Size(), gn, dt, got.Len(), want.Len()) } for i := range want.Len() { if got.RawInts()[i] != want.RawInts()[i] { t.Fatalf("size %d gn %d %s: permutation differs at %d: %d against %d", w.Size(), gn, dt, i, got.RawInts()[i], want.RawInts()[i]) } } return nil }) if err != nil { t.Fatalf("size %d gn %d %s: %v", size, gn, dt, err) } } } } } // TestShardedArgOverTCP runs the arg and the sort over real // connections. func TestShardedArgOverTCP(t *testing.T) { const gn = 65537 whole := fixtureDtypes(t, gn)[core.Float] wantMax, err := core.ArgMax(whole) if err != nil { t.Fatal(err) } wantSort, err := core.ArgSort(whole) if err != nil { t.Fatal(err) } runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error { span := mustPartition(t, gn, w.Size(), w.Rank()) local := narrowSliceFor(t, whole, span) gotMax, err := w.AllReduceArgShards(local, span, Max) if err != nil { return err } if gotMax != wantMax { t.Fatalf("rank %d: arg %d against %d", w.Rank(), gotMax, wantMax) } gotSort, err := w.AllReduceArgSortShards(local, span) if err != nil { return err } for i := range wantSort.Len() { if gotSort.RawInts()[i] != wantSort.RawInts()[i] { t.Fatalf("rank %d: the permutation differs at %d over TCP", w.Rank(), i) } } return nil }) }