Files
tensor/spmd/arg_test.go
T

159 lines
4.5 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 (
"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
})
}