Files

346 lines
11 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// 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 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)
}
}