Skip to main content

solow_stats/
nonparametric_tests.rs

1//! Non-parametric two-sample and k-sample tests.
2//!
3//! * [`mannwhitneyu`] — Mann-Whitney U (Wilcoxon rank-sum).
4//! * [`kruskal`] — Kruskal-Wallis one-way ANOVA on ranks.
5//! * [`mcnemar`] — McNemar's paired-binary test.
6//! * [`chi2_contingency`] — Pearson χ² test of independence.
7
8use solow_core::{Error, Result};
9
10/// A one-degree-of-freedom test statistic + p-value.
11#[derive(Clone, Copy, Debug, PartialEq)]
12pub struct TestResult {
13    /// The test statistic.
14    pub statistic: f64,
15    /// Two-sided p-value.
16    pub pvalue: f64,
17}
18
19/// Mann-Whitney U test — two-sided, normal-approximation p-value with
20/// tie correction (Mann-Whitney 1947).
21pub fn mannwhitneyu(x: &[f64], y: &[f64]) -> Result<TestResult> {
22    if x.is_empty() || y.is_empty() {
23        return Err(Error::Value(
24            "mannwhitneyu: both samples must be non-empty".into(),
25        ));
26    }
27    let n1 = x.len() as f64;
28    let n2 = y.len() as f64;
29    let mut combined: Vec<(f64, u8)> = x
30        .iter()
31        .map(|&v| (v, 0))
32        .chain(y.iter().map(|&v| (v, 1)))
33        .collect();
34    combined.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap());
35    let n = combined.len();
36    let mut ranks = vec![0.0_f64; n];
37    let mut i = 0;
38    let mut ties_sum = 0.0_f64;
39    while i < n {
40        let mut j = i;
41        while j + 1 < n && combined[j + 1].0 == combined[i].0 {
42            j += 1;
43        }
44        let avg = ((i + j) as f64 + 2.0) / 2.0;
45        let t = (j - i + 1) as f64;
46        if t > 1.0 {
47            ties_sum += (t.powi(3) - t) / 12.0;
48        }
49        for k in i..=j {
50            ranks[k] = avg;
51        }
52        i = j + 1;
53    }
54    let mut r1 = 0.0_f64;
55    for i in 0..n {
56        if combined[i].1 == 0 {
57            r1 += ranks[i];
58        }
59    }
60    let u1 = r1 - n1 * (n1 + 1.0) / 2.0;
61    let u2 = n1 * n2 - u1;
62    let u = u1.min(u2);
63    let mean = n1 * n2 / 2.0;
64    let n_all = n1 + n2;
65    let var = n1 * n2 / 12.0 * (n_all + 1.0 - ties_sum / (n_all * (n_all - 1.0) / 12.0));
66    let z = (u - mean) / var.sqrt().max(1e-300);
67    let pvalue = 2.0 * standard_normal_survival(z.abs());
68    Ok(TestResult {
69        statistic: u,
70        pvalue,
71    })
72}
73
74/// Kruskal-Wallis one-way ANOVA on ranks. Returns `(H, p)` where `H`
75/// approximately follows a χ²(k − 1) distribution.
76pub fn kruskal(groups: &[Vec<f64>]) -> Result<TestResult> {
77    if groups.len() < 2 {
78        return Err(Error::Value("kruskal: need ≥ 2 groups".into()));
79    }
80    let n_total: usize = groups.iter().map(|g| g.len()).sum();
81    if n_total < 3 {
82        return Err(Error::Value("kruskal: need ≥ 3 total samples".into()));
83    }
84    let mut combined: Vec<(f64, usize)> = Vec::with_capacity(n_total);
85    for (gi, g) in groups.iter().enumerate() {
86        for &v in g {
87            combined.push((v, gi));
88        }
89    }
90    combined.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap());
91    let mut ranks = vec![0.0_f64; n_total];
92    let mut i = 0;
93    while i < n_total {
94        let mut j = i;
95        while j + 1 < n_total && combined[j + 1].0 == combined[i].0 {
96            j += 1;
97        }
98        let avg = ((i + j) as f64 + 2.0) / 2.0;
99        for k in i..=j {
100            ranks[k] = avg;
101        }
102        i = j + 1;
103    }
104    let k = groups.len() as f64;
105    let mut rank_sum = vec![0.0_f64; groups.len()];
106    let mut group_size = vec![0.0_f64; groups.len()];
107    for i in 0..n_total {
108        rank_sum[combined[i].1] += ranks[i];
109        group_size[combined[i].1] += 1.0;
110    }
111    let n = n_total as f64;
112    let mut h = 0.0_f64;
113    for j in 0..groups.len() {
114        if group_size[j] > 0.0 {
115            h += rank_sum[j] * rank_sum[j] / group_size[j];
116        }
117    }
118    h = 12.0 / (n * (n + 1.0)) * h - 3.0 * (n + 1.0);
119    let pvalue = chi2_survival(h, k - 1.0);
120    Ok(TestResult {
121        statistic: h,
122        pvalue,
123    })
124}
125
126/// McNemar's test on a 2 × 2 paired-binary table.
127pub fn mcnemar(b: usize, c: usize, exact: bool) -> Result<TestResult> {
128    let bf = b as f64;
129    let cf = c as f64;
130    if b + c == 0 {
131        return Err(Error::Value("mcnemar: b + c must be > 0".into()));
132    }
133    if exact {
134        // Two-sided exact binomial p on min(b, c) with n = b + c, p = 0.5.
135        let n = b + c;
136        let k = b.min(c);
137        let mut p = 0.0_f64;
138        for i in 0..=k {
139            p += binomial_pmf(n, i, 0.5);
140        }
141        let pvalue = (2.0 * p).min(1.0);
142        Ok(TestResult {
143            statistic: k as f64,
144            pvalue,
145        })
146    } else {
147        let stat = (bf - cf).powi(2) / (bf + cf);
148        let pvalue = chi2_survival(stat, 1.0);
149        Ok(TestResult {
150            statistic: stat,
151            pvalue,
152        })
153    }
154}
155
156/// Pearson χ² test of independence on a contingency table (rows ×
157/// cols). Returns `(χ², p, dof, expected)`.
158pub fn chi2_contingency(observed: &[Vec<f64>]) -> Result<(f64, f64, usize, Vec<Vec<f64>>)> {
159    if observed.is_empty() || observed[0].is_empty() {
160        return Err(Error::Value("chi2_contingency: empty table".into()));
161    }
162    let r = observed.len();
163    let c = observed[0].len();
164    for row in observed {
165        if row.len() != c {
166            return Err(Error::Value(
167                "chi2_contingency: rows have inconsistent widths".into(),
168            ));
169        }
170    }
171    let mut row_sums = vec![0.0_f64; r];
172    let mut col_sums = vec![0.0_f64; c];
173    let mut total = 0.0_f64;
174    for i in 0..r {
175        for j in 0..c {
176            row_sums[i] += observed[i][j];
177            col_sums[j] += observed[i][j];
178            total += observed[i][j];
179        }
180    }
181    let mut expected = vec![vec![0.0_f64; c]; r];
182    for i in 0..r {
183        for j in 0..c {
184            expected[i][j] = row_sums[i] * col_sums[j] / total.max(1e-300);
185        }
186    }
187    let mut chi2 = 0.0_f64;
188    for i in 0..r {
189        for j in 0..c {
190            let e = expected[i][j];
191            if e > 0.0 {
192                let d = observed[i][j] - e;
193                chi2 += d * d / e;
194            }
195        }
196    }
197    let dof = (r - 1) * (c - 1);
198    let pvalue = chi2_survival(chi2, dof as f64);
199    Ok((chi2, pvalue, dof, expected))
200}
201
202fn binomial_pmf(n: usize, k: usize, p: f64) -> f64 {
203    if k > n {
204        return 0.0;
205    }
206    let ln_c = ln_choose(n, k);
207    let logp = (ln_c + k as f64 * p.ln() + (n - k) as f64 * (1.0 - p).ln()).exp();
208    logp
209}
210
211fn ln_choose(n: usize, k: usize) -> f64 {
212    (1..=k)
213        .map(|i| ((n - i + 1) as f64).ln() - (i as f64).ln())
214        .sum()
215}
216
217fn chi2_survival(x: f64, df: f64) -> f64 {
218    if x <= 0.0 {
219        return 1.0;
220    }
221    1.0 - lower_regularised_gamma(df / 2.0, x / 2.0)
222}
223
224fn lower_regularised_gamma(s: f64, x: f64) -> f64 {
225    if x < 0.0 || s <= 0.0 {
226        return 0.0;
227    }
228    if x < s + 1.0 {
229        gamma_series(s, x)
230    } else {
231        1.0 - gamma_continued_fraction(s, x)
232    }
233}
234
235fn gamma_series(s: f64, x: f64) -> f64 {
236    let mut sum = 1.0 / s;
237    let mut term = sum;
238    for n in 1..200 {
239        term *= x / (s + n as f64);
240        sum += term;
241        if term.abs() < sum.abs() * 3e-15 {
242            break;
243        }
244    }
245    sum * (-x + s * x.ln() - ln_gamma(s)).exp()
246}
247
248fn gamma_continued_fraction(s: f64, x: f64) -> f64 {
249    let mut b = x + 1.0 - s;
250    let mut c = 1.0 / 1e-300;
251    let mut d = 1.0 / b;
252    let mut h = d;
253    for i in 1..200 {
254        let an = -(i as f64) * (i as f64 - s);
255        b += 2.0;
256        d = an * d + b;
257        if d.abs() < 1e-300 {
258            d = 1e-300;
259        }
260        c = b + an / c;
261        if c.abs() < 1e-300 {
262            c = 1e-300;
263        }
264        d = 1.0 / d;
265        let delta = d * c;
266        h *= delta;
267        if (delta - 1.0).abs() < 3e-15 {
268            break;
269        }
270    }
271    (-x + s * x.ln() - ln_gamma(s)).exp() * h
272}
273
274fn ln_gamma(x: f64) -> f64 {
275    let g = 7.0;
276    let cof = [
277        0.999_999_999_999_809_93,
278        676.520_368_121_885_1,
279        -1_259.139_216_722_402_8,
280        771.323_428_777_653_13,
281        -176.615_029_162_140_59,
282        12.507_343_278_686_905,
283        -0.138_571_095_265_720_12,
284        9.984_369_578_019_571_5e-6,
285        1.505_632_735_149_311_6e-7,
286    ];
287    if x < 0.5 {
288        std::f64::consts::PI.ln() - (std::f64::consts::PI * x).sin().ln() - ln_gamma(1.0 - x)
289    } else {
290        let x = x - 1.0;
291        let mut a = cof[0];
292        let t = x + g + 0.5;
293        for (i, &c) in cof.iter().enumerate().skip(1) {
294            a += c / (x + i as f64);
295        }
296        0.5 * (2.0 * std::f64::consts::PI).ln() + (x + 0.5) * t.ln() - t + a.ln()
297    }
298}
299
300fn standard_normal_survival(z: f64) -> f64 {
301    0.5 * erfc(z / std::f64::consts::SQRT_2)
302}
303
304fn erfc(x: f64) -> f64 {
305    let a1 = 0.254_829_592;
306    let a2 = -0.284_496_736;
307    let a3 = 1.421_413_741;
308    let a4 = -1.453_152_027;
309    let a5 = 1.061_405_429;
310    let p = 0.327_591_1;
311    let sign = if x < 0.0 { -1.0 } else { 1.0 };
312    let ax = x.abs();
313    let t = 1.0 / (1.0 + p * ax);
314    let y = 1.0 - (((((a5 * t + a4) * t) + a3) * t + a2) * t + a1) * t * (-ax * ax).exp();
315    1.0 - sign * y
316}
317
318#[cfg(test)]
319mod tests {
320    use super::*;
321
322    #[test]
323    fn mannwhitneyu_flags_a_clear_difference() {
324        let a = vec![1.0_f64, 2.0, 3.0, 4.0, 5.0];
325        let b = vec![10.0_f64, 11.0, 12.0, 13.0, 14.0];
326        let r = mannwhitneyu(&a, &b).unwrap();
327        assert!(r.pvalue < 0.05);
328    }
329
330    #[test]
331    fn kruskal_flags_three_group_difference() {
332        let a = vec![1.0_f64, 2.0, 3.0];
333        let b = vec![10.0_f64, 11.0, 12.0];
334        let c = vec![100.0_f64, 101.0, 102.0];
335        let r = kruskal(&[a, b, c]).unwrap();
336        assert!(r.pvalue < 0.05);
337    }
338
339    #[test]
340    fn mcnemar_returns_the_correct_stat_for_a_2x2_table() {
341        let r = mcnemar(20, 5, false).unwrap();
342        // With b=20, c=5: chi2 = (20-5)^2/25 = 9.
343        assert!((r.statistic - 9.0).abs() < 1e-12);
344        assert!(r.pvalue < 0.01);
345    }
346
347    #[test]
348    fn chi2_contingency_returns_a_finite_pvalue() {
349        let table = vec![vec![10.0, 20.0, 30.0], vec![6.0, 9.0, 17.0]];
350        let (chi2, p, dof, _expected) = chi2_contingency(&table).unwrap();
351        assert!(chi2 >= 0.0);
352        assert!((0.0..=1.0).contains(&p));
353        assert_eq!(dof, 2);
354    }
355}