use scirs2_core::numeric::Float;
use std::cmp::Ordering;
pub fn total_order<A: Float>(a: &A, b: &A) -> Ordering {
crate::utils::total_order(a, b)
}
pub fn sort_ascending<A: Float>(values: &mut [A]) {
values.sort_by(total_order);
}
pub fn mean<A: Float>(values: &[A]) -> Option<A> {
if values.is_empty() {
return None;
}
let sum = values.iter().fold(A::zero(), |acc, &v| acc + v);
A::from(values.len()).map(|n| sum / n)
}
pub fn population_variance<A: Float>(values: &[A]) -> Option<A> {
let mu = mean(values)?;
let n = A::from(values.len())?;
let ss = values.iter().fold(A::zero(), |acc, &v| {
let d = v - mu;
acc + d * d
});
Some(ss / n)
}
pub fn sample_std_dev<A: Float>(values: &[A]) -> Option<A> {
if values.len() < 2 {
return None;
}
let mu = mean(values)?;
let denom = A::from(values.len() - 1)?;
let ss = values.iter().fold(A::zero(), |acc, &v| {
let d = v - mu;
acc + d * d
});
Some((ss / denom).sqrt())
}
pub fn quantile_in_place<A: Float>(values: &mut [A], p: f64) -> Option<A> {
let n = values.len();
if n == 0 {
return None;
}
let p = p.clamp(0.0, 1.0);
let rank = (p * n as f64).ceil();
let index = if rank < 1.0 {
0
} else {
((rank as usize) - 1).min(n - 1)
};
let (_, nth, _) = values.select_nth_unstable_by(index, total_order);
Some(*nth)
}
pub fn median_in_place<A: Float>(values: &mut [A]) -> Option<A> {
let n = values.len();
if n == 0 {
return None;
}
if n % 2 == 1 {
let (_, nth, _) = values.select_nth_unstable_by(n / 2, total_order);
return Some(*nth);
}
let (lower, upper_mid, _) = values.select_nth_unstable_by(n / 2, total_order);
let upper = *upper_mid;
let lower_mid = lower.iter().copied().reduce(|a, b| {
if total_order(&b, &a) == Ordering::Greater {
b
} else {
a
}
})?;
let two = A::from(2.0)?;
Some((lower_mid + upper) / two)
}
pub fn median<A: Float>(values: &[A]) -> Option<A> {
let mut scratch: Vec<A> = values.to_vec();
median_in_place(&mut scratch)
}
pub fn standard_normal_sf(z: f64) -> Result<f64, String> {
let dist = scirs2_stats::distributions::normal::Normal::new(0.0_f64, 1.0_f64)
.map_err(|e| format!("standard normal construction failed: {e}"))?;
Ok((1.0 - dist.cdf(z)).clamp(0.0, 1.0))
}
pub fn normal_two_sided_p(z: f64) -> Result<f64, String> {
let upper = standard_normal_sf(z.abs())?;
Ok((2.0 * upper).clamp(0.0, 1.0))
}
fn ln_gamma(x: f64) -> f64 {
const COEFFICIENTS: [f64; 9] = [
0.999_999_999_999_810,
676.520_368_121_885,
-1_259.139_216_722_403,
771.323_428_777_653,
-176.615_029_162_141,
12.507_343_278_687,
-0.138_571_095_265_720,
0.000_009_984_369_578,
0.000_000_150_563_274,
];
if x < 0.5 {
return (std::f64::consts::PI / (std::f64::consts::PI * x).sin()).ln() - ln_gamma(1.0 - x);
}
let x = x - 1.0;
let mut series = COEFFICIENTS[0];
for (index, coefficient) in COEFFICIENTS.iter().enumerate().skip(1) {
series += coefficient / (x + index as f64);
}
let t = x + 7.5;
0.5 * (2.0 * std::f64::consts::PI).ln() + (x + 0.5) * t.ln() - t + series.ln()
}
fn regularized_upper_gamma(a: f64, x: f64) -> Result<f64, String> {
if !a.is_finite() || a <= 0.0 || !x.is_finite() || x < 0.0 {
return Err(format!(
"regularized_upper_gamma requires a > 0 and x >= 0, got a={a}, x={x}"
));
}
if x == 0.0 {
return Ok(1.0);
}
let log_prefactor = -x + a * x.ln() - ln_gamma(a);
if x < a + 1.0 {
let mut term = 1.0 / a;
let mut sum = term;
let mut n = a;
for _ in 0..1000 {
n += 1.0;
term *= x / n;
sum += term;
if term.abs() < sum.abs() * 1e-16 {
break;
}
}
let lower = sum * log_prefactor.exp();
return Ok((1.0 - lower).clamp(0.0, 1.0));
}
let tiny = 1e-300_f64;
let mut b = x + 1.0 - a;
let mut c = 1.0 / tiny;
let mut d = 1.0 / b;
let mut h = d;
for i in 1..1000 {
let an = -(i as f64) * (i as f64 - a);
b += 2.0;
d = an * d + b;
if d.abs() < tiny {
d = tiny;
}
c = b + an / c;
if c.abs() < tiny {
c = tiny;
}
d = 1.0 / d;
let delta = d * c;
h *= delta;
if (delta - 1.0).abs() < 1e-16 {
break;
}
}
Ok((log_prefactor.exp() * h).clamp(0.0, 1.0))
}
pub fn chi_square_sf(x: f64, df: f64) -> Result<f64, String> {
if !(df.is_finite() && df > 0.0) {
return Err(format!("chi-square requires df > 0, got {df}"));
}
if x <= 0.0 {
return Ok(1.0);
}
regularized_upper_gamma(df / 2.0, x / 2.0)
}
pub fn kolmogorov_sf(z: f64) -> f64 {
if !z.is_finite() || z <= 0.0 {
return 1.0;
}
let a2 = -2.0 * z * z;
let mut sign = 2.0_f64;
let mut sum = 0.0_f64;
let mut previous_magnitude = 0.0_f64;
for k in 1..=200_u32 {
let term = sign * (a2 * f64::from(k * k)).exp();
sum += term;
let magnitude = term.abs();
if magnitude <= 1e-8 * previous_magnitude || magnitude <= 1e-16 * sum.abs() {
break;
}
previous_magnitude = magnitude;
sign = -sign;
}
sum.clamp(0.0, 1.0)
}
pub fn ks_two_sample_p(d: f64, n1: usize, n2: usize) -> f64 {
if n1 == 0 || n2 == 0 {
return 1.0;
}
let n1 = n1 as f64;
let n2 = n2 as f64;
let effective_n = (n1 * n2) / (n1 + n2);
kolmogorov_sf(effective_n.sqrt() * d)
}
pub fn ks_statistic<A: Float>(sample_a: &[A], sample_b: &[A]) -> Option<f64> {
if sample_a.is_empty() || sample_b.is_empty() {
return None;
}
let mut a: Vec<A> = sample_a.to_vec();
let mut b: Vec<A> = sample_b.to_vec();
sort_ascending(&mut a);
sort_ascending(&mut b);
let na = a.len();
let nb = b.len();
let mut i = 0usize;
let mut j = 0usize;
let mut max_diff = 0.0_f64;
while i < na || j < nb {
let next = match (a.get(i), b.get(j)) {
(Some(x), Some(y)) => {
if total_order(x, y) == Ordering::Greater {
*y
} else {
*x
}
}
(Some(x), None) => *x,
(None, Some(y)) => *y,
(None, None) => break,
};
while i < na && total_order(&a[i], &next) != Ordering::Greater {
i += 1;
}
while j < nb && total_order(&b[j], &next) != Ordering::Greater {
j += 1;
}
let ecdf_a = i as f64 / na as f64;
let ecdf_b = j as f64 / nb as f64;
let diff = (ecdf_a - ecdf_b).abs();
if diff > max_diff {
max_diff = diff;
}
}
Some(max_diff)
}
pub fn finite_range<A: Float>(values: &[A]) -> Option<(f64, f64)> {
let mut min = f64::INFINITY;
let mut max = f64::NEG_INFINITY;
for value in values {
let Some(v) = value.to_f64() else { continue };
if !v.is_finite() {
continue;
}
if v < min {
min = v;
}
if v > max {
max = v;
}
}
if min.is_finite() && max.is_finite() {
Some((min, max))
} else {
None
}
}
pub fn histogram_counts<A: Float>(values: &[A], min: f64, max: f64, bins: usize) -> Vec<f64> {
let bins = bins.max(1);
let mut counts = vec![0.0_f64; bins];
let width = if max > min {
(max - min) / bins as f64
} else {
0.0
};
for value in values {
let Some(v) = value.to_f64() else { continue };
if !v.is_finite() {
continue;
}
let index = if width > 0.0 {
(((v - min) / width).floor().max(0.0) as usize).min(bins - 1)
} else {
0
};
counts[index] += 1.0;
}
counts
}
pub fn smoothed_pmf(counts: &[f64], smoothing: f64) -> Vec<f64> {
let smoothing = smoothing.max(f64::MIN_POSITIVE);
let total: f64 = counts.iter().sum::<f64>() + smoothing * counts.len() as f64;
if total <= 0.0 {
let uniform = 1.0 / counts.len().max(1) as f64;
return vec![uniform; counts.len()];
}
counts.iter().map(|&c| (c + smoothing) / total).collect()
}
pub fn kl_divergence(p: &[f64], q: &[f64]) -> Result<f64, String> {
if p.len() != q.len() {
return Err("KL divergence requires equal-length distributions".to_string());
}
let mut sum = 0.0_f64;
for (&pi, &qi) in p.iter().zip(q.iter()) {
if pi <= 0.0 {
continue;
}
if qi <= 0.0 {
return Err("KL divergence is undefined for a zero reference bin".to_string());
}
sum += pi * (pi / qi).ln();
}
Ok(sum.max(0.0))
}
pub fn js_divergence(p: &[f64], q: &[f64]) -> Result<f64, String> {
if p.len() != q.len() {
return Err("JS divergence requires equal-length distributions".to_string());
}
let mixture: Vec<f64> = p
.iter()
.zip(q.iter())
.map(|(&pi, &qi)| 0.5 * (pi + qi))
.collect();
let left = kl_divergence(p, &mixture)?;
let right = kl_divergence(q, &mixture)?;
Ok((0.5 * left + 0.5 * right).clamp(0.0, std::f64::consts::LN_2))
}
pub fn hellinger_distance(p: &[f64], q: &[f64]) -> Result<f64, String> {
if p.len() != q.len() {
return Err("Hellinger distance requires equal-length distributions".to_string());
}
let bhattacharyya: f64 = p
.iter()
.zip(q.iter())
.map(|(&pi, &qi)| (pi.max(0.0) * qi.max(0.0)).sqrt())
.sum();
Ok((1.0 - bhattacharyya).max(0.0).sqrt())
}
pub fn wasserstein_1d<A: Float>(sample_a: &[A], sample_b: &[A]) -> Option<f64> {
if sample_a.is_empty() || sample_b.is_empty() {
return None;
}
let mut a: Vec<f64> = sample_a.iter().filter_map(|v| v.to_f64()).collect();
let mut b: Vec<f64> = sample_b.iter().filter_map(|v| v.to_f64()).collect();
if a.is_empty() || b.is_empty() {
return None;
}
a.sort_by(|x, y| x.partial_cmp(y).unwrap_or(Ordering::Equal));
b.sort_by(|x, y| x.partial_cmp(y).unwrap_or(Ordering::Equal));
let na = a.len();
let nb = b.len();
let mut i = 0usize;
let mut j = 0usize;
let mut previous = a[0].min(b[0]);
let mut total = 0.0_f64;
while i < na || j < nb {
let next = match (a.get(i), b.get(j)) {
(Some(&x), Some(&y)) => x.min(y),
(Some(&x), None) => x,
(None, Some(&y)) => y,
(None, None) => break,
};
let ecdf_a = i as f64 / na as f64;
let ecdf_b = j as f64 / nb as f64;
total += (ecdf_a - ecdf_b).abs() * (next - previous);
previous = next;
while i < na && a[i] <= next {
i += 1;
}
while j < nb && b[j] <= next {
j += 1;
}
}
Some(total)
}
pub fn g_test_statistic(counts_a: &[f64], counts_b: &[f64]) -> Result<(f64, f64), String> {
if counts_a.len() != counts_b.len() {
return Err("G-test requires equal-length histograms".to_string());
}
let total_a: f64 = counts_a.iter().sum();
let total_b: f64 = counts_b.iter().sum();
let grand_total = total_a + total_b;
if total_a <= 0.0 || total_b <= 0.0 {
return Err("G-test requires both samples to be non-empty".to_string());
}
let mut g = 0.0_f64;
let mut non_empty_columns = 0usize;
for (&observed_a, &observed_b) in counts_a.iter().zip(counts_b.iter()) {
let column_total = observed_a + observed_b;
if column_total <= 0.0 {
continue;
}
non_empty_columns += 1;
let expected_a = column_total * total_a / grand_total;
let expected_b = column_total * total_b / grand_total;
if observed_a > 0.0 && expected_a > 0.0 {
g += observed_a * (observed_a / expected_a).ln();
}
if observed_b > 0.0 && expected_b > 0.0 {
g += observed_b * (observed_b / expected_b).ln();
}
}
let degrees_of_freedom = (non_empty_columns.saturating_sub(1)) as f64;
Ok((2.0 * g, degrees_of_freedom))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn total_order_is_nan_safe_and_total() {
let mut values = vec![3.0_f64, f64::NAN, 1.0, 2.0];
sort_ascending(&mut values);
assert_eq!(values[0], 1.0);
assert_eq!(values[1], 2.0);
assert_eq!(values[2], 3.0);
assert!(values[3].is_nan(), "NaN must sort last");
}
#[test]
fn median_matches_definition_for_odd_and_even_lengths() {
assert_eq!(median(&[5.0_f64, 1.0, 3.0]), Some(3.0));
assert_eq!(median(&[4.0_f64, 1.0, 3.0, 2.0]), Some(2.5));
assert_eq!(median::<f64>(&[]), None);
}
#[test]
fn quantile_uses_nearest_rank() {
let mut values = vec![1.0_f64, 2.0, 3.0, 4.0];
assert_eq!(quantile_in_place(&mut values, 0.0), Some(1.0));
let mut values = vec![1.0_f64, 2.0, 3.0, 4.0];
assert_eq!(quantile_in_place(&mut values, 1.0), Some(4.0));
let mut values = vec![1.0_f64, 2.0, 3.0, 4.0];
assert_eq!(quantile_in_place(&mut values, 0.5), Some(2.0));
}
#[test]
fn normal_survival_function_matches_known_quantiles() {
let p = standard_normal_sf(1.959_963_985).expect("normal sf");
assert!(
(p - 0.025).abs() < 1e-4,
"P(Z > 1.96) should be ~0.025, got {p}"
);
let p0 = standard_normal_sf(0.0).expect("normal sf");
assert!((p0 - 0.5).abs() < 1e-9);
}
#[test]
fn chi_square_survival_function_matches_known_quantiles() {
let p = chi_square_sf(3.841_458_8, 1.0).expect("chi2 sf");
assert!((p - 0.05).abs() < 1e-6, "expected 0.05, got {p}");
let p = chi_square_sf(5.991_464_5, 2.0).expect("chi2 sf");
assert!((p - 0.05).abs() < 1e-6, "expected 0.05, got {p}");
let p = chi_square_sf(24.995_79, 15.0).expect("chi2 sf");
assert!((p - 0.05).abs() < 1e-5, "expected 0.05, got {p}");
}
#[test]
fn chi_square_survival_function_is_correct_in_the_series_branch() {
let p = chi_square_sf(1.0, 4.0).expect("chi2 sf");
assert!((p - 0.909_796).abs() < 1e-5, "expected 0.909796, got {p}");
let p = chi_square_sf(2.0, 10.0).expect("chi2 sf");
assert!((p - 0.996_340).abs() < 1e-5, "expected 0.996340, got {p}");
assert!(chi_square_sf(1e-12, 3.0).expect("chi2 sf") > 0.999_999);
assert_eq!(chi_square_sf(0.0, 3.0).expect("chi2 sf"), 1.0);
}
#[test]
fn chi_square_survival_function_is_continuous_across_the_branch_switch() {
let df = 6.0_f64;
let a = df / 2.0; let switch = 2.0 * (a + 1.0);
let below = chi_square_sf(switch - 1e-7, df).expect("chi2 sf");
let above = chi_square_sf(switch + 1e-7, df).expect("chi2 sf");
assert!(
(below - above).abs() < 1e-7,
"the two branches disagree at the switch point ({below} vs {above})"
);
let mut previous = 1.0_f64;
for step in 1..=200 {
let x = step as f64 * 0.15;
let p = chi_square_sf(x, df).expect("chi2 sf");
assert!(
p <= previous + 1e-12,
"survival function increased at x = {x} ({previous} -> {p})"
);
previous = p;
}
}
#[test]
fn ln_gamma_matches_known_values() {
assert!(ln_gamma(1.0).abs() < 1e-12);
assert!(ln_gamma(2.0).abs() < 1e-12);
assert!((ln_gamma(5.0) - 24.0_f64.ln()).abs() < 1e-11);
assert!(
(ln_gamma(0.5) - std::f64::consts::PI.sqrt().ln()).abs() < 1e-11,
"ln G(0.5) = {}",
ln_gamma(0.5)
);
}
#[test]
fn ks_statistic_is_one_for_disjoint_samples() {
let d = ks_statistic(&[0.0_f64, 1.0, 2.0], &[10.0, 11.0, 12.0]).expect("ks");
assert!(
(d - 1.0).abs() < 1e-12,
"disjoint samples must give D = 1, got {d}"
);
let p = ks_two_sample_p(d, 3, 3);
assert!(
p < 0.5,
"D = 1 on n = 3 should be at least mildly significant, got {p}"
);
}
#[test]
fn ks_statistic_is_zero_for_identical_samples() {
let d = ks_statistic(&[1.0_f64, 2.0, 3.0], &[1.0, 2.0, 3.0]).expect("ks");
assert!(
d.abs() < 1e-12,
"identical samples must give D = 0, got {d}"
);
assert!((ks_two_sample_p(d, 3, 3) - 1.0).abs() < 1e-12);
}
#[test]
fn wasserstein_matches_closed_form_for_a_pure_shift() {
let a = [0.0_f64, 1.0, 2.0, 3.0];
let b = [5.0_f64, 6.0, 7.0, 8.0];
let w = wasserstein_1d(&a, &b).expect("w1");
assert!((w - 5.0).abs() < 1e-9, "expected W1 = 5, got {w}");
}
#[test]
fn divergences_are_zero_for_identical_distributions_and_positive_otherwise() {
let p = smoothed_pmf(&[10.0, 10.0, 10.0], 0.5);
let q = smoothed_pmf(&[10.0, 10.0, 10.0], 0.5);
assert!(kl_divergence(&p, &q).expect("kl") < 1e-12);
assert!(js_divergence(&p, &q).expect("js") < 1e-12);
assert!(hellinger_distance(&p, &q).expect("hellinger") < 1e-6);
let r = smoothed_pmf(&[30.0, 0.0, 0.0], 0.5);
assert!(kl_divergence(&p, &r).expect("kl") > 0.1);
assert!(js_divergence(&p, &r).expect("js") > 0.1);
assert!(hellinger_distance(&p, &r).expect("hellinger") > 0.1);
}
#[test]
fn g_test_is_insignificant_for_matching_histograms() {
let (g, df) = g_test_statistic(&[20.0, 20.0, 20.0], &[20.0, 20.0, 20.0]).expect("g-test");
assert!(g.abs() < 1e-9);
assert_eq!(df, 2.0);
let p = chi_square_sf(g, df.max(1.0)).expect("chi2");
assert!(
p > 0.9,
"identical histograms must not be significant, got {p}"
);
}
#[test]
fn g_test_is_significant_for_disjoint_histograms() {
let (g, df) = g_test_statistic(&[60.0, 0.0], &[0.0, 60.0]).expect("g-test");
assert!(
g > 100.0,
"disjoint histograms should give a large G, got {g}"
);
let p = chi_square_sf(g, df.max(1.0)).expect("chi2");
assert!(p < 1e-6, "expected an extremely small p-value, got {p}");
}
}