356 lines
13 KiB
Go
356 lines
13 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package stats
|
||
|
||
import (
|
||
"math"
|
||
"strconv"
|
||
|
||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
)
|
||
|
||
// Contingency tables: the exact and the asymptotic tests of a
|
||
// cross-classification and the effect size that summarises one. The
|
||
// counts enter as an array and are read through the widening
|
||
// accessor, so a table carried in any real dtype answers the same
|
||
// numbers; a negative or fractional entry is refused, because a count
|
||
// it is not. The exact answers travel no further than lgamma and the
|
||
// package's own incomplete beta, and every p-value is computed in log
|
||
// space from the ratios of table probabilities, so a far-from-uniform
|
||
// table cannot lose its answer to an underflow the absolute
|
||
// probabilities would have suffered.
|
||
|
||
// Alternative names the alternative hypothesis a directional test
|
||
// evaluates.
|
||
type Alternative int
|
||
|
||
const (
|
||
// TwoSided is the alternative of a difference in either direction.
|
||
TwoSided Alternative = iota
|
||
// Less is the alternative that the first named quantity is smaller.
|
||
Less
|
||
// Greater is the alternative that the first named quantity is larger.
|
||
Greater
|
||
)
|
||
|
||
// String returns the name of the alternative, the spelling the error
|
||
// messages carry.
|
||
func (a Alternative) String() string {
|
||
switch a {
|
||
case TwoSided:
|
||
return "two-sided"
|
||
case Less:
|
||
return "less"
|
||
case Greater:
|
||
return "greater"
|
||
}
|
||
return "alternative(" + strconv.Itoa(int(a)) + ")"
|
||
}
|
||
|
||
// maxFisherSupport caps the hypergeometric support FisherExactTest
|
||
// enumerates. Each support point costs a pair of lgamma evaluations,
|
||
// and a table whose margins put millions of possible tables between
|
||
// the support's ends has long left the exact regime the test exists
|
||
// for: there the asymptotic ChiSquareIndependence is the honest
|
||
// answer, and the exact test refuses with the cap named rather than
|
||
// spend seconds enumerating a sum whose terms have all underflowed.
|
||
const maxFisherSupport = 4_000_000
|
||
|
||
// maxExactCount bounds a single count of the exact tests' tables. The
|
||
// enumeration walks the support integer by integer, and the format
|
||
// holds those integers stepwise only below 2^52: past it every second
|
||
// one is missing, the walk's increment stops advancing and the loop
|
||
// cannot leave its first table. One count at most 2^51 keeps every
|
||
// margin the four counts can sum to below 2^52 and every support
|
||
// point an exact integer.
|
||
const maxExactCount = 1 << 51
|
||
|
||
// FisherExactTest runs Fisher's exact test on a 2×2 table of counts,
|
||
// the classical answer to a cross-classification too sparse for the
|
||
// χ² approximation. The rows are the two samples and the columns the
|
||
// two outcomes, and the test conditions on the margins: under
|
||
// independence the first cell follows the hypergeometric law, and the
|
||
// p-value sums the probabilities of the tables at least as extreme as
|
||
// the observed one, exactly. The two-sided sum collects every table
|
||
// whose probability does not exceed the observed one's, so an
|
||
// asymmetric table sums its two tails unequally, as the exact
|
||
// definition demands.
|
||
//
|
||
// Returns the p-value for the given alternative and the sample odds
|
||
// ratio (a·d)/(b·c): +∞ where a zero cell sits against a full one,
|
||
// NaN where the table carries no contrast at all (a zero margin),
|
||
// which answers a p-value of 1, the one table its margins allow.
|
||
//
|
||
// Refuses a table that is not 2×2, complex input, a non-finite,
|
||
// negative or fractional count, a count past the maxExactCount the
|
||
// format holds the integers of, an unknown alternative and a support
|
||
// wider than maxFisherSupport.
|
||
func FisherExactTest(table *core.Array, alternative Alternative) (pValue, oddsRatio float64, err error) {
|
||
const name = "FisherExactTest"
|
||
a, b, c, d, err := squareCounts(name, table)
|
||
if err != nil {
|
||
return 0, 0, err
|
||
}
|
||
switch alternative {
|
||
case TwoSided, Less, Greater:
|
||
default:
|
||
return 0, 0, base.Errf("%s: unknown alternative %s", name, alternative)
|
||
}
|
||
oddsRatio = a * d / (b * c)
|
||
row1, row2 := a+b, c+d
|
||
col1, col2 := a+c, b+d
|
||
if row1 == 0 || row2 == 0 || col1 == 0 || col2 == 0 {
|
||
// A zero margin admits exactly one table, the observed one: the
|
||
// conditioning has removed every degree of freedom and no
|
||
// arrangement is more extreme.
|
||
return 1, oddsRatio, nil
|
||
}
|
||
lo := math.Max(0, row1-col2)
|
||
hi := math.Min(row1, col1)
|
||
if hi-lo+1 > maxFisherSupport {
|
||
return 0, 0, base.Errf("%s: the margins span a support of %.0f tables, above the %d cap; use ChiSquareIndependence",
|
||
name, hi-lo+1, maxFisherSupport)
|
||
}
|
||
// The probability of the table with first cell x, up to the factor
|
||
// the margins fix: log C(col1, x) + log C(col2, row1−x). The
|
||
// common constant (the normalising choice of margins) never enters
|
||
// a ratio, so it is left unevaluated, and the two sums are folded
|
||
// in log space: on a lopsided table the ratio the observed table's
|
||
// probability bears to the mode's overflows the format many times
|
||
// over, and the quotients the direct sums would then divide (Inf by
|
||
// Inf on a one-sided tail, a finite extreme by an infinite total on
|
||
// the two-sided one) answer NaN or a flushed zero where the true
|
||
// tail is representable. The log-sum-exp fold keeps every
|
||
// accumulator finite, so the p-value is one quotient of two
|
||
// well-scaled logarithms at the end.
|
||
logTerm := func(x float64) float64 {
|
||
return lchoose(col1, x) + lchoose(col2, row1-x)
|
||
}
|
||
observed := logTerm(a)
|
||
slack := math.Log1p(1e-9)
|
||
logTotal := math.Inf(-1)
|
||
logExtreme := math.Inf(-1)
|
||
for x := lo; x <= hi; x++ {
|
||
lt := logTerm(x)
|
||
logTotal = logAddExp(logTotal, lt)
|
||
counts := false
|
||
switch alternative {
|
||
case TwoSided:
|
||
// A table counts as at least as extreme when its probability
|
||
// does not exceed the observed one's; the slack absorbs the
|
||
// rounding of equal-probability tables onto either side of 1.
|
||
counts = lt <= observed+slack
|
||
case Less:
|
||
counts = x <= a
|
||
case Greater:
|
||
counts = x >= a
|
||
}
|
||
if counts {
|
||
logExtreme = logAddExp(logExtreme, lt)
|
||
}
|
||
}
|
||
return min(1, math.Exp(logExtreme-logTotal)), oddsRatio, nil
|
||
}
|
||
|
||
// logAddExp folds one logarithm into another without leaving log space:
|
||
// the max guard keeps the exponential's argument at most zero, so the
|
||
// fold never overflows however far apart the two terms sit, and a term
|
||
// at negative infinity (a sum still empty) passes through untouched.
|
||
func logAddExp(a, b float64) float64 {
|
||
if a == math.Inf(-1) {
|
||
return b
|
||
}
|
||
if b == math.Inf(-1) {
|
||
return a
|
||
}
|
||
m := math.Max(a, b)
|
||
return m + math.Log1p(math.Exp(-math.Abs(a-b)))
|
||
}
|
||
|
||
// lchoose returns log C(n, k) for integral arguments held as floats.
|
||
// The lgamma sign is never negative on these arguments, all at least 1.
|
||
func lchoose(n, k float64) float64 {
|
||
la, _ := math.Lgamma(n + 1)
|
||
lb, _ := math.Lgamma(k + 1)
|
||
lc, _ := math.Lgamma(n - k + 1)
|
||
return la - lb - lc
|
||
}
|
||
|
||
// ChiSquareIndependence runs Pearson's χ² test of independence on an
|
||
// r×c table of counts: the expected count under independence is the
|
||
// row total times the column total over the grand total, the
|
||
// statistic is Σ(O−E)²/E over the cells, and the p-value is the upper
|
||
// tail of χ² on (r−1)(c−1) degrees of freedom. A row or column that
|
||
// totals zero carries no information and only divides by zero; it is
|
||
// refused by name rather than absorbed.
|
||
//
|
||
// The approximation is the asymptotic one: with expected counts in
|
||
// the single digits the exact FisherExactTest is the honest test, and
|
||
// this one is the answer once the table is dense enough for χ² to
|
||
// hold. Refuses fewer than two rows or columns, complex input, a
|
||
// non-finite, negative or fractional count, and a zero row or column
|
||
// total.
|
||
func ChiSquareIndependence(table *core.Array) (chi2 float64, df int, pValue float64, err error) {
|
||
const name = "ChiSquareIndependence"
|
||
stat, _, nrows, ncols, serr := contingencyStatistic(name, table)
|
||
if serr != nil {
|
||
return 0, 0, 0, serr
|
||
}
|
||
chi2 = stat
|
||
df = (nrows - 1) * (ncols - 1)
|
||
pValue, err = GammaUpper(float64(df)/2, chi2/2)
|
||
if err != nil {
|
||
return 0, 0, 0, base.Errf("%s: %w", name, err)
|
||
}
|
||
return chi2, df, pValue, nil
|
||
}
|
||
|
||
// CramersV returns Cramér's V for an r×c table, the χ² statistic of
|
||
// independence rescaled into [0, 1] by the sample size and the
|
||
// smaller margin: V = √(χ²/(n·(min(r, c)−1))). Zero means the counts
|
||
// sit exactly on independence, one means each row leans on a single
|
||
// column. The input contract is ChiSquareIndependence's own.
|
||
func CramersV(table *core.Array) (float64, error) {
|
||
const name = "CramersV"
|
||
chi2, total, rows, cols, err := contingencyStatistic(name, table)
|
||
if err != nil {
|
||
return 0, err
|
||
}
|
||
return math.Sqrt(chi2 / (total * float64(min(rows, cols)-1))), nil
|
||
}
|
||
|
||
// McNemarTest runs the exact McNemar test on a 2×2 table of paired
|
||
// counts: the diagonal holds the agreements and the off-diagonal pair
|
||
// (b, c) the two directions of disagreement, and under the null the
|
||
// disagreements split evenly, so min(b, c) follows the binomial law
|
||
// with m = b + c trials at p = ½. The two-sided p-value doubles the
|
||
// smaller tail, evaluated through the package's incomplete beta
|
||
// rather than an m-term sum, so a table with millions of discordant
|
||
// pairs costs the same as one with a dozen. A table with no
|
||
// discordant pair has nothing to test and answers 1. Refuses a table
|
||
// that is not 2×2, complex input, a non-finite, negative or fractional
|
||
// count, and a count past the maxExactCount the format holds the
|
||
// integers of.
|
||
func McNemarTest(table *core.Array) (pValue float64, err error) {
|
||
const name = "McNemarTest"
|
||
_, b, c, _, err := squareCounts(name, table)
|
||
if err != nil {
|
||
return 0, err
|
||
}
|
||
m := b + c
|
||
if m == 0 {
|
||
return 1, nil
|
||
}
|
||
k := min(b, c)
|
||
tail, err := BetaIncomplete(0.5, m-k, k+1)
|
||
if err != nil {
|
||
return 0, base.Errf("%s: %w", name, err)
|
||
}
|
||
return min(1, 2*tail), nil
|
||
}
|
||
|
||
// squareCounts reads a 2×2 count table into its cells, the shared
|
||
// reader of FisherExactTest and McNemarTest.
|
||
func squareCounts(name string, table *core.Array) (a, b, c, d float64, err error) {
|
||
if table.NDim() != 2 || table.Shape()[0] != 2 || table.Shape()[1] != 2 {
|
||
return 0, 0, 0, 0, base.Errf("%s: the table must be 2×2, got shape %s", name, base.ShapeText(table.Shape()))
|
||
}
|
||
if table.Dtype() == core.Complex {
|
||
return 0, 0, 0, 0, base.Errf("%s: complex tables are not supported", name)
|
||
}
|
||
counts, err := countValues(name, table)
|
||
if err != nil {
|
||
return 0, 0, 0, 0, err
|
||
}
|
||
for _, v := range counts {
|
||
if v >= maxExactCount {
|
||
return 0, 0, 0, 0, base.Errf("%s: count %g exceeds 2^51, past which float64 no longer holds the integers stepwise and the exact enumeration cannot run", name, v)
|
||
}
|
||
}
|
||
return counts[0], counts[1], counts[2], counts[3], nil
|
||
}
|
||
|
||
// contingencyStatistic computes the Pearson statistic of an r×c count
|
||
// table together with the total the effect sizes normalise by, the
|
||
// shared body of ChiSquareIndependence and CramersV.
|
||
func contingencyStatistic(name string, table *core.Array) (chi2, total float64, rows, cols int, err error) {
|
||
if table.NDim() != 2 {
|
||
return 0, 0, 0, 0, base.Errf("%s: the table must be rank 2, got shape %s", name, base.ShapeText(table.Shape()))
|
||
}
|
||
rows, cols = table.Shape()[0], table.Shape()[1]
|
||
if rows < 2 || cols < 2 {
|
||
return 0, 0, 0, 0, base.Errf("%s: the table needs at least two rows and two columns, got %d×%d", name, rows, cols)
|
||
}
|
||
if table.Dtype() == core.Complex {
|
||
return 0, 0, 0, 0, base.Errf("%s: complex tables are not supported", name)
|
||
}
|
||
counts, err := countValues(name, table)
|
||
if err != nil {
|
||
return 0, 0, 0, 0, err
|
||
}
|
||
rowTotals := make([]float64, rows)
|
||
colTotals := make([]float64, cols)
|
||
for i := range rows {
|
||
for j := range cols {
|
||
v := counts[i*cols+j]
|
||
rowTotals[i] += v
|
||
colTotals[j] += v
|
||
total += v
|
||
}
|
||
}
|
||
if total == 0 {
|
||
return 0, 0, 0, 0, base.Errf("%s: the table totals zero, there is nothing to test", name)
|
||
}
|
||
for i, r := range rowTotals {
|
||
if r == 0 {
|
||
return 0, 0, 0, 0, base.Errf("%s: row %d totals zero and carries no information", name, i+1)
|
||
}
|
||
}
|
||
for j, c := range colTotals {
|
||
if c == 0 {
|
||
return 0, 0, 0, 0, base.Errf("%s: column %d totals zero and carries no information", name, j+1)
|
||
}
|
||
}
|
||
chi2 = 0.0
|
||
for i := range rows {
|
||
for j := range cols {
|
||
expected := rowTotals[i] * colTotals[j] / total
|
||
d := counts[i*cols+j] - expected
|
||
chi2 += d * d / expected
|
||
}
|
||
}
|
||
return chi2, total, rows, cols, nil
|
||
}
|
||
|
||
// countValues reads a real array as counts, refusing complex input,
|
||
// a non-finite entry, a negative entry and a fractional entry: the
|
||
// accessor walk answers the same widened values for every real dtype,
|
||
// so a table carried in any of them reads identically.
|
||
func countValues(name string, table *core.Array) ([]float64, error) {
|
||
n := table.Len()
|
||
counts := make([]float64, n)
|
||
if fs := rawFloats(table); fs != nil {
|
||
// A rebased view's payload may run past its own count: only the
|
||
// visible elements are counts.
|
||
copy(counts, fs[:n])
|
||
} else {
|
||
for i := range counts {
|
||
counts[i] = table.FloatAt(i)
|
||
}
|
||
}
|
||
for i, v := range counts {
|
||
if math.IsNaN(v) || math.IsInf(v, 0) {
|
||
return nil, base.Errf("%s: count [%d] is not finite (%g)", name, i, v)
|
||
}
|
||
if v < 0 {
|
||
return nil, base.Errf("%s: count [%d] is negative (%g)", name, i, v)
|
||
}
|
||
if v != math.Trunc(v) {
|
||
return nil, base.Errf("%s: count [%d] is fractional (%g)", name, i, v)
|
||
}
|
||
}
|
||
return counts, nil
|
||
}
|