Skip to main content

qualia_core_db/solvers/statistics/hypothesis/
anova.rs

1//! One-way ANOVA — the F-test for equality of `k` group means, with a real F-tail
2//! p-value from [`fisher_f`](super::super::distributions::fisher_f).
3
4use super::super::descriptive::mean;
5use super::super::distributions::fisher_f;
6
7/// One-way ANOVA result.
8#[derive(Debug, Clone, Copy, PartialEq)]
9pub struct AnovaResult {
10    pub f_statistic: f64,
11    pub p_value: f64,
12    pub df_between: f64,
13    pub df_within: f64,
14    pub ss_between: f64,
15    pub ss_within: f64,
16    pub ms_between: f64,
17    pub ms_within: f64,
18}
19
20/// One-way ANOVA across `groups` (≥ 2 groups, each with ≥ 1 observation, and total
21/// observations > number of groups so the within-group dof is positive). `None`
22/// otherwise.
23pub fn one_way_anova(groups: &[&[f64]]) -> Option<AnovaResult> {
24    let k = groups.len();
25    if k < 2 || groups.iter().any(|g| g.is_empty()) {
26        return None;
27    }
28    let n_total: usize = groups.iter().map(|g| g.len()).sum();
29    if n_total <= k {
30        return None; // df_within would be ≤ 0
31    }
32
33    let grand_mean = groups.iter().flat_map(|g| g.iter()).sum::<f64>() / n_total as f64;
34
35    let mut ss_between = 0.0;
36    let mut ss_within = 0.0;
37    for g in groups {
38        let gm = mean(g)?;
39        ss_between += g.len() as f64 * (gm - grand_mean).powi(2);
40        for &x in *g {
41            ss_within += (x - gm).powi(2);
42        }
43    }
44
45    let df_between = (k - 1) as f64;
46    let df_within = (n_total - k) as f64;
47    let ms_between = ss_between / df_between;
48    let ms_within = ss_within / df_within;
49
50    let f = if ms_within > 0.0 {
51        ms_between / ms_within
52    } else {
53        // Zero within-group variance: F is +∞ unless between-group variance is also
54        // 0 (all values identical), in which case there is no effect.
55        if ms_between > 0.0 {
56            f64::INFINITY
57        } else {
58            0.0
59        }
60    };
61    let p = if f.is_finite() {
62        fisher_f::upper_p(f, df_between, df_within)
63    } else {
64        0.0
65    };
66
67    Some(AnovaResult {
68        f_statistic: f,
69        p_value: p,
70        df_between,
71        df_within,
72        ss_between,
73        ss_within,
74        ms_between,
75        ms_within,
76    })
77}
78
79#[cfg(test)]
80mod tests {
81    use super::*;
82
83    #[test]
84    fn detects_a_real_difference() {
85        // Three clearly separated groups → large F, tiny p.
86        let g1 = [1.0, 2.0, 1.5, 2.5, 2.0];
87        let g2 = [5.0, 6.0, 5.5, 6.5, 6.0];
88        let g3 = [10.0, 11.0, 10.5, 11.5, 11.0];
89        let r = one_way_anova(&[&g1, &g2, &g3]).unwrap();
90        assert_eq!(r.df_between, 2.0);
91        assert_eq!(r.df_within, 12.0);
92        assert!(r.f_statistic > 50.0);
93        assert!(r.p_value < 1e-6);
94    }
95
96    #[test]
97    fn no_difference_is_not_significant() {
98        let g1 = [4.0, 5.0, 6.0, 5.0];
99        let g2 = [5.0, 6.0, 4.0, 5.0];
100        let g3 = [6.0, 4.0, 5.0, 5.0];
101        let r = one_way_anova(&[&g1, &g2, &g3]).unwrap();
102        assert!(
103            r.p_value > 0.2,
104            "similar groups should not be significant: p={}",
105            r.p_value
106        );
107    }
108
109    #[test]
110    fn matches_known_worked_example() {
111        // Classic textbook example: groups (6,8,4,5,3,4),(8,12,9,11,6,8),(13,9,11,8,7,12).
112        let a = [6.0, 8.0, 4.0, 5.0, 3.0, 4.0];
113        let b = [8.0, 12.0, 9.0, 11.0, 6.0, 8.0];
114        let c = [13.0, 9.0, 11.0, 8.0, 7.0, 12.0];
115        let r = one_way_anova(&[&a, &b, &c]).unwrap();
116        // Known result: F ≈ 9.26, p ≈ 0.0026.
117        assert!((r.f_statistic - 9.264).abs() < 0.05, "F={}", r.f_statistic);
118        assert!((r.p_value - 0.00256).abs() < 5e-4, "p={}", r.p_value);
119    }
120
121    #[test]
122    fn guards_degenerate_input() {
123        assert!(one_way_anova(&[&[1.0, 2.0][..]]).is_none()); // < 2 groups
124        assert!(one_way_anova(&[&[1.0][..], &[2.0][..]]).is_none()); // n_total == k
125    }
126}