231 lines
7.8 KiB
Go
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)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|