Files
tensor/spmd/reduce_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

887 lines
26 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"encoding/binary"
"math"
"strings"
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The sharded reduction tests judge the whole contract in one claim:
// whatever the world's size, the sharded answer carries the
// single-array reduction's exact bits. The single-array answer is
// computed from the same fixture beside the test, so the comparison is
// the contract and not a self-check.
// sparseFixture spreads magnitudes and drops NaNs and an infinity at
// fixed spots, so the folds meet cancellation and the extremum rules
// meet their NaN edges.
func sparseFixture(n int) []float64 {
v := make([]float64, n)
for i := range v {
v[i] = float64((i*6559)%2001-1000) * math.Pow(10, float64(i%7)-3)
if i%401 == 3 {
v[i] = math.NaN()
}
if i%401 == 200 {
v[i] = math.Inf(-1)
}
}
return v
}
// fixtureDtypes builds the same data in every element type the shard
// reductions carry.
func fixtureDtypes(t *testing.T, gn int) map[core.Dtype]*core.Array {
t.Helper()
build := func(a *core.Array, err error) *core.Array {
if err != nil {
t.Fatal(err)
}
return a
}
vals := sparseFixture(gn)
ints := make([]int64, gn)
bools := make([]bool, gn)
for i := range ints {
ints[i] = int64((i*13)%97) - 48
bools[i] = i%3 == 0
}
halves := make([]uint16, gn)
f32s := make([]float32, gn)
for i := range halves {
halves[i] = core.HalfFromFloat64(vals[i])
f32s[i] = float32(vals[i])
}
complexes := make([]complex128, gn)
for i := range complexes {
complexes[i] = complex(vals[i], -vals[i]/2)
}
return map[core.Dtype]*core.Array{
core.Float: build(core.FromFloats(vals, gn)),
core.Float32: build(core.FromFloat32s(f32s, gn)),
core.Float16: build(core.HalvesFromArray(halves, gn)),
core.Complex: build(core.FromComplexes(complexes, gn)),
core.Int: build(core.FromInts(ints, gn)),
core.Int8: build(core.FromInt8s(narrow8s(ints), gn)),
core.Uint8: build(core.FromUint8s(narrowu8s(ints), gn)),
core.Int16: build(core.FromInt16s(narrow16s(ints), gn)),
core.Uint16: build(core.FromUint16s(narrowu16s(ints), gn)),
core.Int32: build(core.FromInt32s(narrow32s(ints), gn)),
core.Uint32: build(core.FromUint32s(narrowu32s(ints), gn)),
core.Bool: build(core.FromBools(bools, gn)),
}
}
func narrow8s(v []int64) []int8 {
out := make([]int8, len(v))
for i := range v {
out[i] = int8(v[i])
}
return out
}
func narrowu8s(v []int64) []uint8 {
out := make([]uint8, len(v))
for i := range v {
out[i] = uint8(v[i])
}
return out
}
func narrow16s(v []int64) []int16 {
out := make([]int16, len(v))
for i := range v {
out[i] = int16(v[i])
}
return out
}
func narrowu16s(v []int64) []uint16 {
out := make([]uint16, len(v))
for i := range v {
out[i] = uint16(v[i])
}
return out
}
func narrow32s(v []int64) []int32 {
out := make([]int32, len(v))
for i := range v {
out[i] = int32(v[i])
}
return out
}
func narrowu32s(v []int64) []uint32 {
out := make([]uint32, len(v))
for i := range v {
out[i] = uint32(v[i])
}
return out
}
// narrowSliceFor deals any dtype's rows off the whole array, keeping
// the bits exactly: what Scatter deals, built without the collectives
// under test.
func narrowSliceFor(t *testing.T, whole *core.Array, span Span) *core.Array {
t.Helper()
wire, err := encodePart(nil, whole, append([]int{span.Len()}, whole.Shape()[1:]...), span.Lo, span.Len())
if err != nil {
t.Fatal(err)
}
a, err := decodeWire(wire)
if err != nil {
t.Fatal(err)
}
return a
}
// scalarBits renders a scalar's exact bits, NaN payloads and signed
// zeros included, so the comparison is bit for bit.
func scalarBits(s core.Scalar) string {
switch {
case s.IsComplex():
c := s.Complex()
return "c" + fBits(real(c)) + "/" + fBits(imag(c))
case s.IsFloat():
return "f" + fBits(s.Float())
default:
return "i" + iFormat(s.Int())
}
}
func fBits(f float64) string { return iFormat(int64(math.Float64bits(f))) }
func iFormat(i int64) string {
if i < 0 {
return "-" + uFormat(-uint64(i))
}
return uFormat(uint64(i))
}
func uFormat(u uint64) string {
if u == 0 {
return "0"
}
var b []byte
for u > 0 {
b = append([]byte{byte('0' + u%10)}, b...)
u /= 10
}
return string(b)
}
// TestShardedSumIsTheSingleArraySum is the headline claim of the whole
// package: shard the data, reduce the shards, and the bits are the
// single array's bits, for every world size, every element type and
// every length that touches the partition's edges.
func TestShardedSumIsTheSingleArraySum(t *testing.T) {
for _, size := range []int{1, 2, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537, 131073, 327681} {
for dt, whole := range fixtureDtypes(t, gn) {
want := core.Sum(whole)
err := Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceShards(local, span, Sum)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("size %d gn %d %s: sharded %s against single-array %s",
w.Size(), gn, dt, scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d %s: %v", size, gn, dt, err)
}
}
}
}
}
// TestShardedExtremaAreTheSingleArrayExtrema pins Min and Max against
// the single-array walk, NaN rules and the all-NaN fallback included.
func TestShardedExtremaAreTheSingleArrayExtrema(t *testing.T) {
for _, size := range []int{1, 2, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537, 131073} {
for _, dt := range []core.Dtype{core.Float, core.Float32, core.Float16, core.Int, core.Int8} {
whole := fixtureDtypes(t, gn)[dt]
for _, op := range []Op{Min, Max} {
want, err := (func() (core.Scalar, error) {
if op == Min {
return core.Min(whole)
}
return core.Max(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.AllReduceShards(local, span, op)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("size %d gn %d %s %s: sharded %s against single-array %s",
w.Size(), gn, dt, op, scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d %s %s: %v", size, gn, dt, op, err)
}
}
}
}
}
// The all-NaN array: every block reports no candidate, so the answer
// is the last block's fallback, the single-array walk's own rule.
const gn = 131073 // two blocks
allNaN := make([]float64, gn)
for i := range allNaN {
allNaN[i] = math.NaN()
}
nans := mk(core.FromFloats(allNaN, gn))
for _, op := range []Op{Min, Max} {
want, err := (func() (core.Scalar, error) {
if op == Min {
return core.Min(nans)
}
return core.Max(nans)
})()
if err != nil {
t.Fatal(err)
}
err = Launch(3, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, nans, span)
got, err := w.AllReduceShards(local, span, op)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("all-NaN %s: sharded %s against single-array %s",
op, scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("all-NaN %s: %v", op, err)
}
}
}
// TestShardedBoolReductions pins Any and All against the single-array
// answers.
func TestShardedBoolReductions(t *testing.T) {
for _, size := range []int{1, 2, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537} {
whole := fixtureDtypes(t, gn)[core.Bool]
anyWant, _ := core.Any(whole)
allWant, _ := core.All(whole)
err := Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
anyGot, err := w.AllReduceShards(local, span, Any)
if err != nil {
return err
}
allGot, err := w.AllReduceShards(local, span, All)
if err != nil {
return err
}
if anyGot.Int() != bToF(anyWant) || allGot.Int() != bToF(allWant) {
t.Fatalf("size %d gn %d: any %d all %d against %v/%v",
w.Size(), gn, anyGot.Int(), allGot.Int(), anyWant, allWant)
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d: %v", size, gn, err)
}
}
}
}
// TestReduceShardsAnswersTheRootAlone pins the root-directed shape of
// the collective: the root holds the answer, nobody else holds
// anything.
func TestReduceShardsAnswersTheRootAlone(t *testing.T) {
const gn = 131074
whole := fixtureDtypes(t, gn)[core.Float]
want := core.Sum(whole)
err := Launch(3, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.ReduceShards(local, span, Sum, 2)
if err != nil {
return err
}
if w.Rank() == 2 {
if got == nil || scalarBits(*got) != scalarBits(want) {
t.Fatalf("the root's sharded sum against the single-array %s", scalarBits(want))
}
return nil
}
if got != nil {
t.Fatalf("rank %d received an answer it should not have", w.Rank())
}
return nil
})
if err != nil {
t.Fatal(err)
}
}
// TestShardsRefuseAForeignSpan: a span the partition did not cut is an
// error naming the canonical boundaries, never a number.
func TestShardsRefuseAForeignSpan(t *testing.T) {
const gn = 65537
whole := fixtureDtypes(t, gn)[core.Float]
err := Launch(3, func(w *World) error {
span := mustPartition(t, gn, 3, w.Rank())
local := narrowSliceFor(t, whole, mustPartition(t, gn, 3, w.Rank()))
if w.Rank() == 1 {
// A plausible but foreign cut: one element over.
span.Lo++
}
_, err := w.AllReduceShards(local, span, Sum)
return err
})
if err == nil || !strings.Contains(err.Error(), "canonical partition") {
t.Fatalf("a foreign span did not name the canonical boundaries: %v", err)
}
}
// TestShardedSumOverTCP runs the headline claim over real connections.
func TestShardedSumOverTCP(t *testing.T) {
const gn = 131073
whole := fixtureDtypes(t, gn)[core.Float]
want := core.Sum(whole)
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)
got, err := w.AllReduceShards(local, span, Sum)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("rank %d: sharded %s against single-array %s", w.Rank(), scalarBits(got), scalarBits(want))
}
return nil
})
}
// TestAllReduceFoldsInRankOrder pins family A: elementwise across the
// ranks' same-shaped arrays, folded in rank index order, which the
// expected answer rebuilds with the core's own pairwise ops.
func TestAllReduceFoldsInRankOrder(t *testing.T) {
for _, size := range []int{1, 2, 3, 5} {
rankArrays := make([]*core.Array, size)
for r := range rankArrays {
vals := make([]float64, 100)
for i := range vals {
vals[i] = float64(r+1) * float64((i*17)%31-15) / 3.0
}
rankArrays[r] = mk(core.FromFloats(vals, 100))
}
want := rankArrays[0]
for _, a := range rankArrays[1:] {
next, err := core.Add(want, a)
if err != nil {
t.Fatal(err)
}
want = next
}
wantMin := rankArrays[0]
for _, a := range rankArrays[1:] {
next, err := core.Minimum(wantMin, a)
if err != nil {
t.Fatal(err)
}
wantMin = next
}
err := Launch(size, func(w *World) error {
got, err := w.AllReduce(rankArrays[w.Rank()], Sum)
if err != nil {
return err
}
if !sameBits(want, got) {
t.Fatalf("size %d: the AllReduce sum differs from the rank-order fold", size)
}
got, err = w.AllReduce(rankArrays[w.Rank()], Min)
if err != nil {
return err
}
if !sameBits(wantMin, got) {
t.Fatalf("size %d: the AllReduce min differs from the rank-order fold", size)
}
return nil
})
if err != nil {
t.Fatalf("size %d: %v", size, err)
}
}
}
// TestShardContributionRefusesTheHostile: the counts the wire names
// are bounded as unsigned values before anything is allocated, so a
// top-bit count is refused rather than turned into a negative length,
// and the candidate flags of an extremum contribution are read back as
// the 0/1 bytes the encoder writes, an honest run and an empty one
// included.
func TestShardContributionRefusesTheHostile(t *testing.T) {
head := make([]byte, 8+1+8*4)
head[8] = kindSumF64
binary.LittleEndian.PutUint64(head[9:], uint64(1)<<63)
if _, err := decodeBlockValues(head); err == nil {
t.Fatal("a top-bit float count decoded")
}
honest := binary.LittleEndian.AppendUint64(nil, 0)
honest = append(honest, kindSumF64)
honest = binary.LittleEndian.AppendUint64(honest, 1)
honest = binary.LittleEndian.AppendUint64(honest, 0)
honest = binary.LittleEndian.AppendUint64(honest, 0)
honest = binary.LittleEndian.AppendUint64(honest, 0)
honest = binary.LittleEndian.AppendUint64(honest, math.Float64bits(1))
bv, err := decodeBlockValues(honest)
if err != nil || len(bv.f) != 1 || bv.f[0] != 1 {
t.Fatalf("an honest contribution was refused: %v", err)
}
extremum := func(flags ...byte) []byte {
buf := binary.LittleEndian.AppendUint64(nil, 0)
buf = append(buf, kindExtF64)
buf = binary.LittleEndian.AppendUint64(buf, uint64(len(flags)))
buf = binary.LittleEndian.AppendUint64(buf, 0)
buf = binary.LittleEndian.AppendUint64(buf, 0)
buf = binary.LittleEndian.AppendUint64(buf, 0)
for range flags {
buf = binary.LittleEndian.AppendUint64(buf, math.Float64bits(1))
}
return append(buf, flags...)
}
if _, err := decodeBlockValues(extremum()); err != nil {
t.Fatalf("an empty flag run was refused: %v", err)
}
bv, err = decodeBlockValues(extremum(1, 0))
if err != nil {
t.Fatalf("honest candidate flags were refused: %v", err)
}
if !bv.oks[0] || bv.oks[1] {
t.Fatalf("the candidate flags came back as %v", bv.oks)
}
if _, err := decodeBlockValues(extremum(1, 2)); err == nil {
t.Fatal("a candidate flag byte of 2 decoded")
}
}
// TestShardedSumAtZeroLength: the degenerate global axis answers the
// single-array reduction's zero at any world size, no exchange needed.
func TestShardedSumAtZeroLength(t *testing.T) {
for _, size := range []int{1, 2, 3, 5} {
empty := mk(core.FromFloats(nil, 0))
want := core.Sum(empty)
err := Launch(size, func(w *World) error {
span := mustPartition(t, 0, w.Size(), w.Rank())
got, err := w.AllReduceShards(empty, span, Sum)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("size %d: sharded %s against single-array %s",
w.Size(), scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("size %d: %v", size, err)
}
}
}
// TestShardedExtremaReachTheRemotePieces pins the candidate flags'
// travel: the extremum living strictly on a non-root rank, without a
// tie to mask it, must win the combine whoever holds it. The saturating
// fixture above cannot tell this, because its extrema repeat in every
// window and the tie rule hands the answer to the root anyway.
func TestShardedExtremaReachTheRemotePieces(t *testing.T) {
// Ascending data: the maximum lives on the last rank alone, and the
// minimum on the root. The root holds its own block here.
const gn = 131073
asc := make([]float64, gn)
for i := range asc {
asc[i] = float64(i)
}
whole := mk(core.FromFloats(asc, gn))
for _, op := range []Op{Min, Max} {
want, err := (func() (core.Scalar, error) {
if op == Min {
return core.Min(whole)
}
return core.Max(whole)
})()
if err != nil {
t.Fatal(err)
}
err = Launch(2, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceShards(local, span, op)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("%s: sharded %s against single-array %s",
op, scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("%s: %v", op, err)
}
}
// A world whose root holds nothing at all (four ranks over three
// blocks) and whose maximum lives only in the middle block: with
// every root-side candidate absent, the combine must still take the
// middle block's value, never the last block's.
const gn2 = 131073
pyr := make([]float64, gn2)
for i := range pyr {
d := i - gn2/2
if d < 0 {
d = -d
}
pyr[i] = -float64(d)
}
peak := mk(core.FromFloats(pyr, gn2))
want, err := core.Max(peak)
if err != nil {
t.Fatal(err)
}
err = Launch(4, func(w *World) error {
span := mustPartition(t, gn2, w.Size(), w.Rank())
local := narrowSliceFor(t, peak, span)
got, err := w.AllReduceShards(local, span, Max)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("peak: sharded %s against single-array %s",
scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatal(err)
}
// A tie in value with different bits: the earlier block's signed zero
// must win, whichever sign each block holds.
const gn3 = 131073 // two blocks
for _, tc := range []struct {
first, second float64
}{
{0, math.Copysign(0, -1)}, // tie keeps the first block's +0
{math.Copysign(0, -1), 0}, // tie keeps the first block's -0
} {
vals := make([]float64, gn3)
for i := range vals {
vals[i] = tc.first
if i >= gn3/2 {
vals[i] = tc.second
}
}
zeros := mk(core.FromFloats(vals, gn3))
want, err := core.Max(zeros)
if err != nil {
t.Fatal(err)
}
err = Launch(2, func(w *World) error {
span := mustPartition(t, gn3, w.Size(), w.Rank())
local := narrowSliceFor(t, zeros, span)
got, err := w.AllReduceShards(local, span, Max)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("tie: sharded max %s against single-array %s",
scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatal(err)
}
}
}
// TestShardedExtremaOverTCP runs the extrema's candidate flags over
// real connections, the pattern of TestShardedSumOverTCP.
func TestShardedExtremaOverTCP(t *testing.T) {
const gn = 131073
pyr := make([]float64, gn)
for i := range pyr {
d := i - gn/2
if d < 0 {
d = -d
}
pyr[i] = -float64(d)
}
peak := mk(core.FromFloats(pyr, gn))
want, err := core.Max(peak)
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, peak, span)
got, err := w.AllReduceShards(local, span, Max)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("rank %d: sharded %s against single-array %s",
w.Rank(), scalarBits(got), scalarBits(want))
}
return nil
})
}
// TestShardedEmptyAxisMatchesTheDtypeRules pins the degenerate length's
// dtype rules against the non-empty path's: a Bool sum counts its trues
// and answers zero, and Any and All keep refusing the non-Bool dtypes.
func TestShardedEmptyAxisMatchesTheDtypeRules(t *testing.T) {
emptyBool := mk(core.FromBools(nil, 0))
if got := core.Sum(emptyBool); got.Int() != 0 {
t.Fatalf("the single-array Bool sum of an empty array: %v", got)
}
err := Launch(3, func(w *World) error {
span := mustPartition(t, 0, w.Size(), w.Rank())
got, err := w.AllReduceShards(emptyBool, span, Sum)
if err != nil {
return err
}
if got.Int() != 0 {
t.Fatalf("the empty Bool sum answered %v", got)
}
if any, err := w.AllReduceShards(emptyBool, span, Any); err != nil || any.Int() != 0 {
t.Fatalf("the empty Bool Any answered %v %v", any, err)
}
if all, err := w.AllReduceShards(emptyBool, span, All); err != nil || all.Int() != 1 {
t.Fatalf("the empty Bool All answered %v %v", all, err)
}
emptyFloat := mk(core.FromFloats(nil, 0))
if _, err := w.AllReduceShards(emptyFloat, span, Any); err == nil {
t.Fatal("the empty Any of a Float shard was accepted")
}
if _, err := w.AllReduceShards(emptyFloat, span, All); err == nil {
t.Fatal("the empty All of a Float shard was accepted")
}
return nil
})
if err != nil {
t.Fatal(err)
}
}
// prodFixture keeps every factor a relative hair away from one, so a
// long product stays finite and meaningful against the single-array
// answer.
func prodFixture(n int) []float64 {
v := make([]float64, n)
for i := range v {
v[i] = 1 + float64(int64(i%2001)-1000)/1e6
}
return v
}
// TestShardedProdIsTheSingleArrayProd: the sharded product carries the
// single-array product's exact bits, the dtype's own per-step rounding
// included.
func TestShardedProdIsTheSingleArrayProd(t *testing.T) {
for _, size := range []int{1, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537, 131073} {
build := func(t *testing.T) map[core.Dtype]*core.Array {
build1 := func(a *core.Array, err error) *core.Array {
if err != nil {
t.Fatal(err)
}
return a
}
vals := prodFixture(gn)
ints := make([]int64, gn)
for i := range ints {
ints[i] = int64(i%5) - 2
}
halves := make([]uint16, gn)
f32s := make([]float32, gn)
for i := range halves {
halves[i] = core.HalfFromFloat64(vals[i])
f32s[i] = float32(vals[i])
}
return map[core.Dtype]*core.Array{
core.Float: build1(core.FromFloats(vals, gn)),
core.Float32: build1(core.FromFloat32s(f32s, gn)),
core.Float16: build1(core.HalvesFromArray(halves, gn)),
core.Int: build1(core.FromInts(ints, gn)),
}
}
for dt, whole := range build(t) {
want, err := core.Prod(whole, 0, false)
if err != nil {
t.Fatal(err)
}
var wantBits string
switch dt {
case core.Int:
wantBits = scalarBits(core.IntScalar(want.RawInts()[0]))
case core.Float16:
wantBits = scalarBits(core.FloatScalar(core.HalfToFloat64(want.RawHalves()[0])))
case core.Float32:
wantBits = scalarBits(core.FloatScalar(float64(want.RawFloat32s()[0])))
default:
wantBits = scalarBits(core.FloatScalar(want.FloatAt(0)))
}
err = Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
local := narrowSliceFor(t, whole, span)
got, err := w.AllReduceShards(local, span, Prod)
if err != nil {
return err
}
if scalarBits(got) != wantBits {
t.Fatalf("size %d gn %d %s: sharded %s against single-array %s",
w.Size(), gn, dt, scalarBits(got), wantBits)
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d %s: %v", size, gn, dt, err)
}
}
}
}
}
// TestShardedNormIsTheSingleArrayNorm pins the power sums and their
// closing against the single-array norm, finite exponents only.
func TestShardedNormIsTheSingleArrayNorm(t *testing.T) {
for _, size := range []int{1, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537, 131073} {
for _, p := range []float64{1, 2, 3.5} {
whole := sparseFixtureOf(gn, core.Float)
want, err := core.Norm(whole, p, 0, false)
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.AllReduceNormShards(local, span, p)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(core.FloatScalar(want.FloatAt(0))) {
t.Fatalf("size %d gn %d p %v: sharded %s against single-array %s",
w.Size(), gn, p, scalarBits(got), scalarBits(core.FloatScalar(want.FloatAt(0))))
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d p %v: %v", size, gn, p, err)
}
}
}
}
}
// sparseFixtureOf is sparseFixture landing in one dtype.
func sparseFixtureOf(gn int, dt core.Dtype) *core.Array {
vals := sparseFixture(gn)
var a *core.Array
var err error
switch dt {
case core.Float32:
f32s := make([]float32, gn)
for i := range f32s {
f32s[i] = float32(vals[i])
}
a, err = core.FromFloat32s(f32s, gn)
default:
a, err = core.FromFloats(vals, gn)
}
if err != nil {
panic(err)
}
return a
}
// TestShardedDotIsTheSingleArrayDot: two equally sharded arrays answer
// the single-array Dot's exact bits.
func TestShardedDotIsTheSingleArrayDot(t *testing.T) {
for _, size := range []int{1, 3, 5, 8} {
for _, gn := range []int{1, 100, 65537, 131073} {
x := fixtureDtypes(t, gn)[core.Float]
yWhole := sparseFixtureOf(gn, core.Float)
want, err := core.Dot(x, yWhole)
if err != nil {
t.Fatal(err)
}
err = Launch(size, func(w *World) error {
span := mustPartition(t, gn, w.Size(), w.Rank())
lx := narrowSliceFor(t, x, span)
ly := narrowSliceFor(t, yWhole, span)
got, err := w.AllReduceDotShards(lx, ly, span)
if err != nil {
return err
}
if scalarBits(got) != scalarBits(want) {
t.Fatalf("size %d gn %d: sharded %s against single-array %s",
w.Size(), gn, scalarBits(got), scalarBits(want))
}
return nil
})
if err != nil {
t.Fatalf("size %d gn %d: %v", size, gn, err)
}
}
}
}
// TestShardedVectorRefusals: the infinity norm belongs to Max, and the
// vector reductions carry 1-D arrays alone.
func TestShardedVectorRefusals(t *testing.T) {
err := Launch(2, func(w *World) error {
span := mustPartition(t, 100, w.Size(), w.Rank())
local := narrowSliceFor(t, sparseFixtureOf(100, core.Float), span)
if _, err := w.AllReduceNormShards(local, span, math.Inf(1)); err == nil {
t.Fatal("the infinity norm was accepted")
}
two, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
if err != nil {
return err
}
span2 := mustPartition(t, 2, w.Size(), w.Rank())
if _, err := w.AllReduceNormShards(two, span2, 2); err == nil {
t.Fatal("a two-dimensional shard was accepted")
}
if _, err := w.AllReduceShards(two, span2, Prod); err == nil {
t.Fatal("a two-dimensional prod shard was accepted")
}
return nil
})
if err != nil {
t.Fatal(err)
}
}