pub mod stats;
use crate::error::FinError;
use std::collections::VecDeque;
pub const DEFAULT_REDUNDANCY_THRESHOLD: f64 = 0.95;
#[derive(Debug)]
pub struct CorrelationMatrix {
n: usize,
window: usize,
threshold: f64,
buf: VecDeque<Vec<f64>>,
}
impl CorrelationMatrix {
pub fn new(n_indicators: usize, window: usize, redundancy_threshold: f64) -> Result<Self, FinError> {
if window < 2 {
return Err(FinError::InvalidPeriod(window));
}
if n_indicators < 2 {
return Err(FinError::InvalidInput(
"CorrelationMatrix requires at least 2 indicators".to_owned(),
));
}
if redundancy_threshold <= 0.0 || redundancy_threshold > 1.0 {
return Err(FinError::InvalidInput(
"redundancy_threshold must be in (0, 1]".to_owned(),
));
}
Ok(Self {
n: n_indicators,
window,
threshold: redundancy_threshold,
buf: VecDeque::with_capacity(window),
})
}
pub fn with_defaults(n_indicators: usize, window: usize) -> Result<Self, FinError> {
Self::new(n_indicators, window, DEFAULT_REDUNDANCY_THRESHOLD)
}
pub fn update(&mut self, values: &[f64]) -> Result<(), FinError> {
if values.len() != self.n {
return Err(FinError::InvalidInput(format!(
"expected {} values, got {}",
self.n,
values.len()
)));
}
self.buf.push_back(values.to_vec());
if self.buf.len() > self.window {
self.buf.pop_front();
}
Ok(())
}
pub fn is_ready(&self) -> bool {
self.buf.len() >= self.window
}
pub fn get(&self, i: usize, j: usize) -> Option<f64> {
if !self.is_ready() {
return None;
}
if i == j {
return Some(1.0);
}
let n = self.buf.len() as f64;
let mut sum_x = 0.0_f64;
let mut sum_y = 0.0_f64;
let mut sum_xy = 0.0_f64;
let mut sum_x2 = 0.0_f64;
let mut sum_y2 = 0.0_f64;
for row in &self.buf {
let x = row[i];
let y = row[j];
sum_x += x;
sum_y += y;
sum_xy += x * y;
sum_x2 += x * x;
sum_y2 += y * y;
}
let num = n * sum_xy - sum_x * sum_y;
let den_sq = (n * sum_x2 - sum_x * sum_x) * (n * sum_y2 - sum_y * sum_y);
if den_sq <= 0.0 {
return None;
}
let r = num / den_sq.sqrt();
Some(r.clamp(-1.0, 1.0))
}
pub fn matrix(&self) -> Option<Vec<f64>> {
if !self.is_ready() {
return None;
}
let mut mat = vec![0.0_f64; self.n * self.n];
for i in 0..self.n {
for j in 0..self.n {
mat[i * self.n + j] = self.get(i, j).unwrap_or(0.0);
}
}
Some(mat)
}
pub fn most_correlated_with(&self, indicator_id: usize) -> Vec<(usize, f64)> {
if !self.is_ready() {
return vec![];
}
let mut result: Vec<(usize, f64)> = (0..self.n)
.filter(|&j| j != indicator_id)
.filter_map(|j| {
self.get(indicator_id, j)
.map(|r| (j, r))
})
.collect();
result.sort_by(|a, b| b.1.abs().partial_cmp(&a.1.abs()).unwrap_or(std::cmp::Ordering::Equal));
result
}
pub fn redundant_pairs(&self) -> Vec<(usize, usize, f64)> {
if !self.is_ready() {
return vec![];
}
let mut pairs = Vec::new();
for i in 0..self.n {
for j in (i + 1)..self.n {
if let Some(r) = self.get(i, j) {
if r.abs() >= self.threshold {
pairs.push((i, j, r));
}
}
}
}
pairs
}
pub fn n_indicators(&self) -> usize {
self.n
}
pub fn window(&self) -> usize {
self.window
}
pub fn sample_count(&self) -> usize {
self.buf.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn feed(cm: &mut CorrelationMatrix, rows: &[[f64; 3]]) {
for row in rows {
cm.update(row).unwrap();
}
}
#[test]
fn test_perfect_positive_correlation() {
let mut cm = CorrelationMatrix::new(3, 5, 0.95).unwrap();
let data = [[1.0, 2.0, 10.0], [2.0, 4.0, 9.0], [3.0, 6.0, 8.0], [4.0, 8.0, 7.0], [5.0, 10.0, 6.0]];
feed(&mut cm, &data);
assert!(cm.is_ready());
let r01 = cm.get(0, 1).unwrap();
assert!((r01 - 1.0).abs() < 1e-9, "r01={r01}");
let r02 = cm.get(0, 2).unwrap();
assert!((r02 + 1.0).abs() < 1e-9, "r02={r02}");
}
#[test]
fn test_not_ready_until_window_filled() {
let mut cm = CorrelationMatrix::new(2, 5, 0.95).unwrap();
for i in 0..4 {
cm.update(&[i as f64, (i * 2) as f64]).unwrap();
}
assert!(!cm.is_ready());
assert!(cm.get(0, 1).is_none());
}
#[test]
fn test_most_correlated_with_sorted() {
let mut cm = CorrelationMatrix::new(3, 5, 0.50).unwrap();
let data = [[1.0, 2.0, 10.0], [2.0, 4.0, 9.0], [3.0, 6.0, 8.0], [4.0, 8.0, 7.0], [5.0, 10.0, 6.0]];
feed(&mut cm, &data);
let corrs = cm.most_correlated_with(0);
assert_eq!(corrs.len(), 2);
assert!(corrs[0].1.abs() >= corrs[1].1.abs());
}
#[test]
fn test_redundant_pairs() {
let mut cm = CorrelationMatrix::new(3, 5, 0.95).unwrap();
let data = [[1.0, 2.0, 10.0], [2.0, 4.0, 9.0], [3.0, 6.0, 8.0], [4.0, 8.0, 7.0], [5.0, 10.0, 6.0]];
feed(&mut cm, &data);
let pairs = cm.redundant_pairs();
assert_eq!(pairs.len(), 3);
}
#[test]
fn test_self_correlation_is_one() {
let mut cm = CorrelationMatrix::new(2, 3, 0.95).unwrap();
for i in 0..3 {
cm.update(&[i as f64, (i * 3) as f64]).unwrap();
}
assert_eq!(cm.get(0, 0).unwrap(), 1.0);
assert_eq!(cm.get(1, 1).unwrap(), 1.0);
}
#[test]
fn test_zero_variance_returns_none() {
let mut cm = CorrelationMatrix::new(2, 3, 0.95).unwrap();
for _ in 0..3 {
cm.update(&[1.0, 5.0]).unwrap();
}
assert!(cm.get(0, 1).is_none());
}
#[test]
fn test_matrix_shape() {
let mut cm = CorrelationMatrix::new(3, 3, 0.95).unwrap();
for i in 0..3 {
cm.update(&[i as f64, (i + 1) as f64, (i * 2) as f64]).unwrap();
}
let mat = cm.matrix().unwrap();
assert_eq!(mat.len(), 9);
assert_eq!(mat[0], 1.0);
assert_eq!(mat[4], 1.0);
assert_eq!(mat[8], 1.0);
}
#[test]
fn test_invalid_period_error() {
assert!(matches!(
CorrelationMatrix::new(2, 1, 0.95).unwrap_err(),
FinError::InvalidPeriod(_)
));
}
#[test]
fn test_invalid_indicator_count_error() {
assert!(matches!(
CorrelationMatrix::new(1, 5, 0.95).unwrap_err(),
FinError::InvalidInput(_)
));
}
#[test]
fn test_window_rolls_old_samples() {
let mut cm = CorrelationMatrix::new(2, 3, 0.95).unwrap();
for i in 0..5 {
cm.update(&[i as f64, (i * 2) as f64]).unwrap();
}
assert_eq!(cm.sample_count(), 3);
}
}