use crate::correlation::{mean, validate_correlation_input};
use crate::error::{Result, StatError};
use statrs::distribution::{ContinuousCDF, StudentsT};
#[derive(Debug, Clone)]
pub struct PartialCorResult {
pub estimate: f64,
pub statistic: f64,
pub df: f64,
pub p_value: f64,
pub n: usize,
pub n_controls: usize,
pub method: String,
}
pub fn partial_cor(x: &[f64], y: &[f64], z: &[&[f64]]) -> Result<PartialCorResult> {
let n = validate_correlation_input(x, y)?;
for (i, zi) in z.iter().enumerate() {
if zi.len() != n {
return Err(StatError::InvalidParameter(format!(
"Control variable {} has length {}, expected {}",
i,
zi.len(),
n
)));
}
for (j, &val) in zi.iter().enumerate() {
if !val.is_finite() {
return Err(StatError::InvalidParameter(format!(
"Non-finite value in control variable {} at index {}",
i, j
)));
}
}
}
let k = z.len();
if k == 0 {
return compute_simple_correlation(x, y, n);
}
if n <= k + 2 {
return Err(StatError::InsufficientData {
needed: k + 3,
got: n,
});
}
let partial_r = if k == 1 {
compute_partial_cor_single(x, y, z[0])?
} else {
let z_rest: Vec<&[f64]> = z[..k - 1].to_vec();
let zk = z[k - 1];
let r_xy_zrest = partial_cor(x, y, &z_rest)?.estimate;
let r_xzk_zrest = partial_cor(x, zk, &z_rest)?.estimate;
let r_yzk_zrest = partial_cor(y, zk, &z_rest)?.estimate;
let numerator = r_xy_zrest - r_xzk_zrest * r_yzk_zrest;
let denom1 = (1.0 - r_xzk_zrest * r_xzk_zrest).sqrt();
let denom2 = (1.0 - r_yzk_zrest * r_yzk_zrest).sqrt();
if denom1 * denom2 > 1e-10 {
numerator / (denom1 * denom2)
} else {
0.0
}
};
let partial_r = partial_r.clamp(-1.0, 1.0);
let df = (n - k - 2) as f64;
let t_stat = if (1.0 - partial_r * partial_r).abs() < 1e-15 {
if partial_r > 0.0 {
f64::INFINITY
} else {
f64::NEG_INFINITY
}
} else {
partial_r * (df / (1.0 - partial_r * partial_r)).sqrt()
};
let p_value = if t_stat.is_infinite() {
0.0
} else if df > 0.0 {
let t_dist = StudentsT::new(0.0, 1.0, df).unwrap();
2.0 * t_dist.sf(t_stat.abs())
} else {
1.0
};
Ok(PartialCorResult {
estimate: partial_r,
statistic: t_stat,
df,
p_value,
n,
n_controls: k,
method: format!("Partial correlation (controlling for {} variable(s))", k),
})
}
fn compute_partial_cor_single(x: &[f64], y: &[f64], z: &[f64]) -> Result<f64> {
let r_xy = pearson_r(x, y);
let r_xz = pearson_r(x, z);
let r_yz = pearson_r(y, z);
let numerator = r_xy - r_xz * r_yz;
let denom1 = (1.0 - r_xz * r_xz).sqrt();
let denom2 = (1.0 - r_yz * r_yz).sqrt();
if denom1 * denom2 > 1e-10 {
Ok(numerator / (denom1 * denom2))
} else {
Ok(0.0)
}
}
fn compute_simple_correlation(x: &[f64], y: &[f64], n: usize) -> Result<PartialCorResult> {
let r = pearson_r(x, y);
let df = (n - 2) as f64;
let t_stat = if (1.0 - r * r).abs() < 1e-15 {
if r > 0.0 {
f64::INFINITY
} else {
f64::NEG_INFINITY
}
} else {
r * (df / (1.0 - r * r)).sqrt()
};
let p_value = if t_stat.is_infinite() {
0.0
} else {
let t_dist = StudentsT::new(0.0, 1.0, df).unwrap();
2.0 * t_dist.sf(t_stat.abs())
};
Ok(PartialCorResult {
estimate: r,
statistic: t_stat,
df,
p_value,
n,
n_controls: 0,
method: "Pearson correlation (no controls)".to_string(),
})
}
fn pearson_r(x: &[f64], y: &[f64]) -> f64 {
let n = x.len();
let mean_x = mean(x);
let mean_y = mean(y);
let mut sum_xy = 0.0;
let mut sum_xx = 0.0;
let mut sum_yy = 0.0;
for i in 0..n {
let dx = x[i] - mean_x;
let dy = y[i] - mean_y;
sum_xy += dx * dy;
sum_xx += dx * dx;
sum_yy += dy * dy;
}
if sum_xx > 0.0 && sum_yy > 0.0 {
sum_xy / (sum_xx * sum_yy).sqrt()
} else {
0.0
}
}
pub fn semi_partial_cor(x: &[f64], y: &[f64], z: &[&[f64]]) -> Result<PartialCorResult> {
let n = validate_correlation_input(x, y)?;
for (i, zi) in z.iter().enumerate() {
if zi.len() != n {
return Err(StatError::InvalidParameter(format!(
"Control variable {} has length {}, expected {}",
i,
zi.len(),
n
)));
}
}
let k = z.len();
if k == 0 {
return compute_simple_correlation(x, y, n);
}
if n <= k + 2 {
return Err(StatError::InsufficientData {
needed: k + 3,
got: n,
});
}
let y_residuals = compute_residuals(y, z);
let r = pearson_r(x, &y_residuals);
let r = r.clamp(-1.0, 1.0);
let df = (n - k - 2) as f64;
let t_stat = if (1.0 - r * r).abs() < 1e-15 {
if r > 0.0 {
f64::INFINITY
} else {
f64::NEG_INFINITY
}
} else {
r * (df / (1.0 - r * r)).sqrt()
};
let p_value = if t_stat.is_infinite() {
0.0
} else if df > 0.0 {
let t_dist = StudentsT::new(0.0, 1.0, df).unwrap();
2.0 * t_dist.sf(t_stat.abs())
} else {
1.0
};
Ok(PartialCorResult {
estimate: r,
statistic: t_stat,
df,
p_value,
n,
n_controls: k,
method: format!(
"Semi-partial correlation (controlling for {} variable(s) on y)",
k
),
})
}
fn compute_residuals(y: &[f64], z: &[&[f64]]) -> Vec<f64> {
let n = y.len();
let k = z.len();
if k == 0 {
return y.to_vec();
}
if k == 1 {
let z0 = z[0];
let mean_y = mean(y);
let mean_z = mean(z0);
let mut sum_zy = 0.0;
let mut sum_zz = 0.0;
for i in 0..n {
let dz = z0[i] - mean_z;
sum_zy += dz * (y[i] - mean_y);
sum_zz += dz * dz;
}
let beta = if sum_zz > 0.0 { sum_zy / sum_zz } else { 0.0 };
let alpha = mean_y - beta * mean_z;
return y
.iter()
.zip(z0.iter())
.map(|(&yi, &zi)| yi - alpha - beta * zi)
.collect();
}
let mut residuals = y.to_vec();
for zi in z {
let mean_r: f64 = residuals.iter().sum::<f64>() / n as f64;
let mean_z = mean(zi);
let mut sum_rz = 0.0;
let mut sum_zz = 0.0;
for i in 0..n {
let dz = zi[i] - mean_z;
sum_rz += dz * (residuals[i] - mean_r);
sum_zz += dz * dz;
}
let beta = if sum_zz > 0.0 { sum_rz / sum_zz } else { 0.0 };
let alpha = mean_r - beta * mean_z;
for i in 0..n {
residuals[i] -= alpha + beta * zi[i];
}
}
residuals
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_partial_cor_no_controls() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
let y = vec![2.0, 4.0, 6.0, 8.0, 10.0, 12.0, 14.0, 16.0, 18.0, 20.0];
let result = partial_cor(&x, &y, &[]).unwrap();
assert!((result.estimate - 1.0).abs() < 1e-10);
assert_eq!(result.n_controls, 0);
}
#[test]
fn test_partial_cor_single_control() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
let y = vec![2.1, 3.9, 6.1, 7.9, 10.1, 11.9, 14.1, 15.9, 18.1, 19.9];
let z = vec![0.5, 1.5, 2.5, 3.5, 4.5, 5.5, 6.5, 7.5, 8.5, 9.5];
let result = partial_cor(&x, &y, &[&z]).unwrap();
assert!(result.estimate >= -1.0 && result.estimate <= 1.0);
assert_eq!(result.n_controls, 1);
}
#[test]
fn test_partial_cor_confounded() {
let z = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
let x: Vec<f64> = z.iter().map(|&zi| zi * 2.0 + 0.1).collect();
let y: Vec<f64> = z.iter().map(|&zi| zi * 3.0 - 0.1).collect();
let result_no_control = partial_cor(&x, &y, &[]).unwrap();
assert!(result_no_control.estimate > 0.9);
let result_controlled = partial_cor(&x, &y, &[&z]).unwrap();
assert!(result_controlled.estimate.abs() < result_no_control.estimate.abs());
}
#[test]
fn test_partial_cor_multiple_controls() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
let y = vec![2.0, 3.5, 5.0, 6.5, 8.0, 9.5, 11.0, 12.5, 14.0, 15.5];
let z1 = vec![0.5, 1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.5, 5.0];
let z2 = vec![1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.5, 5.0, 5.5];
let result = partial_cor(&x, &y, &[&z1, &z2]).unwrap();
assert!(result.estimate.abs() <= 1.0);
assert_eq!(result.n_controls, 2);
}
#[test]
fn test_semi_partial_cor() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
let y = vec![2.0, 4.0, 5.0, 7.0, 9.0, 11.0, 13.0, 15.0, 17.0, 19.0];
let z = vec![0.5, 1.5, 2.5, 3.5, 4.5, 5.5, 6.5, 7.5, 8.5, 9.5];
let result = semi_partial_cor(&x, &y, &[&z]).unwrap();
assert!(result.estimate.abs() <= 1.0);
assert_eq!(result.n_controls, 1);
}
}