use scirs2_core::ndarray::Array2;
use crate::error::{Error, Result};
use crate::stats::descriptive::pearson_correlation;
pub fn correlation_matrix_listwise(columns: &[Vec<Option<f64>>]) -> Result<Array2<f64>> {
let n = columns.len();
if n == 0 {
return Err(Error::InvalidValue(
"correlation matrix requires at least one column".to_string(),
));
}
let n_rows = columns[0].len();
for (i, col) in columns.iter().enumerate() {
if col.len() != n_rows {
return Err(Error::DimensionMismatch(format!(
"column {} has {} rows, expected {} (from column 0)",
i,
col.len(),
n_rows
)));
}
}
let complete_rows: Vec<usize> = (0..n_rows)
.filter(|&row| columns.iter().all(|col| col[row].is_some()))
.collect();
let complete: Vec<Vec<f64>> = columns
.iter()
.map(|col| complete_rows.iter().filter_map(|&row| col[row]).collect())
.collect();
let mut matrix = Array2::<f64>::zeros((n, n));
for i in 0..n {
matrix[[i, i]] = 1.0;
for j in (i + 1)..n {
let r = pearson_correlation(&complete[i], &complete[j]).unwrap_or(f64::NAN);
matrix[[i, j]] = r;
matrix[[j, i]] = r;
}
}
Ok(matrix)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn listwise_deletion_keeps_pairs_aligned() {
let a = vec![Some(1.0), None, Some(3.0), Some(4.0), Some(5.0)];
let b = vec![Some(2.0), Some(100.0), None, Some(8.0), Some(10.0)];
let matrix = correlation_matrix_listwise(&[a, b]).expect("operation should succeed");
assert_eq!(matrix.shape(), &[2, 2]);
assert!((matrix[[0, 0]] - 1.0).abs() < 1e-12);
assert!((matrix[[1, 1]] - 1.0).abs() < 1e-12);
assert!(
(matrix[[0, 1]] - 1.0).abs() < 1e-9,
"expected perfect correlation on the complete rows, got {}",
matrix[[0, 1]]
);
assert!(
(matrix[[1, 0]] - matrix[[0, 1]]).abs() < 1e-15,
"matrix must be symmetric"
);
}
#[test]
fn empty_columns_errs() {
assert!(correlation_matrix_listwise(&[]).is_err());
}
#[test]
fn mismatched_lengths_err() {
let a = vec![Some(1.0), Some(2.0)];
let b = vec![Some(1.0)];
assert!(correlation_matrix_listwise(&[a, b]).is_err());
}
}