887 lines
26 KiB
Go
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)
|
|
}
|
|
}
|