use crate::error::{Result, StatError};
use statrs::distribution::{ContinuousCDF, Normal};
#[derive(Debug, Clone)]
pub struct ShapiroWilkResult {
pub statistic: f64,
pub p_value: f64,
}
pub fn shapiro_wilk(data: &[f64]) -> Result<ShapiroWilkResult> {
let n = data.len();
if n < 3 {
return Err(StatError::InsufficientData { needed: 3, got: n });
}
if n > 5000 {
return Err(StatError::InvalidParameter(
"Shapiro-Wilk test is limited to n <= 5000".to_string(),
));
}
let mut x = data.to_vec();
x.sort_by(|a, b| a.partial_cmp(b).unwrap());
let range = x[n - 1] - x[0];
if range < 1e-10 {
return Ok(ShapiroWilkResult {
statistic: 1.0,
p_value: 1.0,
});
}
let (w, p_value) = swilk(&x);
Ok(ShapiroWilkResult {
statistic: w,
p_value,
})
}
fn swilk(x: &[f64]) -> (f64, f64) {
let n = x.len();
let n_f = n as f64;
let mean: f64 = x.iter().sum::<f64>() / n_f;
let ss: f64 = x.iter().map(|xi| (xi - mean).powi(2)).sum();
if ss < 1e-30 {
return (1.0, 1.0);
}
let a = compute_coefficients(n);
let nn2 = n / 2;
let mut w_num = 0.0;
for i in 0..nn2 {
w_num += a[i] * (x[n - 1 - i] - x[i]);
}
let w = (w_num * w_num) / ss;
let w = w.clamp(0.0, 1.0);
let p_value = compute_p_value(w, n);
(w, p_value)
}
fn compute_order_statistics(n: usize) -> Vec<f64> {
let n_f = n as f64;
let normal = Normal::new(0.0, 1.0).unwrap();
(0..n)
.map(|i| {
let p = (i as f64 + 1.0 - 0.375) / (n_f + 0.25);
normal.inverse_cdf(p)
})
.collect()
}
fn normalize_coefficients(a: &mut [f64]) {
let a_sum_sq: f64 = a.iter().map(|x| x * x).sum();
if a_sum_sq > 1e-10 {
let target = 0.5;
let scale = (target / a_sum_sq).sqrt();
for ai in a.iter_mut() {
*ai *= scale;
}
}
}
fn compute_coefficients_small(m: &[f64], n: usize, nn2: usize) -> Vec<f64> {
let mut a = vec![0.0; nn2];
for i in 0..nn2 {
a[i] = m[n - 1 - i] - m[i];
}
normalize_coefficients(&mut a);
a
}
fn compute_first_two_coefficients(m: &[f64], sqrt_m2: f64, n: usize) -> (f64, f64) {
let sqrtn = (n as f64).sqrt();
let c1 = [
-2.706056, 4.434685, -2.07119, -0.147981, 0.221157, -0.0006714,
];
let an = m[n - 1] / sqrt_m2 + poly_eval(&c1, 1.0 / sqrtn);
let c2 = [-3.582633, 5.682633, -1.752461, -0.293762, 0.042981, 0.0];
let an1 = if n > 6 {
m[n - 2] / sqrt_m2 + poly_eval(&c2, 1.0 / sqrtn)
} else {
m[n - 2] / sqrt_m2
};
(an, an1)
}
fn compute_coefficients_large(m: &[f64], m2: f64, n: usize, nn2: usize) -> Vec<f64> {
let sqrt_m2 = m2.sqrt();
let (an, an1) = compute_first_two_coefficients(m, sqrt_m2, n);
let sum_first_two_sq = 2.0 * (an * an + an1 * an1);
let sum_middle_m_sq = m2 - 2.0 * m[n - 1].powi(2) - 2.0 * m[n - 2].powi(2);
let phi_sq = sum_middle_m_sq / (1.0 - sum_first_two_sq);
let phi = if phi_sq > 0.0 { phi_sq.sqrt() } else { 1.0 };
let mut a = vec![0.0; nn2];
a[0] = an;
if nn2 > 1 {
a[1] = an1;
}
for i in 2..nn2 {
a[i] = m[n - 1 - i] / phi;
}
normalize_coefficients(&mut a);
a
}
fn compute_coefficients(n: usize) -> Vec<f64> {
let nn2 = n / 2;
let m = compute_order_statistics(n);
let m2: f64 = m.iter().map(|x| x * x).sum();
if n <= 5 {
compute_coefficients_small(&m, n, nn2)
} else {
compute_coefficients_large(&m, m2, n, nn2)
}
}
fn poly_eval(c: &[f64; 6], u: f64) -> f64 {
c[0] * u.powi(5) + c[1] * u.powi(4) + c[2] * u.powi(3) + c[3] * u.powi(2) + c[4] * u + c[5]
}
fn compute_p_value(w: f64, n: usize) -> f64 {
let n_f = n as f64;
let normal = Normal::new(0.0, 1.0).unwrap();
let p = if n == 3 {
let pi = std::f64::consts::PI;
let p = 6.0 / pi * (w.sqrt().asin() - (3.0_f64 / 4.0).sqrt().asin());
p.clamp(0.0, 1.0)
} else if n <= 11 {
let y = if w >= 1.0 - 1e-10 {
return 1.0;
} else {
(1.0 - w).ln()
};
let ln_n = n_f.ln();
let mu = -1.2725 - 1.0521 * ln_n - 0.26758 * ln_n * ln_n;
let sigma = (0.4803 + 0.082676 * ln_n + 0.0030302 * ln_n * ln_n).exp();
let z = (y - mu) / sigma;
normal.sf(z)
} else {
let y = (1.0 - w).ln();
let ln_n = n_f.ln();
let mu = poly_mu_large(ln_n);
let sigma = poly_sigma_large(ln_n).exp();
let z = (y - mu) / sigma;
normal.sf(z)
};
p.clamp(0.0, 1.0)
}
fn poly_mu_large(ln_n: f64) -> f64 {
0.0038915 * ln_n.powi(3) - 0.083751 * ln_n.powi(2) - 0.31082 * ln_n - 1.5861
}
fn poly_sigma_large(ln_n: f64) -> f64 {
0.0030302 * ln_n.powi(2) - 0.082676 * ln_n - 0.4803
}