use crate::correlation::{validate_correlation_input, CorrelationMethod, CorrelationResult};
use crate::error::Result;
use statrs::distribution::{ContinuousCDF, Normal};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum KendallVariant {
TauA,
#[default]
TauB,
TauC,
}
pub fn kendall(x: &[f64], y: &[f64], variant: KendallVariant) -> Result<CorrelationResult> {
let n = validate_correlation_input(x, y)?;
let (concordant, discordant, ties_x, ties_y, ties_xy) = count_pairs(x, y);
let n_pairs = (n * (n - 1)) / 2;
let tau = match variant {
KendallVariant::TauA => {
(concordant as f64 - discordant as f64) / n_pairs as f64
}
KendallVariant::TauB => {
let c_minus_d = concordant as f64 - discordant as f64;
let denom1 = n_pairs as f64 - ties_x as f64;
let denom2 = n_pairs as f64 - ties_y as f64;
if denom1 == 0.0 || denom2 == 0.0 {
0.0
} else {
c_minus_d / (denom1 * denom2).sqrt()
}
}
KendallVariant::TauC => {
let unique_x = count_unique(x);
let unique_y = count_unique(y);
let m = unique_x.min(unique_y);
if m <= 1 {
0.0
} else {
let c_minus_d = concordant as f64 - discordant as f64;
2.0 * c_minus_d * m as f64 / ((n * n * (m - 1)) as f64)
}
}
};
let tau = tau.clamp(-1.0, 1.0);
let (z_stat, p_value) = compute_kendall_significance(
tau, n, concordant, discordant, ties_x, ties_y, ties_xy, variant,
);
Ok(CorrelationResult {
estimate: tau,
statistic: z_stat,
df: None, p_value,
conf_int: None, method: CorrelationMethod::Kendall,
n,
})
}
fn count_pairs(x: &[f64], y: &[f64]) -> (usize, usize, usize, usize, usize) {
let n = x.len();
let mut concordant = 0usize;
let mut discordant = 0usize;
let mut ties_x = 0usize; let mut ties_y = 0usize; let mut ties_xy = 0usize;
for i in 0..n {
for j in (i + 1)..n {
let dx = x[i] - x[j];
let dy = y[i] - y[j];
let tied_x = dx == 0.0;
let tied_y = dy == 0.0;
if tied_x {
ties_x += 1;
}
if tied_y {
ties_y += 1;
}
if tied_x && tied_y {
ties_xy += 1;
}
if !tied_x && !tied_y {
if (dx > 0.0 && dy > 0.0) || (dx < 0.0 && dy < 0.0) {
concordant += 1;
} else {
discordant += 1;
}
}
}
}
(concordant, discordant, ties_x, ties_y, ties_xy)
}
fn count_unique(data: &[f64]) -> usize {
let mut sorted = data.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap());
sorted.dedup();
sorted.len()
}
#[allow(clippy::too_many_arguments)]
fn compute_kendall_significance(
_tau: f64,
n: usize,
concordant: usize,
discordant: usize,
ties_x: usize,
ties_y: usize,
_ties_xy: usize,
_variant: KendallVariant,
) -> (f64, f64) {
let n_f = n as f64;
let n_pairs = (n * (n - 1)) / 2;
let s = concordant as f64 - discordant as f64;
let variance = if ties_x == 0 && ties_y == 0 {
n_f * (n_f - 1.0) * (2.0 * n_f + 5.0) / 18.0
} else {
let t1 = ties_x as f64;
let t2 = ties_y as f64;
let n0 = n_pairs as f64;
let _n1 = n0 - t1;
let _n2 = n0 - t2;
let v0 = n_f * (n_f - 1.0) * (2.0 * n_f + 5.0) / 18.0;
let adj = 1.0 - (t1 + t2) / (2.0 * n0);
v0 * adj * adj
};
let z_stat = if variance <= 0.0 {
0.0
} else {
s / variance.sqrt()
};
let p_value = if z_stat == 0.0 {
1.0
} else {
let normal = Normal::new(0.0, 1.0).unwrap();
2.0 * normal.sf(z_stat.abs())
};
(z_stat, p_value)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_kendall_basic() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let y = vec![2.0, 4.0, 6.0, 8.0, 10.0];
let result = kendall(&x, &y, KendallVariant::TauB).unwrap();
assert!((result.estimate - 1.0).abs() < 1e-10);
assert_eq!(result.method, CorrelationMethod::Kendall);
assert_eq!(result.n, 5);
}
#[test]
fn test_kendall_negative() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let y = vec![5.0, 4.0, 3.0, 2.0, 1.0];
let result = kendall(&x, &y, KendallVariant::TauB).unwrap();
assert!((result.estimate - (-1.0)).abs() < 1e-10);
}
#[test]
fn test_kendall_tau_a_vs_tau_b() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let y = vec![1.0, 3.0, 2.0, 5.0, 4.0];
let tau_a = kendall(&x, &y, KendallVariant::TauA).unwrap();
let tau_b = kendall(&x, &y, KendallVariant::TauB).unwrap();
assert!((tau_a.estimate - tau_b.estimate).abs() < 1e-10);
}
#[test]
fn test_kendall_with_ties() {
let x = vec![1.0, 2.0, 2.0, 4.0, 5.0, 3.0];
let y = vec![1.0, 3.0, 2.0, 4.0, 5.0, 4.0];
let result = kendall(&x, &y, KendallVariant::TauB).unwrap();
assert!(result.estimate > 0.0);
assert!(result.estimate < 1.0);
}
#[test]
fn test_kendall_zero_correlation() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let y = vec![3.0, 1.0, 4.0, 2.0, 6.0, 5.0];
let result = kendall(&x, &y, KendallVariant::TauB).unwrap();
assert!(result.estimate.abs() < 0.5);
}
}