Files
tensor/internal/core/reduce_accuracy_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

358 lines
10 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
}