// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import ( "math" "math/big" "testing" "sourcedock.dev/petrbalvin/tensor/internal/engine" ) // The sum and the dot product are partitioned by length alone, so their // result must not depend on how many workers the engine runs, and the // interleaved partials must round no worse than the single chain they // replaced. Both claims are pinned here against a machine-independent // reference: big.Float at 200 bits. func foldFixture(n int) []float64 { s := uint64(20260921) v := make([]float64, n) for i := range v { s = s*6364136223846793005 + 1442695040888963407 // Alternating signs and a wide magnitude spread, so cancellation // is the dominant error source rather than a rounding curiosity. mag := math.Pow(10, float64(int((s>>50)%13)-6)) v[i] = mag * float64(int64((s>>20)%2001)-1000) / 1000 } return v } func bigSum(v []float64) *big.Float { acc := new(big.Float).SetPrec(200) for _, x := range v { acc.Add(acc, new(big.Float).SetPrec(200).SetFloat64(x)) } return acc } // chainSum is the shape the fold replaced: one accumulator, one chain. func chainSum(v []float64) float64 { var s float64 for _, x := range v { s += x } return s } func chainDot(x, y []float64) float64 { var s float64 for i := range x { s += x[i] * y[i] } return s } func relErr(got float64, want *big.Float) float64 { w, _ := want.Float64() if w == 0 { return math.Abs(got) } return math.Abs(got-w) / math.Abs(w) } func TestSumIndependentOfWorkerCount(t *testing.T) { v := foldFixture(1 << 20) a, err := FromFloats(v, 1<<20) if err != nil { t.Fatal(err) } prev := engine.SetNumWorkers(1) defer engine.SetNumWorkers(prev) one := Sum(a).Float() for _, w := range []int{2, 4, 8, 32} { engine.SetNumWorkers(w) if got := Sum(a).Float(); got != one { t.Fatalf("Sum with %d workers = %v, with 1 worker = %v: the fold must not depend on the worker count", w, got, one) } } engine.SetNumWorkers(prev) // The same for the dot product and the mean. b, err := FromFloats(v, 1<<20) if err != nil { t.Fatal(err) } engine.SetNumWorkers(1) dOne, err := Dot(a, b) if err != nil { t.Fatal(err) } mOne, err := Mean(a) if err != nil { t.Fatal(err) } for _, w := range []int{3, 7, 16} { engine.SetNumWorkers(w) d, err := Dot(a, b) if err != nil { t.Fatal(err) } m, err := Mean(a) if err != nil { t.Fatal(err) } if d.Float() != dOne.Float() || m != mOne { t.Fatalf("worker count %d changed the result: dot %v vs %v, mean %v vs %v", w, d.Float(), dOne.Float(), m, mOne) } } } func TestSumDotAccuracyAgainstBigFloat(t *testing.T) { for _, n := range []int{1 << 12, 1 << 16, 1 << 20} { v := foldFixture(n) a, err := FromFloats(v, n) if err != nil { t.Fatal(err) } want := bigSum(v) got := Sum(a).Float() old := chainSum(v) newErr, oldErr := relErr(got, want), relErr(old, want) t.Logf("n=%d: fold %.3e, one chain %.3e", n, newErr, oldErr) // The two forms round differently, not always in the same // direction: the interleaved partials shorten each dependency // chain, the combination adds three roundings, and on some // lengths the chain wins by luck. What is pinned is the order: // the fold stays within a small factor of the chain and beats it // where the chain is long enough for its own roundings to // accumulate. if newErr > 4*oldErr && newErr > 1e-16 { t.Errorf("n=%d: the fold rounds an order worse than the chain: %.3e against %.3e", n, newErr, oldErr) } // The dot product of the fixture with itself: the same claim. d, err := Dot(a, a) if err != nil { t.Fatal(err) } acc := new(big.Float).SetPrec(200) for _, x := range v { bx := new(big.Float).SetPrec(200).SetFloat64(x) acc.Add(acc, new(big.Float).SetPrec(200).Mul(bx, bx)) } dOld := chainDot(v, v) if dErr, dOldErr := relErr(d.Float(), acc), relErr(dOld, acc); dErr > 4*dOldErr { t.Errorf("n=%d: the dot fold rounds an order worse than the chain: %.3e against %.3e", n, dErr, dOldErr) } } } // TestFoldComposesAcrossBlockCuts pins the property the distributed // reductions stand on: partials folded per block of the canonical // partition combine, through the same tree the whole fold uses, to the // whole fold's exact bits, whatever contiguous cuts of the block range // produced them. func TestFoldComposesAcrossBlockCuts(t *testing.T) { for _, n := range []int{foldChunk + 1, 3*foldChunk + 17, 17 * foldChunk} { v := foldFixture(n) a, err := FromFloats(v, n) if err != nil { t.Fatal(err) } whole := Sum(a).Float() parts := foldParts(n) cuts := [][]int{{0, parts}, {0, 1, parts}, {0, parts / 3, (2 * parts) / 3, parts}} if parts >= 13 { cuts = append(cuts, []int{0, 2, 3, 5, 8, 13, parts}) } for _, cut := range cuts { partials := make([]float64, 0, parts) for j := 0; j+1 < len(cut); j++ { lo, hi := cut[j], cut[j+1] for c := lo; c < hi; c++ { partials = append(partials, foldRange(v[c*n/parts:(c+1)*n/parts])) } } if got := treeSum(partials); math.Float64bits(got) != math.Float64bits(whole) { t.Fatalf("n=%d cut %v: sharded fold %v (%b) against whole fold %v (%b)", n, cut, got, math.Float64bits(got), whole, math.Float64bits(whole)) } } } } // serialExtreme is the walk the partitioned fold replaced, kept here as // the reference: seed past the leading NaNs, then keep the first strictly // better element. func serialExtreme(src []float64, greater bool) float64 { best, k := src[0], 1 for math.IsNaN(best) && k < len(src) { best = src[k] k++ } for _, v := range src[k:] { if (greater && v > best) || (!greater && v < best) { best = v } } return best } func TestExtremesMatchSerialWalk(t *testing.T) { nan := math.NaN() shapes := [][]float64{ {1, 2, 3}, {nan, nan, 5, 1}, {5, nan, nan}, {nan, nan, nan}, {0, math.Copysign(0, -1)}, {math.Copysign(0, -1), 0}, {0, math.Copysign(0, -1), 0}, {math.Inf(1), math.Inf(-1), 1}, {-1, -1, -1}, } // A large array with NaNs and zeros pinned to the chunk boundaries. big := make([]float64, 3*foldChunk+17) for i := range big { big[i] = float64((i*37)%101) - 50 } for _, idx := range []int{0, foldChunk - 1, foldChunk, 2 * foldChunk, len(big) - 1} { big[idx] = nan } big[foldChunk+3] = 0 big[foldChunk+4] = math.Copysign(0, -1) shapes = append(shapes, big) prev := engine.SetNumWorkers(1) defer engine.SetNumWorkers(prev) for si, vals := range shapes { a, err := FromFloats(vals, len(vals)) if err != nil { t.Fatal(err) } for _, w := range []int{1, 2, 3, 8, 32} { engine.SetNumWorkers(w) for _, greater := range []bool{true, false} { var got Scalar var err error if greater { got, err = Max(a) } else { got, err = Min(a) } if err != nil { t.Fatalf("shape %d: %v", si, err) } want := serialExtreme(vals, greater) if math.Float64bits(got.Float()) != math.Float64bits(want) { t.Fatalf("shape %d workers %d greater=%v: %v (%b), serial %v (%b)", si, w, greater, got.Float(), math.Float64bits(got.Float()), want, math.Float64bits(want)) } } } } // The integer path, which has no NaN or zero subtlety but must still // be partition-independent. ivals := make([]int64, foldChunk+5) for i := range ivals { ivals[i] = int64((i*13)%97) - 48 } ia, err := FromInts(ivals, len(ivals)) if err != nil { t.Fatal(err) } engine.SetNumWorkers(1) imin, _ := Min(ia) imax, _ := Max(ia) for _, w := range []int{2, 5, 16} { engine.SetNumWorkers(w) lo, _ := Min(ia) hi, _ := Max(ia) if lo.Int() != imin.Int() || hi.Int() != imax.Int() { t.Fatalf("integer extremes at %d workers: %v/%v against %v/%v", w, lo.Int(), hi.Int(), imin.Int(), imax.Int()) } } } // prodFixture keeps every factor a relative hair away from one, so a // million-fold product stays finite and the error the roundings make // is measurable against the exact referent. func prodFixture(n int) []float64 { s := uint64(20260922) v := make([]float64, n) for i := range v { s = s*6364136223846793005 + 1442695040888963407 v[i] = 1 + float64(int64((s>>40)%2001)-1000)/1e6 } return v } func bigProd(v []float64) *big.Float { acc := new(big.Float).SetPrec(200) acc.SetFloat64(1) for _, x := range v { acc.Mul(acc, new(big.Float).SetPrec(200).SetFloat64(x)) } return acc } func chainProd(v []float64) float64 { m := 1.0 for _, x := range v { m *= x } return m } func TestProdNormAccuracyAgainstBigFloat(t *testing.T) { const n = 1 << 20 // The product: the block tree against the single chain. v := prodFixture(n) a, err := FromFloats(v, n) if err != nil { t.Fatal(err) } want := bigProd(v) got, err := Prod(a, 0, false) if err != nil { t.Fatal(err) } old := chainProd(v) if newErr, oldErr := relErr(got.FloatAt(0), want), relErr(old, want); newErr > 4*oldErr && newErr > 1e-16 { t.Errorf("n=%d: the product tree rounds an order worse than the chain: %.3e against %.3e", n, newErr, oldErr) } else { t.Logf("n=%d: product tree %.3e, one chain %.3e", n, newErr, oldErr) } // The two norm, whose power sum folds through the same tree: the // squared factors keep the fixture near one, so the referent is // meaningful. sq := make([]float64, n) for i := range sq { sq[i] = v[i] * v[i] } acc := new(big.Float).SetPrec(200) for _, x := range sq { acc.Add(acc, new(big.Float).SetPrec(200).SetFloat64(x)) } nrm, err := Norm(a, 2, 0, false) if err != nil { t.Fatal(err) } // The norm closes through Sqrt; compare the power sums, where the // rounding the tree moves lives. chainSum := 0.0 for _, x := range sq { chainSum += x } // The norm answers Sqrt of its power sum; recover the sum to keep // the comparison on the folded quantity. gotSum := nrm.FloatAt(0) * nrm.FloatAt(0) if newErr, oldErr := relErr(gotSum, acc), relErr(chainSum, acc); newErr > 4*oldErr && newErr > 1e-16 { t.Errorf("n=%d: the norm's power sum rounds an order worse than the chain: %.3e against %.3e", n, newErr, oldErr) } else { t.Logf("n=%d: norm power sum %.3e, one chain %.3e", n, newErr, oldErr) } _ = want }