use super::normal_cdf;
use crate::error::{Error, Result};
use crate::tests_stat::exact::mann_whitney as exact;
use crate::tests_stat::parametric::len_f64;
use crate::tests_stat::ranks::{mid_ranks, tie_correction};
use crate::tests_stat::{Alternative, Mode, TestResult};
const EXACT_THRESHOLD: usize = 8;
pub fn mann_whitney_u(
a: &[f64],
b: &[f64],
alternative: Alternative,
use_continuity: bool,
) -> Result<TestResult> {
if a.is_empty() || b.is_empty() {
return Err(Error::EmptyInput);
}
let n1 = a.len();
let n2 = b.len();
let mut pooled = Vec::with_capacity(n1 + n2);
pooled.extend_from_slice(a);
pooled.extend_from_slice(b);
let ranks = mid_ranks(&pooled);
let rank_sum_a: f64 = ranks.iter().take(n1).sum();
let n1_f = len_f64(n1);
let n2_f = len_f64(n2);
let u1 = rank_sum_a - n1_f * (n1_f + 1.0) / 2.0;
let u2 = n1_f.mul_add(n2_f, -u1);
let n = n1_f + n2_f;
let tie = tie_correction(&pooled);
let variance = n1_f * n2_f / 12.0 * ((n + 1.0) - tie / (n * (n - 1.0)));
if variance <= 0.0 {
return Err(Error::DegenerateInput("all values tied".to_owned()));
}
let sigma = variance.sqrt();
let mu = n1_f * n2_f / 2.0;
let p_value = p_value(u1, u2, mu, sigma, alternative, use_continuity);
let rank_biserial = 1.0 - 2.0 * u2 / (n1_f * n2_f);
Ok(TestResult {
statistic: u1,
p_value,
log_p_value: None,
df: None,
effect_size: Some(rank_biserial),
})
}
pub fn mann_whitney_u_mode(
a: &[f64],
b: &[f64],
alternative: Alternative,
use_continuity: bool,
mode: Mode,
) -> Result<TestResult> {
if a.is_empty() || b.is_empty() {
return Err(Error::EmptyInput);
}
let n1 = a.len();
let n2 = b.len();
let mut pooled = Vec::with_capacity(n1 + n2);
pooled.extend_from_slice(a);
pooled.extend_from_slice(b);
let has_ties = tie_correction(&pooled) > 0.0;
let resolved = match mode.resolve(n1.max(n2), EXACT_THRESHOLD) {
Mode::Exact if has_ties && matches!(mode, Mode::Auto) => Mode::Asymptotic,
other => other,
};
if matches!(resolved, Mode::Exact) && has_ties {
return Err(Error::InvalidInput(
"exact Mann–Whitney requires untied data".to_owned(),
));
}
if matches!(resolved, Mode::Asymptotic) {
return mann_whitney_u(a, b, alternative, use_continuity);
}
let ranks = mid_ranks(&pooled);
let rank_sum_a: f64 = ranks.iter().take(n1).sum();
let n1_f = len_f64(n1);
let n2_f = len_f64(n2);
let u1 = rank_sum_a - n1_f * (n1_f + 1.0) / 2.0;
let u2 = n1_f.mul_add(n2_f, -u1);
let p_value = exact::p_value(u1, n1, n2, alternative);
let rank_biserial = 1.0 - 2.0 * u2 / (n1_f * n2_f);
Ok(TestResult {
statistic: u1,
p_value,
log_p_value: None,
df: None,
effect_size: Some(rank_biserial),
})
}
fn p_value(
u1: f64,
u2: f64,
mu: f64,
sigma: f64,
alternative: Alternative,
use_continuity: bool,
) -> f64 {
let cc = if use_continuity { 0.5 } else { 0.0 };
match alternative {
Alternative::Greater => {
let z = (u1 - mu - cc) / sigma;
(1.0 - normal_cdf(z)).clamp(0.0, 1.0)
}
Alternative::Less => {
let z = (u2 - mu - cc) / sigma;
(1.0 - normal_cdf(z)).clamp(0.0, 1.0)
}
Alternative::TwoSided => {
let u = u1.max(u2);
let z = (u - mu - cc) / sigma;
(2.0 * (1.0 - normal_cdf(z))).clamp(0.0, 1.0)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn separated_samples_are_significant() -> Result<()> {
let a = [1.0, 2.0, 3.0, 4.0];
let b = [10.0, 11.0, 12.0, 13.0];
let r = mann_whitney_u(&a, &b, Alternative::TwoSided, true)?;
assert!(r.p_value < 0.05, "p was {}", r.p_value);
Ok(())
}
#[test]
fn empty_sample_is_error() {
assert!(matches!(
mann_whitney_u(&[], &[1.0], Alternative::TwoSided, true),
Err(Error::EmptyInput)
));
}
}