use super::decomp::eigen::symmetric_eigen;
use crate::core::errors::RustyQLibError;
pub fn nearest_correlation(a: &[Vec<f64>], tol: f64, max_iter: usize) -> Result<Vec<Vec<f64>>, RustyQLibError> {
let n = a.len();
if a.iter().any(|row| row.len() != n) {
return Err(RustyQLibError::NumericalError("matrix must be square".to_string()));
}
for i in 0..n {
for j in 0..n {
if (a[i][j] - a[j][i]).abs() > 1e-8 {
return Err(RustyQLibError::NumericalError("matrix must be symmetric".to_string()));
}
}
}
let mut y = a.to_vec();
for i in 0..n {
for j in 0..i {
let s = 0.5 * (y[i][j] + y[j][i]);
y[i][j] = s;
y[j][i] = s;
}
}
let mut dykstra = vec![vec![0.0; n]; n];
for _ in 0..max_iter {
let mut r = y.clone();
for i in 0..n {
for j in 0..n {
r[i][j] -= dykstra[i][j];
}
}
let x = psd_projection(&r);
for i in 0..n {
for j in 0..n {
dykstra[i][j] = x[i][j] - r[i][j];
}
}
let mut y_next = x.clone();
for (i, row) in y_next.iter_mut().enumerate() {
row[i] = 1.0;
}
let delta = max_abs_diff(&y_next, &y);
y = y_next;
if delta <= tol {
break;
}
}
let mut out = psd_projection(&y);
for i in 0..n {
out[i][i] = 1.0;
for j in 0..n {
if i != j {
out[i][j] = out[i][j].clamp(-1.0, 1.0);
}
}
}
Ok(out)
}
fn max_abs_diff(a: &[Vec<f64>], b: &[Vec<f64>]) -> f64 {
a.iter()
.zip(b)
.flat_map(|(ra, rb)| ra.iter().zip(rb).map(|(x, y)| (x - y).abs()))
.fold(0.0, f64::max)
}
fn psd_projection(a: &[Vec<f64>]) -> Vec<Vec<f64>> {
let n = a.len();
let (mut vals, vecs) = symmetric_eigen(a);
for v in vals.iter_mut() {
*v = v.max(0.0);
}
let mut out = vec![vec![0.0; n]; n];
for i in 0..n {
for j in 0..=i {
let mut s = 0.0;
for (k, &val) in vals.iter().enumerate() {
s += vecs[i][k] * val * vecs[j][k];
}
out[i][j] = s;
out[j][i] = s;
}
}
out
}
#[cfg(test)]
mod tests {
use super::super::cholesky::cholesky;
use super::*;
fn eigenvalues(a: &[Vec<f64>]) -> Vec<f64> {
symmetric_eigen(a).0
}
#[test]
fn higham_2002_example_matches_the_published_answer() {
let a = vec![vec![1.0, 1.0, 0.0], vec![1.0, 1.0, 1.0], vec![0.0, 1.0, 1.0]];
let x = nearest_correlation(&a, 1e-12, 200).unwrap();
assert!((x[0][1] - 0.7607).abs() < 1e-3, "x01 {}", x[0][1]);
assert!((x[1][2] - 0.7607).abs() < 1e-3, "x12 {}", x[1][2]);
assert!((x[0][2] - 0.1573).abs() < 1e-3, "x02 {}", x[0][2]);
}
#[test]
fn output_is_a_valid_correlation_matrix() {
let a = vec![
vec![1.0, 0.9, -0.9],
vec![0.9, 1.0, 0.9],
vec![-0.9, 0.9, 1.0],
];
assert!(eigenvalues(&a).iter().any(|&v| v < -1e-6), "test input should be indefinite");
let x = nearest_correlation(&a, 1e-12, 200).unwrap();
for i in 0..3 {
assert_eq!(x[i][i], 1.0);
for j in 0..3 {
assert!((x[i][j] - x[j][i]).abs() < 1e-12);
assert!(x[i][j].abs() <= 1.0 + 1e-12);
}
}
assert!(eigenvalues(&x).iter().all(|&v| v >= -1e-10), "not PSD: {:?}", eigenvalues(&x));
assert!(cholesky(&x).is_ok());
}
#[test]
fn valid_matrices_pass_through_unchanged() {
let a = vec![vec![1.0, 0.3, 0.1], vec![0.3, 1.0, 0.2], vec![0.1, 0.2, 1.0]];
let x = nearest_correlation(&a, 1e-12, 200).unwrap();
for i in 0..3 {
for j in 0..3 {
assert!((x[i][j] - a[i][j]).abs() < 1e-8, "[{i}][{j}]");
}
}
}
}