// Copyright (c) 2026 Petr BalvĂ­n (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) } }