use super::{f_upper_log_tail, f_upper_tail, len_f64};
use crate::error::{Error, Result};
use crate::tests_stat::TestResult;
#[derive(Debug, Clone)]
pub struct TwoWayAnovaResult {
pub factor_a: TestResult,
pub factor_b: TestResult,
pub interaction: TestResult,
pub ss_within: f64,
pub df_within: f64,
}
pub fn two_way_anova(cells: &[Vec<Vec<f64>>]) -> Result<TwoWayAnovaResult> {
let a = cells.len();
let first_row = cells.first().ok_or(Error::InsufficientData)?;
let b = first_row.len();
if a < 2 || b < 2 {
return Err(Error::InsufficientData);
}
let n = first_row.first().map_or(0, Vec::len);
if n < 2 {
return Err(Error::InsufficientData);
}
for row in cells {
if row.len() != b {
return Err(Error::InvalidInput(
"unbalanced design: factor B levels differ across factor A".to_owned(),
));
}
if row.iter().any(|cell| cell.len() != n) {
return Err(Error::InvalidInput(
"unbalanced design: unequal replicate counts".to_owned(),
));
}
}
let (a_f, b_f, n_f) = (len_f64(a), len_f64(b), len_f64(n));
let cells_per_a = b_f * n_f;
let cells_per_b = a_f * n_f;
let total = a_f * b_f * n_f;
let row_sums: Vec<f64> = cells.iter().map(|row| row.iter().flatten().sum()).collect();
let col_sums: Vec<f64> = (0..b)
.map(|j| cells.iter().filter_map(|row| row.get(j)).flatten().sum())
.collect();
let grand_mean = row_sums.iter().sum::<f64>() / total;
let row_means: Vec<f64> = row_sums.iter().map(|&s| s / cells_per_a).collect();
let col_means: Vec<f64> = col_sums.iter().map(|&s| s / cells_per_b).collect();
let ss_a = cells_per_a
* row_means
.iter()
.map(|&rm| (rm - grand_mean) * (rm - grand_mean))
.sum::<f64>();
let ss_b = cells_per_b
* col_means
.iter()
.map(|&cm| (cm - grand_mean) * (cm - grand_mean))
.sum::<f64>();
let mut ss_interaction = 0.0;
let mut ss_within = 0.0;
for (row, &rm) in cells.iter().zip(&row_means) {
for (cell, &cm) in row.iter().zip(&col_means) {
let cell_mean = cell.iter().sum::<f64>() / n_f;
let interaction_dev = cell_mean - rm - cm + grand_mean;
ss_interaction = interaction_dev.mul_add(interaction_dev, ss_interaction);
ss_within += cell
.iter()
.map(|&y| (y - cell_mean) * (y - cell_mean))
.sum::<f64>();
}
}
ss_interaction *= n_f;
if ss_within <= 0.0 {
return Err(Error::DegenerateInput(
"zero within-cell variation".to_owned(),
));
}
let (df_a, df_b) = (a - 1, b - 1);
let df_interaction = df_a * df_b;
let df_within = a * b * (n - 1);
let ms_within = ss_within / len_f64(df_within);
Ok(TwoWayAnovaResult {
factor_a: effect_result(ss_a, df_a, ss_within, ms_within, df_within),
factor_b: effect_result(ss_b, df_b, ss_within, ms_within, df_within),
interaction: effect_result(
ss_interaction,
df_interaction,
ss_within,
ms_within,
df_within,
),
ss_within,
df_within: len_f64(df_within),
})
}
fn effect_result(
ss: f64,
df_effect: usize,
ss_within: f64,
ms_within: f64,
df_within: usize,
) -> TestResult {
let df_num = i64::try_from(df_effect).unwrap_or(i64::MAX);
let df_den = i64::try_from(df_within).unwrap_or(i64::MAX);
let f = (ss / len_f64(df_effect)) / ms_within;
TestResult {
statistic: f,
p_value: f_upper_tail(f, df_num, df_den),
log_p_value: Some(f_upper_log_tail(f, df_num, df_den)),
df: Some(len_f64(df_effect)),
effect_size: Some(ss / (ss + ss_within)),
}
}
#[cfg(kani)]
mod verification {
include!("two_way_anova_verification.rs");
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn one_level_factor_a_is_insufficient() {
let cells = vec![vec![vec![1.0, 2.0], vec![3.0, 4.0]]];
assert!(
matches!(two_way_anova(&cells), Err(Error::InsufficientData)),
"one A level should be InsufficientData"
);
}
#[test]
fn unbalanced_cells_are_invalid() {
let cells = vec![
vec![vec![1.0, 2.0], vec![3.0, 4.0]],
vec![vec![5.0, 6.0], vec![7.0]],
];
assert!(
matches!(two_way_anova(&cells), Err(Error::InvalidInput(_))),
"unequal replicate counts should be InvalidInput"
);
}
#[test]
fn single_replicate_is_insufficient() {
let cells = vec![vec![vec![1.0], vec![2.0]], vec![vec![3.0], vec![4.0]]];
assert!(
matches!(two_way_anova(&cells), Err(Error::InsufficientData)),
"one replicate per cell should be InsufficientData"
);
}
fn rel_close(got: f64, want: f64, rel: f64) -> bool {
(got - want).abs() <= rel * want.abs()
}
#[test]
fn golden_2x2x3_matches_statsmodels() -> Result<()> {
let cells = vec![
vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]],
vec![vec![7.0, 8.0, 9.0], vec![10.0, 11.0, 13.0]],
];
let r = two_way_anova(&cells)?;
assert!(
rel_close(r.ss_within, 10.666_666_666_666_666, 1e-12),
"ss_within {}",
r.ss_within
);
assert!(
(r.df_within - 8.0).abs() < 1e-12,
"df_within {}",
r.df_within
);
assert!(
rel_close(r.factor_a.statistic, 85.562_500_000_000_1, 1e-9),
"F_a {}",
r.factor_a.statistic
);
assert!(
(r.factor_a.p_value - 1.514_373_310_913_598e-5).abs() < 1e-9,
"p_a {}",
r.factor_a.p_value
);
assert_eq!(r.factor_a.df, Some(1.0), "df_a");
assert!(
r.factor_a
.effect_size
.is_some_and(|e| rel_close(e, 0.914_495_657_982_632, 1e-9)),
"eta_a {:?}",
r.factor_a.effect_size
);
assert!(
rel_close(r.factor_b.statistic, 22.562_500_000_000_032, 1e-9),
"F_b {}",
r.factor_b.statistic
);
assert!(
(r.factor_b.p_value - 0.001_445_249_130_436_945_8).abs() < 1e-9,
"p_b {}",
r.factor_b.p_value
);
assert_eq!(r.factor_b.df, Some(1.0), "df_b");
assert!(
rel_close(r.interaction.statistic, 0.062_500_000_000_002_04, 1e-9),
"F_ab {}",
r.interaction.statistic
);
assert!(
(r.interaction.p_value - 0.808_887_445_493_532_1).abs() < 1e-9,
"p_ab {}",
r.interaction.p_value
);
assert_eq!(r.interaction.df, Some(1.0), "df_ab");
Ok(())
}
#[test]
fn golden_3x2x4_matches_statsmodels() -> Result<()> {
let cells = vec![
vec![vec![2.0, 3.0, 5.0, 4.0], vec![11.0, 10.0, 12.0, 9.0]],
vec![vec![1.0, 0.0, 2.0, 3.0], vec![8.0, 9.0, 7.0, 10.0]],
vec![vec![9.0, 11.0, 10.0, 12.0], vec![6.0, 5.0, 7.0, 4.0]],
];
let r = two_way_anova(&cells)?;
assert!(
rel_close(r.ss_within, 30.0, 1e-12),
"ss_within {}",
r.ss_within
);
assert!(
(r.df_within - 18.0).abs() < 1e-12,
"df_within {}",
r.df_within
);
assert!(
rel_close(r.factor_a.statistic, 11.200_000_000_000_03, 1e-9),
"F_a {}",
r.factor_a.statistic
);
assert!(
(r.factor_a.p_value - 0.000_691_863_245_740_617_3).abs() < 1e-9,
"p_a {}",
r.factor_a.p_value
);
assert_eq!(r.factor_a.df, Some(2.0), "df_a");
assert!(
r.factor_a
.effect_size
.is_some_and(|e| rel_close(e, 0.554_455_445_544_555_2, 1e-9)),
"eta_a {:?}",
r.factor_a.effect_size
);
assert!(
rel_close(r.factor_b.statistic, 32.400_000_000_000_06, 1e-9),
"F_b {}",
r.factor_b.statistic
);
assert!(
(r.factor_b.p_value - 2.130_126_127_468_567_5e-5).abs() < 1e-9,
"p_b {}",
r.factor_b.p_value
);
assert_eq!(r.factor_b.df, Some(1.0), "df_b");
assert!(
r.factor_b
.effect_size
.is_some_and(|e| rel_close(e, 0.642_857_142_857_143_2, 1e-9)),
"eta_b {:?}",
r.factor_b.effect_size
);
assert!(
rel_close(r.interaction.statistic, 57.599_999_999_999_98, 1e-9),
"F_ab {}",
r.interaction.statistic
);
assert!(
rel_close(r.interaction.p_value, 1.502_846_147_704_456_7e-8, 1e-8),
"p_ab {}",
r.interaction.p_value
);
assert_eq!(r.interaction.df, Some(2.0), "df_ab");
assert!(
r.interaction
.effect_size
.is_some_and(|e| rel_close(e, 0.864_864_864_864_864_8, 1e-9)),
"eta_ab {:?}",
r.interaction.effect_size
);
Ok(())
}
#[test]
fn partial_eta_squared_matches_definition() -> Result<()> {
let cells = vec![
vec![vec![2.0, 3.0, 5.0, 4.0], vec![11.0, 10.0, 12.0, 9.0]],
vec![vec![1.0, 0.0, 2.0, 3.0], vec![8.0, 9.0, 7.0, 10.0]],
vec![vec![9.0, 11.0, 10.0, 12.0], vec![6.0, 5.0, 7.0, 4.0]],
];
let r = two_way_anova(&cells)?;
let ms_within = r.ss_within / r.df_within;
for eff in [&r.factor_a, &r.factor_b, &r.interaction] {
let df_eff = eff.df.ok_or(Error::InsufficientData)?;
let ss_eff = eff.statistic * df_eff * ms_within;
let want = ss_eff / (ss_eff + r.ss_within);
assert!(
eff.effect_size.is_some_and(|e| rel_close(e, want, 1e-9)),
"partial eta^2 {:?} != {want}",
eff.effect_size
);
}
Ok(())
}
}