Files
tensor/stats/inference_test.go
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

199 lines
6.5 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package stats
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
// TestCovarianceMatrixExact checks the covariance of y = 2x + 1 over
// x = 1..4: every sample variance is 5/3 and the cross covariance
// 10/3, all exact fractions.
func TestCovarianceMatrixExact(t *testing.T) {
obs := mustFloats(t, []float64{
1, 3,
2, 5,
3, 7,
4, 9,
}, 4, 2)
cov, err := CovarianceMatrix(obs)
if err != nil {
t.Fatalf("CovarianceMatrix: %v", err)
}
if cov.Shape()[0] != 2 || cov.Shape()[1] != 2 {
t.Fatalf("covariance shape %v, want (2, 2)", cov.Shape())
}
for _, want := range []struct {
i, j int
v float64
}{
{0, 0, 5.0 / 3}, {1, 1, 20.0 / 3}, {0, 1, 10.0 / 3}, {1, 0, 10.0 / 3},
} {
if math.Abs(cov.FloatAt(want.i*2+want.j)-want.v) > 1e-14 {
t.Fatalf("cov(%d, %d) = %.16g, want %.16g", want.i, want.j,
cov.FloatAt(want.i*2+want.j), want.v)
}
}
corr, err := CorrelationMatrix(obs)
if err != nil {
t.Fatalf("CorrelationMatrix: %v", err)
}
if math.Abs(corr.FloatAt(0)-1) > 1e-14 || math.Abs(corr.FloatAt(3)-1) > 1e-14 {
t.Fatalf("correlation diagonal = %g, %g, want exactly 1",
corr.FloatAt(0), corr.FloatAt(3))
}
if math.Abs(corr.FloatAt(1)-1) > 1e-14 {
t.Fatalf("perfect linear relation has correlation %g, want 1", corr.FloatAt(1))
}
}
func TestCovarianceMatrixErrors(t *testing.T) {
if _, err := CovarianceMatrix(mustFloats(t, []float64{1, 2, 3, 4}, 4)); err == nil {
t.Fatal("1-D input: want an error")
}
single := mustFloats(t, []float64{1, 2}, 1, 2)
if _, err := CovarianceMatrix(single); err == nil {
t.Fatal("one observation: want an error")
}
constant := mustFloats(t, []float64{1, 2, 1, 2, 1, 2, 1, 2}, 4, 2)
if _, err := CorrelationMatrix(constant); err == nil {
t.Fatal("zero-variance column under correlation: want an error")
}
}
// TestWelchTTestSeparated checks the statistic and degrees of freedom
// on samples with an exact analytic answer: a = [1, 2, 3],
// b = [4, 5, 6] gives t = −3√(3/2) and df = 2(n−1) = 4.
func TestWelchTTestSeparated(t *testing.T) {
a := mustFloats(t, []float64{1, 2, 3})
b := mustFloats(t, []float64{4, 5, 6})
stat, df, p, err := WelchTTest(a, b)
if err != nil {
t.Fatalf("WelchTTest: %v", err)
}
wantT := -3 / math.Sqrt(2.0/3)
if math.Abs(stat-wantT) > 1e-14 {
t.Fatalf("t = %.16g, want %.16g", stat, wantT)
}
if math.Abs(df-4) > 1e-14 {
t.Fatalf("df = %.16g, want 4", df)
}
// The p-value must equal the closed-form tail of the same
// distribution.
wantP, err := BetaIncomplete(4.0/(4.0+wantT*wantT), 2, 0.5)
if err != nil || math.Abs(p-wantP) > 1e-14 {
t.Fatalf("p = %v (%v), want %.16g", p, err, wantP)
}
if p > 0.2 {
t.Fatalf("shifted samples keep p = %g, want a small tail", p)
}
// Identical samples: statistic 0, p 1.
stat, _, p, err = WelchTTest(a, a)
if err != nil || stat != 0 || p < 1 {
t.Fatalf("identical samples: t = %v, p = %v, %v", stat, p, err)
}
if _, _, _, err := WelchTTest(mustFloats(t, []float64{1}), b); err == nil {
t.Fatal("one-observation sample: want an error")
}
}
// TestChiSquareGoodnessOfFairDie pins the statistic of a near-fair
// die: (16, 18, 22, 21, 19, 24) against 20 each gives exactly 2.1.
func TestChiSquareGoodnessOfFairDie(t *testing.T) {
observed := mustFloats(t, []float64{16, 18, 22, 21, 19, 24})
expected := mustFloats(t, []float64{20, 20, 20, 20, 20, 20})
chi2, df, p, err := ChiSquareGoodnessOfFit(observed, expected)
if err != nil {
t.Fatalf("ChiSquareGoodnessOfFit: %v", err)
}
if math.Abs(chi2-2.1) > 1e-14 {
t.Fatalf("chi² = %.16g, want 2.1", chi2)
}
if df != 5 {
t.Fatalf("df = %d, want 5", df)
}
wantP, err := GammaUpper(2.5, 1.05)
if err != nil || math.Abs(p-wantP) > 1e-14 {
t.Fatalf("p = %v (%v), want %.16g", p, err, wantP)
}
if p < 0.5 || p > 0.95 {
t.Fatalf("near-fair die p = %g outside the plausible band", p)
}
if _, _, _, err := ChiSquareGoodnessOfFit(observed, mustFloats(t, []float64{20, 0, 20, 20, 20, 20})); err == nil {
t.Fatal("zero expected bin: want an error")
}
}
// TestKolmogorovSmirnovShifted pins the statistic on a known shift:
// samples {1..4} and {2..5} have sup distance exactly 1/4.
func TestKolmogorovSmirnovShifted(t *testing.T) {
a := mustFloats(t, []float64{1, 2, 3, 4})
b := mustFloats(t, []float64{2, 3, 4, 5})
d, p, err := KolmogorovSmirnovTest(a, b)
if err != nil {
t.Fatalf("KolmogorovSmirnovTest: %v", err)
}
if math.Abs(d-0.25) > 1e-15 {
t.Fatalf("d = %.16g, want 0.25", d)
}
if p < 0.8 {
t.Fatalf("a one-step shift of four points keeps p = %g, want a large value", p)
}
// Identical samples: d = 0, p = 1.
d, p, err = KolmogorovSmirnovTest(b, b)
if err != nil || d != 0 || p < 1 {
t.Fatalf("identical samples: d = %v, p = %v, %v", d, p, err)
}
if _, _, err := KolmogorovSmirnovTest(mustFloats(t, []float64{}), b); err == nil {
t.Fatal("empty sample: want an error")
}
}
// TestBootstrapCIMean resamples a sample whose mean is exactly 10:
// the interval must bracket 10 with a plausible width and stay
// deterministic under the seed.
func TestBootstrapCIMean(t *testing.T) {
data := mustFloats(t, []float64{
10.7, 9.3, 10.1, 8.9, 11.2, 9.8, 10.4, 9.1, 10.9, 9.6,
})
mean := func(v *core.Array) (float64, error) {
return core.Mean(v)
}
lower, upper, err := BootstrapCI(data, mean, 0.9, 500, 42)
if err != nil {
t.Fatalf("BootstrapCI: %v", err)
}
m, merr := core.Mean(data)
if merr != nil {
t.Fatalf("Mean: %v", merr)
}
if !(lower <= m && m <= upper) {
t.Fatalf("interval [%g, %g] misses the sample mean %g", lower, upper, m)
}
if upper-lower <= 0 || upper-lower > 2 {
t.Fatalf("interval width %g is implausible", upper-lower)
}
l2, u2, err := BootstrapCI(data, mean, 0.9, 500, 42)
if err != nil || l2 != lower || u2 != upper {
t.Fatalf("seeded run not deterministic: [%g, %g] vs [%g, %g], %v",
l2, u2, lower, upper, err)
}
// A higher level must not be tighter.
l3, u3, err := BootstrapCI(data, mean, 0.99, 500, 42)
if err != nil || l3 > lower || u3 < upper {
t.Fatalf("99 %% interval [%g, %g] tighter than 90 %% [%g, %g], %v",
l3, u3, lower, upper, err)
}
boom := func(*core.Array) (float64, error) { return 0, base.Errf("statistic failed") }
if _, _, err := BootstrapCI(data, boom, 0.9, 10, 1); err == nil {
t.Fatal("statistic error: want an error")
}
if _, _, err := BootstrapCI(data, mean, 1.5, 10, 1); err == nil {
t.Fatal("level outside (0, 1): want an error")
}
}