Files
tensor/stats/rolling_prec_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

231 lines
7.8 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Precision pins for the incremental rolling totals: the carried
// two-float update is held against an exact big.Float reference on the
// adversarial data the drift shows on, beside the per-window rescan it
// replaced and the uncompensated incremental walk it must beat.
package stats
import (
"math"
"math/big"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// rollingPrecReference folds one window of vals exactly, at 200 bits.
func rollingPrecReference(vals []float64, start, window int) *big.Float {
sum := new(big.Float).SetPrec(200)
for _, v := range vals[start : start+window] {
sum.Add(sum, new(big.Float).SetPrec(200).SetFloat64(v))
}
return sum
}
// rollingPrecWalk carries the exact total across the whole series with
// the entering-minus-leaving update at 200 bits: one reference value
// per window position, at a precision the float64 answers cannot see.
func rollingPrecWalk(vals []float64, window int) []*big.Float {
refs := make([]*big.Float, len(vals)-window+1)
refs[0] = rollingPrecReference(vals, 0, window)
for i := 1; i < len(refs); i++ {
ref := new(big.Float).SetPrec(200).Set(refs[i-1])
ref.Add(ref, new(big.Float).SetPrec(200).SetFloat64(vals[i+window-1]))
ref.Sub(ref, new(big.Float).SetPrec(200).SetFloat64(vals[i-1]))
refs[i] = ref
}
return refs
}
// rollingPrecWorst returns the worst relative deviation of got from
// the reference walk, with the position it sat at.
func rollingPrecWorst(t *testing.T, label string, got []float64, refs []*big.Float) (float64, int) {
t.Helper()
worst, worstAt := 0.0, 0
for i, ref := range refs {
refF, _ := ref.Float64()
dev := math.Abs(got[i]-refF) / math.Abs(refF)
if dev > worst {
worst, worstAt = dev, i
}
}
t.Logf("%s: worst relative deviation %.3g at position %d", label, worst, worstAt)
return worst, worstAt
}
// rollingPrecOldFold is the algorithm the incremental totals replaced:
// every window folded from scratch, exactly as rolling.go carried it
// before. It is the accuracy bar the carried update may not drop
// below.
func rollingPrecOldFold(vals []float64, window int, mean bool) []float64 {
outLen := len(vals) - window + 1
out := make([]float64, outLen)
for i := range outLen {
total := 0.0
for _, v := range vals[i : i+window] {
total += v
}
if mean {
total /= float64(window)
}
out[i] = total
}
return out
}
// rollingPrecPureIncremental is the uncompensated carried update, the
// rejected alternative: enter minus leave on a bare running sum. It is
// kept here, measured, because the drift it suffers is the reason the
// shipped walk carries a two-float pair.
func rollingPrecPureIncremental(vals []float64, window int) []float64 {
outLen := len(vals) - window + 1
out := make([]float64, outLen)
total := 0.0
for _, v := range vals[:window] {
total += v
}
out[0] = total
for i := 1; i < outLen; i++ {
total += vals[i+window-1] - vals[i-1]
out[i] = total
}
return out
}
// TestRollingSumMeanPrecisionAgainstBigFloat holds the incremental
// rolling totals against the exact reference on data chosen for the
// drift: values near 1e8 beside values near 1, signs mixed, windows
// long. The acceptance bar is the accuracy of the per-window rescan:
// the carried pair must sit no further from the reference than the
// rescan does, worst position against worst position. The
// uncompensated walk is measured beside them and reported, and the
// reference itself is spot-checked against direct per-window folds.
func TestRollingSumMeanPrecisionAgainstBigFloat(t *testing.T) {
const n = 40000
// Alternating magnitudes: the running total sits near 1e11 while
// single units enter and leave it.
alternating := make([]float64, n)
// Mostly small with sparse large spikes, so the carried total is
// dominated by samples long gone from the window.
spiked := make([]float64, n)
// A slow ramp of near-equal large values: every window sum is a
// wide cancellation the plain update is worst at.
ramp := make([]float64, n)
for i := range n {
sign := 1.0
if (i/64)%2 == 1 {
sign = -1
}
alternating[i] = sign * 1e8 * (1 + float64(i%3)*0.25)
if i%2 == 1 {
alternating[i] = float64(i%7) - 3
}
spiked[i] = float64(i%11) - 5
if i%97 == 0 {
spiked[i] = 1e8 * sign
}
ramp[i] = 1e8 + 0.001*float64(i)
}
for _, tc := range []struct {
name string
vals []float64
}{
{"alternating magnitudes", alternating},
{"sparse spikes", spiked},
{"near-equal ramp", ramp},
} {
t.Run(tc.name, func(t *testing.T) {
for _, window := range []int{8, 64, 4096} {
t.Run("window", func(t *testing.T) {
refs := rollingPrecWalk(tc.vals, window)
// The reference walk is itself incremental; prove it
// carries no drift by refolding sampled windows whole.
for start := 0; start < len(refs); start += 2048 {
if diff := rollingPrecReference(tc.vals, start, window); diff.Cmp(refs[start]) != 0 {
t.Fatalf("reference drift at %d: %v vs %v", start, diff, refs[start])
}
}
sumArr, err := core.FromFloats(tc.vals, len(tc.vals))
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
newSum, err := RollingSum(sumArr, window)
if err != nil {
t.Fatalf("RollingSum: %v", err)
}
newMean, err := RollingMean(sumArr, window)
if err != nil {
t.Fatalf("RollingMean: %v", err)
}
gotSum := newSum.RawFloats()[:newSum.Len()]
gotMean := newMean.RawFloats()[:newMean.Len()]
oldSum := rollingPrecOldFold(tc.vals, window, false)
oldMean := rollingPrecOldFold(tc.vals, window, true)
pureSum := rollingPrecPureIncremental(tc.vals, window)
worstSum, _ := rollingPrecWorst(t, "new sum", gotSum, refs)
worstOldSum, _ := rollingPrecWorst(t, "rescan sum", oldSum, refs)
rollingPrecWorst(t, "uncompensated sum", pureSum, refs)
refMeans := make([]*big.Float, len(refs))
for i, ref := range refs {
refMeans[i] = new(big.Float).SetPrec(200).Quo(ref, big.NewFloat(float64(window)))
}
worstMean, _ := rollingPrecWorst(t, "new mean", gotMean, refMeans)
worstOldMean, _ := rollingPrecWorst(t, "rescan mean", oldMean, refMeans)
if worstSum > worstOldSum {
t.Fatalf("RollingSum window %d: worst relative deviation %.3g is past the rescan's %.3g",
window, worstSum, worstOldSum)
}
if worstMean > worstOldMean {
t.Fatalf("RollingMean window %d: worst relative deviation %.3g is past the rescan's %.3g",
window, worstMean, worstOldMean)
}
})
}
})
}
}
// TestRollingSumMeanNonFiniteMatchesRescan pins the non-finite route:
// a window carrying a NaN or an infinity, and the windows that recover
// after one, answer the per-window rescan's value bit for bit, and the
// carried walk resumes cleanly once every non-finite sample has left.
func TestRollingSumMeanNonFiniteMatchesRescan(t *testing.T) {
vals := make([]float64, 512)
for i := range vals {
vals[i] = float64(i%13) - 6 + 0.125*float64(i%7)
}
vals[5] = math.NaN()
vals[100] = math.Inf(1)
vals[201] = math.Inf(-1)
vals[300] = 1.5e308
vals[301] = 1.5e308
a, err := core.FromFloats(vals, len(vals))
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
for _, window := range []int{1, 2, 4, 64} {
for _, mean := range []bool{false, true} {
var got *core.Array
if mean {
got, err = RollingMean(a, window)
} else {
got, err = RollingSum(a, window)
}
if err != nil {
t.Fatalf("window %d: %v", window, err)
}
want := rollingPrecOldFold(vals, window, mean)
for i := range want {
gotV, wantV := got.RawFloats()[i], want[i]
if gotV != wantV && !(math.IsNaN(gotV) && math.IsNaN(wantV)) {
t.Fatalf("window %d mean %v: position %d = %.17g, the rescan answers %.17g",
window, mean, i, gotV, wantV)
}
}
}
}
}