225 lines
7.2 KiB
Go
225 lines
7.2 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package stats
|
||
|
||
import (
|
||
"math"
|
||
"strings"
|
||
"testing"
|
||
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
)
|
||
|
||
// TestANOVAOneWayIdenticalMeans pins the null behaviour: two groups
|
||
// with the same mean give F = 0 exactly and p = 1 exactly, since the
|
||
// between-group sum of squares vanishes.
|
||
func TestANOVAOneWayIdenticalMeans(t *testing.T) {
|
||
groups := []*core.Array{
|
||
mustFloats(t, []float64{1, 2, 3}),
|
||
mustFloats(t, []float64{3, 1, 2}),
|
||
}
|
||
f, p, err := ANOVAOneWay(groups)
|
||
if err != nil {
|
||
t.Fatalf("ANOVAOneWay: %v", err)
|
||
}
|
||
if f != 0 {
|
||
t.Errorf("F = %v, want 0", f)
|
||
}
|
||
if p != 1 {
|
||
t.Errorf("p = %v, want 1", p)
|
||
}
|
||
}
|
||
|
||
// TestANOVAOneWayHandComputed pins a fully hand-computable case:
|
||
// groups {1,2,3} and {11,12,13} give SSB = 150, SSW = 4, F = 150 with
|
||
// (1, 4) degrees of freedom, and the upper tail has the closed form
|
||
// P(F > f) = 1 − 1.5·s + 0.5·s³ with s = √(1 − 2/77), which the test
|
||
// evaluates independently of the incomplete-beta continued fraction.
|
||
func TestANOVAOneWayHandComputed(t *testing.T) {
|
||
f, p, err := ANOVAOneWay([]*core.Array{
|
||
mustFloats(t, []float64{1, 2, 3}),
|
||
mustFloats(t, []float64{11, 12, 13}),
|
||
})
|
||
if err != nil {
|
||
t.Fatalf("ANOVAOneWay: %v", err)
|
||
}
|
||
if math.Abs(f-150) > 1e-9 {
|
||
t.Errorf("F = %v, want 150", f)
|
||
}
|
||
s := math.Sqrt(75.0 / 77.0)
|
||
want := 1 - 1.5*s + 0.5*s*s*s
|
||
if math.Abs(p-want) > 1e-9 {
|
||
t.Errorf("p = %.16g, want %.16g", p, want)
|
||
}
|
||
}
|
||
|
||
// TestANOVAOneWayThreeGroups pins a three-group case whose p-value
|
||
// reduces to elementary arithmetic: F = 3 with (2, 6) degrees of
|
||
// freedom has upper tail I_{1/2}(3, 1) = (1/2)³ = 1/8 exactly.
|
||
func TestANOVAOneWayThreeGroups(t *testing.T) {
|
||
f, p, err := ANOVAOneWay([]*core.Array{
|
||
mustFloats(t, []float64{1, 2, 3}),
|
||
mustFloats(t, []float64{2, 3, 4}),
|
||
mustFloats(t, []float64{3, 4, 5}),
|
||
})
|
||
if err != nil {
|
||
t.Fatalf("ANOVAOneWay: %v", err)
|
||
}
|
||
if math.Abs(f-3) > 1e-9 {
|
||
t.Errorf("F = %v, want 3", f)
|
||
}
|
||
if math.Abs(p-0.125) > 1e-12 {
|
||
t.Errorf("p = %v, want 0.125", p)
|
||
}
|
||
}
|
||
|
||
// TestANOVAOneWayDegenerate pins the zero-variance limits: distinct
|
||
// means with no noise send F to +Inf with p = 0, while one repeated
|
||
// value everywhere is refused instead of answered 0/0.
|
||
func TestANOVAOneWayDegenerate(t *testing.T) {
|
||
f, p, err := ANOVAOneWay([]*core.Array{
|
||
mustFloats(t, []float64{1, 1}),
|
||
mustFloats(t, []float64{2, 2}),
|
||
})
|
||
if err != nil {
|
||
t.Fatalf("ANOVAOneWay: %v", err)
|
||
}
|
||
if !math.IsInf(f, 1) || p != 0 {
|
||
t.Errorf("F = %v, p = %v, want +Inf and 0", f, p)
|
||
}
|
||
_, _, err = ANOVAOneWay([]*core.Array{
|
||
mustFloats(t, []float64{5, 5}),
|
||
mustFloats(t, []float64{5, 5}),
|
||
})
|
||
if err == nil || !strings.Contains(err.Error(), "identical value") {
|
||
t.Errorf("all-identical observations: error %v, want the degeneracy refusal", err)
|
||
}
|
||
}
|
||
|
||
// TestANOVAOneWayRejects pins the input contracts.
|
||
func TestANOVAOneWayRejects(t *testing.T) {
|
||
ok := mustFloats(t, []float64{1, 2})
|
||
cases := []struct {
|
||
name string
|
||
groups []*core.Array
|
||
want string
|
||
}{
|
||
{"one group", []*core.Array{ok}, "at least two groups"},
|
||
{"empty group", []*core.Array{ok, mustFloats(t, nil)}, "empty"},
|
||
{"NaN observation", []*core.Array{ok, mustFloats(t, []float64{1, math.NaN()})}, "non-finite"},
|
||
{"infinite observation", []*core.Array{ok, mustFloats(t, []float64{1, math.Inf(1)})}, "non-finite"},
|
||
{"singletons", []*core.Array{mustFloats(t, []float64{1}), mustFloats(t, []float64{2})}, "no within-group degrees of freedom"},
|
||
{"complex group", []*core.Array{ok, mustFromComplexes(t, []complex128{1, 2}, 2)}, "complex"},
|
||
}
|
||
for _, c := range cases {
|
||
if _, _, err := ANOVAOneWay(c.groups); err == nil {
|
||
t.Errorf("%s: expected an error", c.name)
|
||
} else if !strings.Contains(err.Error(), c.want) {
|
||
t.Errorf("%s: error %q lacks %q", c.name, err, c.want)
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestMannWhitneyUSeparated pins the fully separated, tie-free case
|
||
// a = {1,2,3}, b = {4,5,6} against hand arithmetic: every rank goes
|
||
// to b, so u = 0; the tie-free variance is σ² = nm(N+1)/12 = 9·7/12 =
|
||
// 5.25, and the continuity-corrected z is 4/√5.25.
|
||
func TestMannWhitneyUSeparated(t *testing.T) {
|
||
u, p, err := MannWhitneyU(mustFloats(t, []float64{1, 2, 3}), mustFloats(t, []float64{4, 5, 6}))
|
||
if err != nil {
|
||
t.Fatalf("MannWhitneyU: %v", err)
|
||
}
|
||
if u != 0 {
|
||
t.Errorf("u = %v, want 0", u)
|
||
}
|
||
want := 2 * (1 - NormalCDF(4/math.Sqrt(5.25)))
|
||
if math.Abs(p-want) > 1e-12 {
|
||
t.Errorf("p = %.16g, want %.16g", p, want)
|
||
}
|
||
// Independently computed reference (an independent 30-digit calculation).
|
||
if math.Abs(p-0.080855598370052291) > 1e-15 {
|
||
t.Errorf("p = %.16g, want the tabulated 0.080855598370052291", p)
|
||
}
|
||
}
|
||
|
||
// TestMannWhitneyUTies pins the tie-corrected variance on a sample
|
||
// pair with two tied blocks of three: u = 3, σ² = 76/7 by hand, and
|
||
// the continuity-corrected z is 4.5/√(76/7).
|
||
func TestMannWhitneyUTies(t *testing.T) {
|
||
u, p, err := MannWhitneyU(mustFloats(t, []float64{1, 2, 2, 3}), mustFloats(t, []float64{2, 3, 3, 4}))
|
||
if err != nil {
|
||
t.Fatalf("MannWhitneyU: %v", err)
|
||
}
|
||
if u != 3 {
|
||
t.Errorf("u = %v, want 3", u)
|
||
}
|
||
want := 2 * (1 - NormalCDF(4.5/math.Sqrt(76.0/7.0)))
|
||
if math.Abs(p-want) > 1e-12 {
|
||
t.Errorf("p = %.16g, want %.16g", p, want)
|
||
}
|
||
// Independently computed reference (an independent 30-digit calculation).
|
||
if math.Abs(p-0.17203370892182298) > 1e-15 {
|
||
t.Errorf("p = %.16g, want the tabulated 0.17203370892182298", p)
|
||
}
|
||
}
|
||
|
||
// TestMannWhitneyUSymmetry checks the companion statistic and the
|
||
// shared p-value: swapping the samples gives n·m − u and the same p.
|
||
func TestMannWhitneyUSymmetry(t *testing.T) {
|
||
a := mustFloats(t, []float64{1, 2, 2, 3})
|
||
b := mustFloats(t, []float64{2, 3, 3, 4})
|
||
u1, p1, err := MannWhitneyU(a, b)
|
||
if err != nil {
|
||
t.Fatalf("MannWhitneyU(a, b): %v", err)
|
||
}
|
||
u2, p2, err := MannWhitneyU(b, a)
|
||
if err != nil {
|
||
t.Fatalf("MannWhitneyU(b, a): %v", err)
|
||
}
|
||
if u1+u2 != 16 {
|
||
t.Errorf("u + u' = %v + %v, want 16", u1, u2)
|
||
}
|
||
if p1 != p2 {
|
||
t.Errorf("p-values differ: %v vs %v", p1, p2)
|
||
}
|
||
}
|
||
|
||
// TestMannWhitneyUCentral pins the exact-centre behaviour: identical
|
||
// samples put u at n·m/2, where the clamped z = 0 gives p = 1.
|
||
func TestMannWhitneyUCentral(t *testing.T) {
|
||
u, p, err := MannWhitneyU(mustFloats(t, []float64{1, 2, 3, 4}), mustFloats(t, []float64{1, 2, 3, 4}))
|
||
if err != nil {
|
||
t.Fatalf("MannWhitneyU: %v", err)
|
||
}
|
||
if u != 8 {
|
||
t.Errorf("u = %v, want 8", u)
|
||
}
|
||
if p != 1 {
|
||
t.Errorf("p = %v, want 1", p)
|
||
}
|
||
}
|
||
|
||
// TestMannWhitneyURejects pins the input contracts.
|
||
func TestMannWhitneyURejects(t *testing.T) {
|
||
ok := mustFloats(t, []float64{1, 2})
|
||
cases := []struct {
|
||
name string
|
||
a, b *core.Array
|
||
want string
|
||
}{
|
||
{"empty a", mustFloats(t, nil), ok, "non-empty"},
|
||
{"empty b", ok, mustFloats(t, nil), "non-empty"},
|
||
{"complex", ok, mustFromComplexes(t, []complex128{1, 2}, 2), "complex"},
|
||
{"NaN", ok, mustFloats(t, []float64{1, math.NaN()}), "non-finite"},
|
||
{"all tied", mustFloats(t, []float64{1, 1}), mustFloats(t, []float64{1, 1}), "no spread"},
|
||
}
|
||
for _, c := range cases {
|
||
if _, _, err := MannWhitneyU(c.a, c.b); err == nil {
|
||
t.Errorf("%s: expected an error", c.name)
|
||
} else if !strings.Contains(err.Error(), c.want) {
|
||
t.Errorf("%s: error %q lacks %q", c.name, err, c.want)
|
||
}
|
||
}
|
||
}
|