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