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

356 lines
13 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 (
"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
}