346 lines
11 KiB
Go
346 lines
11 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 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)
|
||
|
|
}
|
||
|
|
}
|