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

229 lines
8.0 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 (
"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
}