1use 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(
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
74pub 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
126pub 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 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
156pub 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 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}