use crate::error::{Error, Result};
use crate::special::ln_gamma;
use crate::tests_stat::parametric::floor_to_i64;
use crate::tests_stat::{Alternative, TestResult};
pub fn fisher_exact(table: &[Vec<f64>], alternative: Alternative) -> Result<TestResult> {
let cell = cells(table)?;
let observed = cell.a;
let row1 = cell.a + cell.b;
let row2 = cell.c + cell.d;
let col1 = cell.a + cell.c;
let lo = floor_to_i64((col1 - row2).max(0.0));
let hi = floor_to_i64(col1.min(row1));
let p_observed = log_hypergeom(observed, row1, row2, col1).exp();
let threshold = p_observed * (1.0 + 1e-7);
let mut p_two = 0.0;
let mut p_less = 0.0;
let mut p_greater = 0.0;
for step in lo..=hi {
let k = f64::from(i32::try_from(step).unwrap_or(i32::MAX));
let pk = log_hypergeom(k, row1, row2, col1).exp();
if k <= observed + 0.5 {
p_less += pk;
}
if k >= observed - 0.5 {
p_greater += pk;
}
if pk <= threshold {
p_two += pk;
}
}
let odds_ratio = odds(cell.a, cell.b, cell.c, cell.d);
let p_value = match alternative {
Alternative::Less => p_less.clamp(0.0, 1.0),
Alternative::Greater => p_greater.clamp(0.0, 1.0),
Alternative::TwoSided => p_two.clamp(0.0, 1.0),
};
Ok(TestResult {
statistic: odds_ratio,
p_value,
log_p_value: None,
df: None,
effect_size: Some(odds_ratio),
})
}
#[derive(Debug, Clone, Copy)]
struct Cells {
a: f64,
b: f64,
c: f64,
d: f64,
}
fn cells(table: &[Vec<f64>]) -> Result<Cells> {
let bad = || Error::InvalidInput("Fisher exact requires a 2x2 table".to_owned());
if table.len() != 2 {
return Err(bad());
}
let row0 = table.first().ok_or_else(bad)?;
let row1 = table.get(1).ok_or_else(bad)?;
if row0.len() != 2 || row1.len() != 2 {
return Err(bad());
}
let cell = Cells {
a: *row0.first().ok_or_else(bad)?,
b: *row0.get(1).ok_or_else(bad)?,
c: *row1.first().ok_or_else(bad)?,
d: *row1.get(1).ok_or_else(bad)?,
};
for value in [cell.a, cell.b, cell.c, cell.d] {
if !value.is_finite() || value < 0.0 {
return Err(Error::InvalidInput(
"counts must be finite and non-negative".to_owned(),
));
}
}
Ok(cell)
}
fn odds(a: f64, b: f64, c: f64, d: f64) -> f64 {
let denom = b * c;
if denom == 0.0 {
if a * d == 0.0 {
f64::NAN
} else {
f64::INFINITY
}
} else {
a * d / denom
}
}
fn log_hypergeom(k: f64, row1: f64, row2: f64, col1: f64) -> f64 {
log_choose(row1, k) + log_choose(row2, col1 - k) - log_choose(row1 + row2, col1)
}
fn log_choose(n: f64, k: f64) -> f64 {
if k < 0.0 || k > n {
return f64::NEG_INFINITY;
}
ln_gamma(n + 1.0) - ln_gamma(k + 1.0) - ln_gamma(n - k + 1.0)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn balanced_table_odds_ratio_is_one() -> Result<()> {
let table = vec![vec![5.0, 5.0], vec![5.0, 5.0]];
let r = fisher_exact(&table, Alternative::TwoSided)?;
assert!(
(r.statistic - 1.0).abs() < 1e-12,
"odds ratio was {}",
r.statistic
);
Ok(())
}
#[test]
fn non_2x2_is_invalid() {
let table = vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]];
assert!(matches!(
fisher_exact(&table, Alternative::TwoSided),
Err(Error::InvalidInput(_))
));
}
}