// Copyright (c) 2026 Petr Balvín (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") } }