159 lines
4.5 KiB
Go
159 lines
4.5 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
|
|
})
|
|
}
|