use super::{len_f64, mean, one_way_anova, variance};
use crate::error::{Error, Result};
use crate::tests_stat::{TestResult, chi_squared_upper_log_tail, chi_squared_upper_tail};
pub fn levene(groups: &[&[f64]]) -> Result<TestResult> {
if groups.len() < 2 {
return Err(Error::InsufficientData);
}
let deviations: Vec<Vec<f64>> = groups
.iter()
.map(|g| {
let gm = mean(g);
g.iter().map(|&x| (x - gm).abs()).collect()
})
.collect();
let views: Vec<&[f64]> = deviations.iter().map(Vec::as_slice).collect();
let anova = one_way_anova(&views)?;
Ok(TestResult {
statistic: anova.statistic,
p_value: anova.p_value,
log_p_value: anova.log_p_value,
df: anova.df,
effect_size: None,
})
}
pub fn bartlett(groups: &[&[f64]]) -> Result<TestResult> {
let k = groups.len();
if k < 2 || groups.iter().any(|g| g.len() < 2) {
return Err(Error::InsufficientData);
}
let mut total_df = 0.0;
let mut pooled_num = 0.0;
let mut sum_log = 0.0;
let mut sum_inv_df = 0.0;
for group in groups {
let var = variance(group);
if var <= 0.0 {
return Err(Error::DegenerateInput("zero-variance group".to_owned()));
}
let df = len_f64(group.len() - 1);
total_df += df;
pooled_num += df * var;
sum_log += df * var.ln();
sum_inv_df += 1.0 / df;
}
let pooled_var = pooled_num / total_df;
let numerator = total_df * pooled_var.ln() - sum_log;
let k_f = len_f64(k);
let correction = 1.0 + (sum_inv_df - 1.0 / total_df) / (3.0 * (k_f - 1.0));
let statistic = numerator / correction;
let df_int = i64::try_from(k - 1).unwrap_or(i64::MAX);
let p_value = chi_squared_upper_tail(statistic, df_int);
Ok(TestResult {
statistic,
p_value,
log_p_value: Some(chi_squared_upper_log_tail(statistic, df_int)),
df: Some(k_f - 1.0),
effect_size: None,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn levene_zero_for_equal_spread() -> Result<()> {
let a = [1.0, 2.0, 3.0, 4.0];
let b = [5.0, 6.0, 7.0, 8.0];
let r = levene(&[&a[..], &b[..]])?;
assert!(r.statistic.abs() < 1e-9, "Levene F was {}", r.statistic);
Ok(())
}
#[test]
fn bartlett_rejects_zero_variance() {
let a = [2.0, 2.0, 2.0];
let b = [1.0, 2.0, 3.0];
assert!(matches!(
bartlett(&[&a[..], &b[..]]),
Err(Error::DegenerateInput(_))
));
}
}