use crate::error::{Result, StatError};
use statrs::distribution::{ContinuousCDF, FisherSnedecor};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ICCType {
ICC1,
#[default]
ICC2,
ICC3,
ICC1k,
ICC2k,
ICC3k,
}
impl std::fmt::Display for ICCType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ICCType::ICC1 => write!(f, "ICC(1)"),
ICCType::ICC2 => write!(f, "ICC(2,1)"),
ICCType::ICC3 => write!(f, "ICC(3,1)"),
ICCType::ICC1k => write!(f, "ICC(1,k)"),
ICCType::ICC2k => write!(f, "ICC(2,k)"),
ICCType::ICC3k => write!(f, "ICC(3,k)"),
}
}
}
#[derive(Debug, Clone)]
pub struct ICCResult {
pub icc: f64,
pub icc_type: ICCType,
pub f_value: f64,
pub df1: f64,
pub df2: f64,
pub p_value: f64,
pub conf_int_lower: f64,
pub conf_int_upper: f64,
pub n_subjects: usize,
pub n_raters: usize,
pub method: String,
}
pub fn icc(data: &[Vec<f64>], icc_type: ICCType) -> Result<ICCResult> {
validate_icc_input(data)?;
let n = data.len(); let k = data[0].len();
let (ms_r, ms_c, ms_e, ms_w) = compute_anova_components(data);
let (icc_value, f_value, df1, df2) = match icc_type {
ICCType::ICC1 => {
let icc = (ms_r - ms_w) / (ms_r + (k as f64 - 1.0) * ms_w);
let f = ms_r / ms_w;
(icc, f, (n - 1) as f64, n as f64 * (k as f64 - 1.0))
}
ICCType::ICC2 => {
let numer = ms_r - ms_e;
let denom = ms_r + (k as f64 - 1.0) * ms_e + k as f64 * (ms_c - ms_e) / n as f64;
let icc = if denom > 0.0 { numer / denom } else { 0.0 };
let f = ms_r / ms_e;
(icc, f, (n - 1) as f64, ((n - 1) * (k - 1)) as f64)
}
ICCType::ICC3 => {
let icc = (ms_r - ms_e) / (ms_r + (k as f64 - 1.0) * ms_e);
let f = ms_r / ms_e;
(icc, f, (n - 1) as f64, ((n - 1) * (k - 1)) as f64)
}
ICCType::ICC1k => {
let icc = (ms_r - ms_w) / ms_r;
let f = ms_r / ms_w;
(icc, f, (n - 1) as f64, n as f64 * (k as f64 - 1.0))
}
ICCType::ICC2k => {
let numer = ms_r - ms_e;
let denom = ms_r + (ms_c - ms_e) / n as f64;
let icc = if denom > 0.0 { numer / denom } else { 0.0 };
let f = ms_r / ms_e;
(icc, f, (n - 1) as f64, ((n - 1) * (k - 1)) as f64)
}
ICCType::ICC3k => {
let icc = (ms_r - ms_e) / ms_r;
let f = ms_r / ms_e;
(icc, f, (n - 1) as f64, ((n - 1) * (k - 1)) as f64)
}
};
let icc_value = icc_value.clamp(-1.0, 1.0);
let p_value = if f_value > 0.0 && df1 > 0.0 && df2 > 0.0 {
let f_dist = FisherSnedecor::new(df1, df2).unwrap();
f_dist.sf(f_value)
} else {
1.0
};
let (conf_int_lower, conf_int_upper) =
compute_icc_ci(icc_value, f_value, df1, df2, n, k, icc_type);
Ok(ICCResult {
icc: icc_value,
icc_type,
f_value,
df1,
df2,
p_value,
conf_int_lower,
conf_int_upper,
n_subjects: n,
n_raters: k,
method: format!("{} - Intraclass Correlation Coefficient", icc_type),
})
}
fn validate_icc_input(data: &[Vec<f64>]) -> Result<()> {
if data.is_empty() {
return Err(StatError::EmptyData);
}
let n = data.len();
if n < 2 {
return Err(StatError::InsufficientData { needed: 2, got: n });
}
let k = data[0].len();
if k < 2 {
return Err(StatError::InsufficientData { needed: 2, got: k });
}
for (i, row) in data.iter().enumerate() {
if row.len() != k {
return Err(StatError::InvalidParameter(format!(
"Row {} has {} columns, expected {}",
i,
row.len(),
k
)));
}
for (j, &val) in row.iter().enumerate() {
if !val.is_finite() {
return Err(StatError::InvalidParameter(format!(
"Non-finite value at row {}, column {}",
i, j
)));
}
}
}
Ok(())
}
fn compute_anova_components(data: &[Vec<f64>]) -> (f64, f64, f64, f64) {
let n = data.len(); let k = data[0].len(); let n_f = n as f64;
let k_f = k as f64;
let total_n = (n * k) as f64;
let grand_mean: f64 = data.iter().flat_map(|row| row.iter()).sum::<f64>() / total_n;
let row_means: Vec<f64> = data
.iter()
.map(|row| row.iter().sum::<f64>() / k_f)
.collect();
let col_means: Vec<f64> = (0..k)
.map(|j| data.iter().map(|row| row[j]).sum::<f64>() / n_f)
.collect();
let ss_total: f64 = data
.iter()
.flat_map(|row| row.iter())
.map(|&x| (x - grand_mean).powi(2))
.sum();
let ss_r: f64 = k_f
* row_means
.iter()
.map(|&m| (m - grand_mean).powi(2))
.sum::<f64>();
let ss_c: f64 = n_f
* col_means
.iter()
.map(|&m| (m - grand_mean).powi(2))
.sum::<f64>();
let ss_e = (ss_total - ss_r - ss_c).max(0.0);
let ss_w = (ss_total - ss_r).max(0.0);
let df_r = n_f - 1.0;
let df_c = k_f - 1.0;
let df_e = (n_f - 1.0) * (k_f - 1.0);
let df_w = n_f * (k_f - 1.0);
let ms_r = if df_r > 0.0 { ss_r / df_r } else { 0.0 };
let ms_c = if df_c > 0.0 { ss_c / df_c } else { 0.0 };
let ms_e = if df_e > 0.0 { ss_e / df_e } else { 0.0 };
let ms_w = if df_w > 0.0 { ss_w / df_w } else { 0.0 };
(ms_r, ms_c, ms_e, ms_w)
}
fn compute_icc_ci(
_icc: f64,
f_value: f64,
df1: f64,
df2: f64,
n: usize,
k: usize,
icc_type: ICCType,
) -> (f64, f64) {
if df1 <= 0.0 || df2 <= 0.0 || !f_value.is_finite() {
return (f64::NEG_INFINITY, f64::INFINITY);
}
let alpha = 0.05;
let f_dist = FisherSnedecor::new(df1, df2).unwrap();
let f_lower = f_dist.inverse_cdf(alpha / 2.0);
let f_upper = f_dist.inverse_cdf(1.0 - alpha / 2.0);
let _n = n; let k_f = k as f64;
let (lower, upper) = match icc_type {
ICCType::ICC1 | ICCType::ICC2 | ICCType::ICC3 => {
let f_l = f_value / f_upper;
let f_u = f_value / f_lower;
let lower = (f_l - 1.0) / (f_l + k_f - 1.0);
let upper = (f_u - 1.0) / (f_u + k_f - 1.0);
(lower.max(-1.0), upper.min(1.0))
}
ICCType::ICC1k | ICCType::ICC2k | ICCType::ICC3k => {
let f_l = f_value / f_upper;
let f_u = f_value / f_lower;
let lower = 1.0 - 1.0 / f_l;
let upper = 1.0 - 1.0 / f_u;
(lower.max(-1.0), upper.min(1.0))
}
};
(lower, upper)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_icc_perfect_agreement() {
let data = vec![
vec![1.0, 1.0, 1.0],
vec![2.0, 2.0, 2.0],
vec![3.0, 3.0, 3.0],
vec![4.0, 4.0, 4.0],
vec![5.0, 5.0, 5.0],
];
let result = icc(&data, ICCType::ICC2).unwrap();
assert!((result.icc - 1.0).abs() < 1e-10);
}
#[test]
fn test_icc_no_agreement() {
let data = vec![
vec![1.0, 5.0, 3.0],
vec![5.0, 1.0, 2.0],
vec![3.0, 4.0, 1.0],
vec![2.0, 3.0, 5.0],
vec![4.0, 2.0, 4.0],
];
let result = icc(&data, ICCType::ICC2).unwrap();
assert!(result.icc < 0.5);
}
#[test]
fn test_icc_types() {
let data = vec![
vec![9.0, 2.0, 5.0, 8.0],
vec![6.0, 1.0, 3.0, 2.0],
vec![8.0, 4.0, 6.0, 8.0],
vec![7.0, 1.0, 2.0, 6.0],
vec![10.0, 5.0, 6.0, 9.0],
vec![6.0, 2.0, 4.0, 7.0],
];
let icc1 = icc(&data, ICCType::ICC1).unwrap();
let icc2 = icc(&data, ICCType::ICC2).unwrap();
let icc3 = icc(&data, ICCType::ICC3).unwrap();
assert!(icc1.icc >= -1.0 && icc1.icc <= 1.0);
assert!(icc2.icc >= -1.0 && icc2.icc <= 1.0);
assert!(icc3.icc >= -1.0 && icc3.icc <= 1.0);
assert!(icc3.icc >= icc2.icc - 0.01);
}
#[test]
fn test_icc_average_raters() {
let data = vec![
vec![9.0, 2.0, 5.0, 8.0],
vec![6.0, 1.0, 3.0, 2.0],
vec![8.0, 4.0, 6.0, 8.0],
vec![7.0, 1.0, 2.0, 6.0],
vec![10.0, 5.0, 6.0, 9.0],
vec![6.0, 2.0, 4.0, 7.0],
];
let icc2_single = icc(&data, ICCType::ICC2).unwrap();
let icc2_avg = icc(&data, ICCType::ICC2k).unwrap();
assert!(icc2_avg.icc >= icc2_single.icc);
}
#[test]
fn test_icc_ci() {
let data = vec![
vec![9.0, 2.0, 5.0, 8.0],
vec![6.0, 1.0, 3.0, 2.0],
vec![8.0, 4.0, 6.0, 8.0],
vec![7.0, 1.0, 2.0, 6.0],
vec![10.0, 5.0, 6.0, 9.0],
vec![6.0, 2.0, 4.0, 7.0],
];
let result = icc(&data, ICCType::ICC2).unwrap();
assert!(result.conf_int_lower >= -1.0);
assert!(result.conf_int_upper <= 1.0);
assert!(result.conf_int_lower < result.conf_int_upper);
}
}