use super::normal_cdf;
use crate::error::{Error, Result};
use crate::tests_stat::exact::wilcoxon 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 = 25;
pub fn wilcoxon_signed_rank(
a: &[f64],
b: &[f64],
alternative: Alternative,
use_continuity: bool,
) -> Result<TestResult> {
if a.len() != b.len() {
return Err(Error::InvalidInput(
"paired samples differ in length".to_owned(),
));
}
let diffs: Vec<f64> = a
.iter()
.zip(b)
.map(|(&x, &y)| x - y)
.filter(|d| d.abs() > 0.0)
.collect();
if diffs.is_empty() {
return Err(Error::InsufficientData);
}
let abs: Vec<f64> = diffs.iter().map(|d| d.abs()).collect();
let ranks = mid_ranks(&abs);
let mut r_plus = 0.0;
let mut r_minus = 0.0;
for (d, r) in diffs.iter().zip(&ranks) {
if *d > 0.0 {
r_plus += r;
} else {
r_minus += r;
}
}
let count = len_f64(diffs.len());
let mean = count * (count + 1.0) * 0.25;
let tie = tie_correction(&abs);
let var = count * (count + 1.0) * (2.0f64.mul_add(count, 1.0)) / 24.0 - tie / 48.0;
if var <= 0.0 {
return Err(Error::InsufficientData);
}
let se = var.sqrt();
let statistic = match alternative {
Alternative::TwoSided => r_plus.min(r_minus),
_ => r_plus,
};
let p_value = p_value(r_plus, mean, se, alternative, use_continuity);
Ok(TestResult {
statistic,
p_value,
log_p_value: None,
df: None,
effect_size: None,
})
}
pub fn wilcoxon_signed_rank_mode(
a: &[f64],
b: &[f64],
alternative: Alternative,
use_continuity: bool,
mode: Mode,
) -> Result<TestResult> {
if a.len() != b.len() {
return Err(Error::InvalidInput(
"paired samples differ in length".to_owned(),
));
}
let diffs: Vec<f64> = a
.iter()
.zip(b)
.map(|(&x, &y)| x - y)
.filter(|d| d.abs() > 0.0)
.collect();
if diffs.is_empty() {
return Err(Error::InsufficientData);
}
let abs: Vec<f64> = diffs.iter().map(|d| d.abs()).collect();
let nonzero = diffs.len();
let has_ties = tie_correction(&abs) > 0.0;
let resolved = match mode.resolve(nonzero, 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 Wilcoxon requires untied absolute differences".to_owned(),
));
}
if matches!(resolved, Mode::Asymptotic) {
return wilcoxon_signed_rank(a, b, alternative, use_continuity);
}
let ranks = mid_ranks(&abs);
let mut r_plus = 0.0;
let mut r_minus = 0.0;
for (d, r) in diffs.iter().zip(&ranks) {
if *d > 0.0 {
r_plus += r;
} else {
r_minus += r;
}
}
let statistic = match alternative {
Alternative::TwoSided => r_plus.min(r_minus),
_ => r_plus,
};
let p_value = exact::p_value(r_plus, nonzero, alternative);
Ok(TestResult {
statistic,
p_value,
log_p_value: None,
df: None,
effect_size: None,
})
}
fn p_value(r_plus: f64, mean: f64, se: f64, alternative: Alternative, use_continuity: bool) -> f64 {
let diff = r_plus - mean;
let cc = if use_continuity { 0.5 } else { 0.0 };
match alternative {
Alternative::Greater => {
let z = (diff - cc) / se;
(1.0 - normal_cdf(z)).clamp(0.0, 1.0)
}
Alternative::Less => {
let z = (diff + cc) / se;
normal_cdf(z).clamp(0.0, 1.0)
}
Alternative::TwoSided => {
let z = (diff.abs() - cc) / se;
(2.0 * (1.0 - normal_cdf(z))).clamp(0.0, 1.0)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn length_mismatch_is_invalid() {
assert!(matches!(
wilcoxon_signed_rank(&[1.0, 2.0], &[1.0], Alternative::TwoSided, true),
Err(Error::InvalidInput(_))
));
}
#[test]
fn all_zero_differences_are_insufficient() {
let a = [3.0, 4.0];
assert!(matches!(
wilcoxon_signed_rank(&a, &a, Alternative::TwoSided, true),
Err(Error::InsufficientData)
));
}
}