229 lines
8.0 KiB
Go
229 lines
8.0 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package stats
|
||
|
||
import (
|
||
"cmp"
|
||
"math"
|
||
"slices"
|
||
|
||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
)
|
||
|
||
// Group comparison: the parametric one-way analysis of variance and
|
||
// the rank-based Mann-Whitney U test. Both are classical frequentist
|
||
// tooling built on the distribution machinery of cdf.go, so a
|
||
// p-value never leaves the library's own incomplete beta and normal
|
||
// CDF.
|
||
|
||
// ANOVAOneWay runs the classical one-way analysis of variance F test
|
||
// over the groups: the statistic is the between-group mean square
|
||
// divided by the within-group mean square, and the p-value is the
|
||
// upper tail of F with (k−1, N−k) degrees of freedom, evaluated
|
||
// through the regularised incomplete beta the same way WelchTTest's
|
||
// p-value is. A large F says the group means spread further than
|
||
// within-group noise explains.
|
||
//
|
||
// Every group must be a real, non-empty array of finite values, at
|
||
// least two groups are needed, and the observations must leave at
|
||
// least one within-group degree of freedom (N > k). Degenerate
|
||
// samples are refused rather than answered with a NaN: groups whose
|
||
// observations all share one identical value carry no within-group
|
||
// variance to divide by.
|
||
//
|
||
// Errors: fewer than two groups, an empty or complex group, a
|
||
// non-finite observation, N = k, or vanishing within-group variance.
|
||
func ANOVAOneWay(groups []*core.Array) (fStat, pValue float64, err error) {
|
||
if len(groups) < 2 {
|
||
return 0, 0, base.Errf("ANOVAOneWay: at least two groups are needed, got %d", len(groups))
|
||
}
|
||
for i, g := range groups {
|
||
if g == nil {
|
||
return 0, 0, base.Errf("ANOVAOneWay: group %d is nil", i)
|
||
}
|
||
}
|
||
k := len(groups)
|
||
sizes := make([]int, k)
|
||
means := make([]float64, k)
|
||
total := 0
|
||
sumAll := 0.0
|
||
for i, g := range groups {
|
||
if g.Dtype() == core.Complex {
|
||
return 0, 0, base.Errf("ANOVAOneWay: group %d is complex", i)
|
||
}
|
||
s := 0.0
|
||
// The walk is bounded by the group's own element count: a rebased
|
||
// view's payload may run past its visible elements, and those
|
||
// invisible tail slots are nobody's observations.
|
||
if fs := rawFloats(g); fs != nil {
|
||
for _, v := range fs[:g.Len()] {
|
||
if math.IsNaN(v) || math.IsInf(v, 0) {
|
||
return 0, 0, base.Errf("ANOVAOneWay: group %d holds the non-finite value %g", i, v)
|
||
}
|
||
s += v
|
||
}
|
||
} else {
|
||
for j := range g.Len() {
|
||
v := g.FloatAt(j)
|
||
if math.IsNaN(v) || math.IsInf(v, 0) {
|
||
return 0, 0, base.Errf("ANOVAOneWay: group %d holds the non-finite value %g", i, v)
|
||
}
|
||
s += v
|
||
}
|
||
}
|
||
sizes[i] = g.Len()
|
||
if sizes[i] == 0 {
|
||
return 0, 0, base.Errf("ANOVAOneWay: group %d is empty", i)
|
||
}
|
||
means[i] = s / float64(sizes[i])
|
||
total += sizes[i]
|
||
sumAll += s
|
||
}
|
||
if total == k {
|
||
return 0, 0, base.Errf("ANOVAOneWay: one observation per group leaves no within-group degrees of freedom")
|
||
}
|
||
grand := sumAll / float64(total)
|
||
ssb, ssw := 0.0, 0.0
|
||
for i, g := range groups {
|
||
d := means[i] - grand
|
||
ssb += float64(sizes[i]) * d * d
|
||
// Each group's within sum of squares folds through the same
|
||
// canonical partition the variance keeps, and the group totals
|
||
// then join in the groups' own order: the fold is a function of
|
||
// the group's length alone, so the split cannot move a bit.
|
||
if fs := rawFloats(g); fs != nil {
|
||
ssw += sqDeviations(fs[:g.Len()], means[i])
|
||
} else {
|
||
ssw += sqDeviationsAt(g, means[i])
|
||
}
|
||
}
|
||
if ssw == 0 {
|
||
if ssb == 0 {
|
||
return 0, 0, base.Errf("ANOVAOneWay: every observation shares one identical value, so the F statistic is undefined")
|
||
}
|
||
// Distinct means with no within-group noise: F grows without bound.
|
||
return math.Inf(1), 0, nil
|
||
}
|
||
fStat = (ssb / float64(k-1)) / (ssw / float64(total-k))
|
||
// Two-sided upper tail of F(dfB, dfW) through the identity
|
||
// P(F > f) = I_{dfW/(dfW + dfB·f)}(dfW/2, dfB/2), the same
|
||
// closed form the Student-t tail uses.
|
||
x := float64(total-k) / (float64(total-k) + float64(k-1)*fStat)
|
||
pValue, err = BetaIncomplete(x, float64(total-k)/2, float64(k-1)/2)
|
||
if err != nil {
|
||
return 0, 0, base.Errf("ANOVAOneWay: %w", err)
|
||
}
|
||
return fStat, pValue, nil
|
||
}
|
||
|
||
// MannWhitneyU runs the Mann-Whitney U test on two independent
|
||
// samples, the rank-based comparison of location that asks no
|
||
// normality: all observations are ranked together with midranks
|
||
// splitting ties, and u counts how strongly sample a outranks
|
||
// sample b through u = R_a − n(n+1)/2, with the companion statistic
|
||
// for sample b always n·m − u. The p-value is two-sided from the
|
||
// normal approximation with the continuity correction and the
|
||
// tie-corrected variance
|
||
//
|
||
// σ² = (n·m/12)·[(N+1) − Σ t(t²−1)/(N(N−1))],
|
||
//
|
||
// summing t over the tied blocks; with fewer than a dozen
|
||
// observations per sample the exact permutation distribution is
|
||
// noticeably discrete and the approximation only frames the answer.
|
||
//
|
||
// Errors: an empty or complex sample, a non-finite observation, or
|
||
// every observation in both samples tied, which leaves the U
|
||
// statistic no spread.
|
||
func MannWhitneyU(a, b *core.Array) (u, pValue float64, err error) {
|
||
if a.Len() == 0 || b.Len() == 0 {
|
||
return 0, 0, base.Errf("MannWhitneyU: both samples must be non-empty")
|
||
}
|
||
if a.Dtype() == core.Complex || b.Dtype() == core.Complex {
|
||
return 0, 0, base.Errf("MannWhitneyU: complex samples are not supported")
|
||
}
|
||
type tagged struct {
|
||
v float64
|
||
inA bool
|
||
}
|
||
all := make([]tagged, 0, a.Len()+b.Len())
|
||
// Both walks are bounded by the samples' own element counts: a
|
||
// rebased view's payload may run past its visible elements, and the
|
||
// invisible tail slots are nobody's observations.
|
||
if fs := rawFloats(a); fs != nil {
|
||
for _, v := range fs[:a.Len()] {
|
||
if math.IsNaN(v) || math.IsInf(v, 0) {
|
||
return 0, 0, base.Errf("MannWhitneyU: sample a holds the non-finite value %g", v)
|
||
}
|
||
all = append(all, tagged{v, true})
|
||
}
|
||
} else {
|
||
for i := range a.Len() {
|
||
v := a.FloatAt(i)
|
||
if math.IsNaN(v) || math.IsInf(v, 0) {
|
||
return 0, 0, base.Errf("MannWhitneyU: sample a holds the non-finite value %g", v)
|
||
}
|
||
all = append(all, tagged{v, true})
|
||
}
|
||
}
|
||
if fs := rawFloats(b); fs != nil {
|
||
for _, v := range fs[:b.Len()] {
|
||
if math.IsNaN(v) || math.IsInf(v, 0) {
|
||
return 0, 0, base.Errf("MannWhitneyU: sample b holds the non-finite value %g", v)
|
||
}
|
||
all = append(all, tagged{v, false})
|
||
}
|
||
} else {
|
||
for i := range b.Len() {
|
||
v := b.FloatAt(i)
|
||
if math.IsNaN(v) || math.IsInf(v, 0) {
|
||
return 0, 0, base.Errf("MannWhitneyU: sample b holds the non-finite value %g", v)
|
||
}
|
||
all = append(all, tagged{v, false})
|
||
}
|
||
}
|
||
slices.SortFunc(all, func(p, q tagged) int { return cmp.Compare(p.v, q.v) })
|
||
// Walk the sorted pool one tie block at a time; each block shares
|
||
// the midrank of its 1-based rank span.
|
||
rankSumA := 0.0
|
||
tieTerm := 0.0
|
||
for i := 0; i < len(all); {
|
||
j := i
|
||
for j < len(all) && all[j].v == all[i].v {
|
||
j++
|
||
}
|
||
mid := float64(i+j+1) / 2
|
||
for k := i; k < j; k++ {
|
||
if all[k].inA {
|
||
rankSumA += mid
|
||
}
|
||
}
|
||
if t := j - i; t > 1 {
|
||
// Accumulate in float64 from the start: the integer form
|
||
// t·(t²−1) overflows int64 for tie blocks of 2^21 and up,
|
||
// and the wrapped term defeats the all-tied refusal below.
|
||
tieTerm += float64(t) * (float64(t)*float64(t) - 1)
|
||
}
|
||
i = j
|
||
}
|
||
na, nb := a.Len(), b.Len()
|
||
u = rankSumA - float64(na)*float64(na+1)/2
|
||
total := float64(na + nb)
|
||
variance := float64(na) * float64(nb) / 12 *
|
||
(total + 1 - tieTerm/(total*(total-1)))
|
||
if variance <= 0 {
|
||
return 0, 0, base.Errf("MannWhitneyU: every observation is tied, so the U statistic has no spread")
|
||
}
|
||
// Continuity-corrected z from the upper side; an exact central U
|
||
// clamps to z = 0 rather than a small negative value.
|
||
z := (math.Abs(u-float64(na*nb)/2) - 0.5) / math.Sqrt(variance)
|
||
if z < 0 {
|
||
z = 0
|
||
}
|
||
// The two-sided normal tail in one Erfc call on the magnitude, as
|
||
// the GLM fits compute it: 2·(1−Φ(z)) cancels to exactly zero once
|
||
// z passes about 8.3, where the true tail is still representable.
|
||
return u, math.Erfc(math.Abs(z) / math.Sqrt2), nil
|
||
}
|