// 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 axis sums fold each line through the canonical partition the // global sums use: fixed blocks of the line, the block partials combined // through the balanced tree. The claims pinned here against big.Float at // 200 bits: on lines long enough for a chain's roundings to pile up the // tree holds its accuracy where the chain drifts; on the large-plus-small // counterpoint the chain's running total swallows the small elements // outright and the tree's blocks keep them; the partition follows from // the line length alone, so no worker count moves a bit; and a // single-line fold answers Sum's own bits. // axisChain is the shape the axis fold replaced: one accumulator per // line, addends in ascending element order. func axisChain(v []float64) float64 { var s float64 for _, x := range v { s += x } return s } func TestAxisSumAccuracyAgainstBigFloat(t *testing.T) { for _, n := range []int{1 << 13, 1 << 16, 1 << 20} { v := foldFixture(n) a, err := FromFloats(v, 1, n) if err != nil { t.Fatal(err) } got, err := SumAxis(a, 1) if err != nil { t.Fatal(err) } want := bigSum(v) axisErr := relErr(got.FloatAt(0), want) chainErr := relErr(axisChain(v), want) t.Logf("n=%d: axis fold %.3e, one chain %.3e", n, axisErr, chainErr) // A single contiguous line is one canonical fold: the global // Sum over the same elements must give the same bits. one, err := FromFloats(v, n) if err != nil { t.Fatal(err) } if math.Float64bits(got.FloatAt(0)) != math.Float64bits(Sum(one).Float()) { t.Fatalf("n=%d: the single-line axis fold %v disagrees with Sum %v", n, got.FloatAt(0), Sum(one).Float()) } // The tree never rounds an order of magnitude past the chain // where the chain happens to win by luck. if axisErr > 4*chainErr && axisErr > 1e-16 { t.Errorf("n=%d: the axis fold rounds an order worse than the chain: %.3e against %.3e", n, axisErr, chainErr) } } } // TestAxisSumCounterpoint pins the large-plus-small line: one element at // 1e16 and the rest ones. The chain's running total sits on the 1e16 // grid, where an ulp is 2, and swallows every one it meets; the tree's // blocks sum the ones among themselves before they meet the giant. The // exact referent is the big.Float sum. func TestAxisSumCounterpoint(t *testing.T) { const n = 1 << 13 v := make([]float64, n) for i := range v { v[i] = 1 } v[0] = 1e16 want := bigSum(v) a, err := FromFloats(v, 1, n) if err != nil { t.Fatal(err) } got, err := SumAxis(a, 1) if err != nil { t.Fatal(err) } w, _ := want.Float64() chainLoss := math.Abs(axisChain(v) - w) axisLoss := math.Abs(got.FloatAt(0) - w) t.Logf("counterpoint: axis fold loses %.0f, one chain loses %.0f", axisLoss, chainLoss) // The chain drops all 8191 ones; the tree's three clean blocks keep // three quarters of them and lose only the ones sharing the giant's // own block. if axisLoss >= chainLoss { t.Fatalf("the axis fold loses %.0f against the chain's %.0f on the counterpoint", axisLoss, chainLoss) } } // TestAxisSumStridedAccuracyAgainstBigFloat holds the strided walk // (reducing the leading dimension) against the same reference: the // stride must not reintroduce a chain. func TestAxisSumStridedAccuracyAgainstBigFloat(t *testing.T) { const rows, cols = 8, 1 << 13 v := foldFixture(rows * cols) a, err := FromFloats(v, rows, cols) if err != nil { t.Fatal(err) } got, err := SumAxis(a, 0) if err != nil { t.Fatal(err) } for c := range cols { // The strided fold and the packed fold share one canonical // partition, so they must agree bit for bit. column := make([]float64, rows) scale := 0.0 for r := range rows { column[r] = v[r*cols+c] scale = math.Max(scale, math.Abs(column[r])) } packed, err := FromFloats(column, 1, rows) if err != nil { t.Fatal(err) } want, err := SumAxis(packed, 1) if err != nil { t.Fatal(err) } if math.Float64bits(got.FloatAt(c)) != math.Float64bits(want.FloatAt(0)) { t.Fatalf("column %d: the strided fold %v disagrees with the packed fold %v", c, got.FloatAt(c), want.FloatAt(0)) } // The reference check is absolute against the addend scale: a // cancelling column makes the relative error meaningless. acc := new(big.Float).SetPrec(200) for r := range rows { acc.Add(acc, new(big.Float).SetPrec(200).SetFloat64(column[r])) } w, _ := acc.Float64() if d := math.Abs(got.FloatAt(c) - w); d > 1e-13*scale { t.Fatalf("column %d: the strided fold is %.3e from the reference (scale %.3g)", c, d, scale) } } // The same bits under a different worker count: the strided lines // split across workers the same way the contiguous ones do. prev := engine.SetNumWorkers(1) defer engine.SetNumWorkers(prev) one, err := SumAxis(a, 0) if err != nil { t.Fatal(err) } for _, w := range []int{2, 8, 32} { engine.SetNumWorkers(w) many, err := SumAxis(a, 0) if err != nil { t.Fatal(err) } for c := range cols { if math.Float64bits(many.FloatAt(c)) != math.Float64bits(one.FloatAt(c)) { t.Fatalf("workers=%d moved strided column %d", w, c) } } } } // TestAxisSumMeanDeterministic pins the fixed partition across worker // counts for the fold and the mean, contiguous and strided, float64 and // float32. func TestAxisSumMeanDeterministic(t *testing.T) { prev := engine.SetNumWorkers(1) defer engine.SetNumWorkers(prev) build := func() (*Array, *Array) { v := foldFixture(16 * (1 << 13)) contig, err := FromFloats(v, 16, 1<<13) if err != nil { t.Fatal(err) } strid, err := FromFloats(v, 1<<13, 16) if err != nil { t.Fatal(err) } return contig, strid } contig, strid := build() snapshot := func(a *Array, dim int, mean bool) []float64 { var out *Array var err error if mean { out, err = MeanAxis(a, dim) } else { out, err = SumAxis(a, dim) } if err != nil { t.Fatal(err) } vals := make([]float64, out.Len()) copy(vals, out.floats) return vals } for _, mean := range []bool{false, true} { cOne := snapshot(contig, 1, mean) sOne := snapshot(strid, 0, mean) for _, w := range []int{2, 7, 32} { engine.SetNumWorkers(w) for i, got := range snapshot(contig, 1, mean) { if math.Float64bits(got) != math.Float64bits(cOne[i]) { t.Fatalf("workers=%d mean=%v moved contiguous slot %d", w, mean, i) } } for i, got := range snapshot(strid, 0, mean) { if math.Float64bits(got) != math.Float64bits(sOne[i]) { t.Fatalf("workers=%d mean=%v moved strided slot %d", w, mean, i) } } } } } // TestAxisMeanMatchesMean pins the mean's consistency: a single-line // MeanAxis divides the canonical fold by the line length, which is // Mean's own computation. func TestAxisMeanMatchesMean(t *testing.T) { const n = 1 << 13 v := foldFixture(n) a, err := FromFloats(v, 1, n) if err != nil { t.Fatal(err) } got, err := MeanAxis(a, 1) if err != nil { t.Fatal(err) } one, err := FromFloats(v, n) if err != nil { t.Fatal(err) } m, err := Mean(one) if err != nil { t.Fatal(err) } if math.Float64bits(got.FloatAt(0)) != math.Float64bits(m) { t.Fatalf("MeanAxis %.17g against Mean %.17g", got.FloatAt(0), m) } // And the mean keeps the reference's accuracy: the fold error // divided by n. want := new(big.Float).SetPrec(200).Quo(bigSum(v), new(big.Float).SetPrec(200).SetFloat64(n)) if d := relErr(got.FloatAt(0), want); d > 1e-15 { t.Fatalf("MeanAxis rounds %.3e from the big.Float mean", d) } } // TestAxisSumComplexAgainstBigFloat pins the complex fold: the real and // imaginary parts accumulate through the same canonical partition, so // each part sits at the reference's floor. func TestAxisSumComplexAgainstBigFloat(t *testing.T) { const n = 1 << 13 re := foldFixture(n) im := foldFixture(n / 2) im = append(im, im...) c := make([]complex128, n) for i := range c { c[i] = complex(re[i], im[i]) } a, err := FromComplexes(c, 1, n) if err != nil { t.Fatal(err) } got, err := SumAxis(a, 1) if err != nil { t.Fatal(err) } v := got.complexAt(0) for _, part := range []struct { name string got float64 want *big.Float chain float64 }{ {"real", real(v), bigSum(re), axisChain(re)}, {"imaginary", imag(v), bigSum(im), axisChain(im)}, } { axisErr, chainErr := relErr(part.got, part.want), relErr(part.chain, part.want) t.Logf("%s line: axis fold %.3e, one chain %.3e", part.name, axisErr, chainErr) if axisErr > 4*chainErr && axisErr > 1e-16 { t.Fatalf("the %s part rounds an order worse than the chain: %.3e against %.3e", part.name, axisErr, chainErr) } } } // TestAxisSumWidenedAccuracy pins the float32 line. Every widening is // exact, so the scratch fold sums exactly the values an accessor walk // would; the accumulator then narrows to float32, which caps the // published answer at the dtype's own resolution. Pinned here: the // published value is the narrowed scratch fold, and the scratch fold // itself never rounds an order past the chain the fold replaced. func TestAxisSumWidenedAccuracy(t *testing.T) { const n = 1 << 16 fv := foldFixture(n) f32 := make([]float32, n) for i, x := range fv { f32[i] = float32(x) } a32, err := FromFloat32s(f32, 1, n) if err != nil { t.Fatal(err) } got, err := SumAxis(a32, 1) if err != nil { t.Fatal(err) } want := new(big.Float).SetPrec(200) for _, x := range f32 { want.Add(want, new(big.Float).SetPrec(200).SetFloat64(float64(x))) } // The scratch fold is the canonical one; the published value is its // float32 narrowing. scratch := floatFoldSumF32(f32) if math.Float64bits(float64(float32(scratch))) != math.Float64bits(got.FloatAt(0)) { t.Fatalf("the published float32 sum %.9g disagrees with the narrowed scratch fold %.9g", got.FloatAt(0), float32(scratch)) } // One float32 ulp of the reference bounds the published answer. w, _ := want.Float64() ulp := math.Abs(w) * 1.19e-7 if d := math.Abs(float64(got.FloatAt(0)) - w); d > 2*ulp { t.Fatalf("the published float32 sum is %.3e from the reference, above two float32 ulp (%.3e)", d, ulp) } // And the scratch fold holds the chain's level on the long line. var chain float64 for _, x := range f32 { chain += float64(x) } scratchErr, chainErr := relErr(scratch, want), relErr(chain, want) // Both forms sit ten orders below the float32 resolution the // published answer carries, so which of them is luckier on one // length is noise: pin only the level, not the race. t.Logf("float32 scratch fold %.3e, one chain %.3e", scratchErr, chainErr) if scratchErr > 1e-13 { t.Fatalf("the float32 scratch fold rounds %.3e from the reference, far past the dtype's floor", scratchErr) } }