Skip to main content

fin_primitives/correlation/
mod.rs

1//! Streaming Pearson correlation matrix for indicator redundancy detection.
2//!
3//! ## Responsibility
4//! Streaming Pearson correlation matrix across a configurable window of indicator outputs.
5//! Identifies redundant signals (pairs with |r| > a configurable threshold, default 0.95).
6//!
7//! Also exposes [`stats`] for standalone Pearson, Spearman, Kendall tau-b functions,
8//! a symbol-based [`stats::SymbolCorrelationMatrix`], and a [`stats::RollingCorrelation`]
9//! rolling-window tracker.
10//!
11//! ## Guarantees
12//! - Returns `None` from [`CorrelationMatrix::get`] until `window` samples have been seen
13//! - Correlation values are clamped to `[-1, 1]` to absorb floating-point rounding errors
14//! - `most_correlated_with` results are sorted descending by absolute correlation
15//!
16//! ## NOT Responsible For
17//! - Causal inference or feature selection policy
18//! - Persistence
19
20/// Standalone correlation measures: Pearson, Spearman rank, Kendall tau-b,
21/// symbol-keyed `SymbolCorrelationMatrix`, and `RollingCorrelation`.
22pub mod stats;
23
24use crate::error::FinError;
25use std::collections::VecDeque;
26
27/// Default redundancy threshold: pairs with |r| above this are flagged as redundant.
28pub const DEFAULT_REDUNDANCY_THRESHOLD: f64 = 0.95;
29
30/// A streaming Pearson correlation matrix for a fixed set of indicators.
31///
32/// Feed one sample per bar via [`CorrelationMatrix::update`]. Once `window` samples
33/// have been accumulated the full `n × n` correlation matrix is available.
34///
35/// # Example
36/// ```rust
37/// use fin_primitives::correlation::CorrelationMatrix;
38///
39/// let mut cm = CorrelationMatrix::new(3, 5, 0.95).unwrap();
40/// for i in 0..5 {
41///     let vals = vec![i as f64, (i * 2) as f64, (10 - i) as f64];
42///     cm.update(&vals).unwrap();
43/// }
44/// // indicators 0 and 1 are perfectly correlated
45/// let r = cm.get(0, 1).unwrap();
46/// assert!((r - 1.0).abs() < 1e-9);
47/// ```
48#[derive(Debug)]
49pub struct CorrelationMatrix {
50    /// Number of indicators tracked.
51    n: usize,
52    /// Rolling window size.
53    window: usize,
54    /// Redundancy threshold.
55    threshold: f64,
56    /// Circular buffer: each entry is one bar's vector of `n` values.
57    buf: VecDeque<Vec<f64>>,
58}
59
60impl CorrelationMatrix {
61    /// Constructs a new `CorrelationMatrix`.
62    ///
63    /// # Parameters
64    /// - `n_indicators`: number of indicators (must be >= 2)
65    /// - `window`: rolling window in bars (must be >= 2)
66    /// - `redundancy_threshold`: absolute correlation above which a pair is flagged
67    ///
68    /// # Errors
69    /// Returns [`FinError::InvalidPeriod`] if `window < 2`.
70    /// Returns [`FinError::InvalidInput`] if `n_indicators < 2` or threshold is not in `(0, 1]`.
71    pub fn new(n_indicators: usize, window: usize, redundancy_threshold: f64) -> Result<Self, FinError> {
72        if window < 2 {
73            return Err(FinError::InvalidPeriod(window));
74        }
75        if n_indicators < 2 {
76            return Err(FinError::InvalidInput(
77                "CorrelationMatrix requires at least 2 indicators".to_owned(),
78            ));
79        }
80        if redundancy_threshold <= 0.0 || redundancy_threshold > 1.0 {
81            return Err(FinError::InvalidInput(
82                "redundancy_threshold must be in (0, 1]".to_owned(),
83            ));
84        }
85        Ok(Self {
86            n: n_indicators,
87            window,
88            threshold: redundancy_threshold,
89            buf: VecDeque::with_capacity(window),
90        })
91    }
92
93    /// Constructs a `CorrelationMatrix` with the default redundancy threshold (0.95).
94    ///
95    /// # Errors
96    /// See [`CorrelationMatrix::new`].
97    pub fn with_defaults(n_indicators: usize, window: usize) -> Result<Self, FinError> {
98        Self::new(n_indicators, window, DEFAULT_REDUNDANCY_THRESHOLD)
99    }
100
101    /// Feeds one bar's worth of indicator values.
102    ///
103    /// `values.len()` must equal the `n_indicators` supplied at construction.
104    ///
105    /// # Errors
106    /// Returns [`FinError::InvalidInput`] if `values.len() != n_indicators`.
107    pub fn update(&mut self, values: &[f64]) -> Result<(), FinError> {
108        if values.len() != self.n {
109            return Err(FinError::InvalidInput(format!(
110                "expected {} values, got {}",
111                self.n,
112                values.len()
113            )));
114        }
115        self.buf.push_back(values.to_vec());
116        if self.buf.len() > self.window {
117            self.buf.pop_front();
118        }
119        Ok(())
120    }
121
122    /// Returns `true` when enough samples have been accumulated to compute correlations.
123    pub fn is_ready(&self) -> bool {
124        self.buf.len() >= self.window
125    }
126
127    /// Returns the Pearson correlation between indicators `i` and `j`.
128    ///
129    /// Returns `None` when fewer than `window` samples have been seen, or when either
130    /// indicator has zero variance (correlation is undefined).
131    pub fn get(&self, i: usize, j: usize) -> Option<f64> {
132        if !self.is_ready() {
133            return None;
134        }
135        if i == j {
136            return Some(1.0);
137        }
138        let n = self.buf.len() as f64;
139        let mut sum_x = 0.0_f64;
140        let mut sum_y = 0.0_f64;
141        let mut sum_xy = 0.0_f64;
142        let mut sum_x2 = 0.0_f64;
143        let mut sum_y2 = 0.0_f64;
144        for row in &self.buf {
145            let x = row[i];
146            let y = row[j];
147            sum_x += x;
148            sum_y += y;
149            sum_xy += x * y;
150            sum_x2 += x * x;
151            sum_y2 += y * y;
152        }
153        let num = n * sum_xy - sum_x * sum_y;
154        let den_sq = (n * sum_x2 - sum_x * sum_x) * (n * sum_y2 - sum_y * sum_y);
155        if den_sq <= 0.0 {
156            return None;
157        }
158        let r = num / den_sq.sqrt();
159        // Clamp to [-1, 1] to absorb floating-point rounding errors
160        Some(r.clamp(-1.0, 1.0))
161    }
162
163    /// Returns the full `n × n` correlation matrix as a flat `Vec<f64>` (row-major).
164    ///
165    /// Element at row `i`, column `j` is at index `i * n + j`.
166    /// Returns `None` until ready.
167    pub fn matrix(&self) -> Option<Vec<f64>> {
168        if !self.is_ready() {
169            return None;
170        }
171        let mut mat = vec![0.0_f64; self.n * self.n];
172        for i in 0..self.n {
173            for j in 0..self.n {
174                mat[i * self.n + j] = self.get(i, j).unwrap_or(0.0);
175            }
176        }
177        Some(mat)
178    }
179
180    /// Returns all indicators whose absolute correlation with `indicator_id` exceeds
181    /// `threshold`, sorted descending by absolute correlation value.
182    ///
183    /// Returns an empty `Vec` if the matrix is not yet ready.
184    pub fn most_correlated_with(&self, indicator_id: usize) -> Vec<(usize, f64)> {
185        if !self.is_ready() {
186            return vec![];
187        }
188        let mut result: Vec<(usize, f64)> = (0..self.n)
189            .filter(|&j| j != indicator_id)
190            .filter_map(|j| {
191                self.get(indicator_id, j)
192                    .map(|r| (j, r))
193            })
194            .collect();
195        result.sort_by(|a, b| b.1.abs().partial_cmp(&a.1.abs()).unwrap_or(std::cmp::Ordering::Equal));
196        result
197    }
198
199    /// Returns all pairs `(i, j)` where `i < j` and `|r| >= threshold`.
200    ///
201    /// These pairs are considered redundant signals.
202    /// Returns an empty `Vec` if the matrix is not ready.
203    pub fn redundant_pairs(&self) -> Vec<(usize, usize, f64)> {
204        if !self.is_ready() {
205            return vec![];
206        }
207        let mut pairs = Vec::new();
208        for i in 0..self.n {
209            for j in (i + 1)..self.n {
210                if let Some(r) = self.get(i, j) {
211                    if r.abs() >= self.threshold {
212                        pairs.push((i, j, r));
213                    }
214                }
215            }
216        }
217        pairs
218    }
219
220    /// Returns the number of indicators tracked.
221    pub fn n_indicators(&self) -> usize {
222        self.n
223    }
224
225    /// Returns the configured window size.
226    pub fn window(&self) -> usize {
227        self.window
228    }
229
230    /// Returns the number of samples currently buffered.
231    pub fn sample_count(&self) -> usize {
232        self.buf.len()
233    }
234}
235
236#[cfg(test)]
237mod tests {
238    use super::*;
239
240    fn feed(cm: &mut CorrelationMatrix, rows: &[[f64; 3]]) {
241        for row in rows {
242            cm.update(row).unwrap();
243        }
244    }
245
246    #[test]
247    fn test_perfect_positive_correlation() {
248        let mut cm = CorrelationMatrix::new(3, 5, 0.95).unwrap();
249        // indicators 0 and 1 are y=2x (perfect positive correlation)
250        // indicator 2 is negatively correlated with 0
251        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]];
252        feed(&mut cm, &data);
253        assert!(cm.is_ready());
254        let r01 = cm.get(0, 1).unwrap();
255        assert!((r01 - 1.0).abs() < 1e-9, "r01={r01}");
256        let r02 = cm.get(0, 2).unwrap();
257        assert!((r02 + 1.0).abs() < 1e-9, "r02={r02}");
258    }
259
260    #[test]
261    fn test_not_ready_until_window_filled() {
262        let mut cm = CorrelationMatrix::new(2, 5, 0.95).unwrap();
263        for i in 0..4 {
264            cm.update(&[i as f64, (i * 2) as f64]).unwrap();
265        }
266        assert!(!cm.is_ready());
267        assert!(cm.get(0, 1).is_none());
268    }
269
270    #[test]
271    fn test_most_correlated_with_sorted() {
272        let mut cm = CorrelationMatrix::new(3, 5, 0.50).unwrap();
273        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]];
274        feed(&mut cm, &data);
275        let corrs = cm.most_correlated_with(0);
276        assert_eq!(corrs.len(), 2);
277        // highest abs correlation first
278        assert!(corrs[0].1.abs() >= corrs[1].1.abs());
279    }
280
281    #[test]
282    fn test_redundant_pairs() {
283        let mut cm = CorrelationMatrix::new(3, 5, 0.95).unwrap();
284        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]];
285        feed(&mut cm, &data);
286        let pairs = cm.redundant_pairs();
287        // pairs (0,1) r≈1.0, (0,2) r≈-1.0, (1,2) r≈-1.0 all exceed 0.95
288        assert_eq!(pairs.len(), 3);
289    }
290
291    #[test]
292    fn test_self_correlation_is_one() {
293        let mut cm = CorrelationMatrix::new(2, 3, 0.95).unwrap();
294        for i in 0..3 {
295            cm.update(&[i as f64, (i * 3) as f64]).unwrap();
296        }
297        assert_eq!(cm.get(0, 0).unwrap(), 1.0);
298        assert_eq!(cm.get(1, 1).unwrap(), 1.0);
299    }
300
301    #[test]
302    fn test_zero_variance_returns_none() {
303        let mut cm = CorrelationMatrix::new(2, 3, 0.95).unwrap();
304        // indicator 1 is constant → zero variance
305        for _ in 0..3 {
306            cm.update(&[1.0, 5.0]).unwrap();
307        }
308        assert!(cm.get(0, 1).is_none());
309    }
310
311    #[test]
312    fn test_matrix_shape() {
313        let mut cm = CorrelationMatrix::new(3, 3, 0.95).unwrap();
314        for i in 0..3 {
315            cm.update(&[i as f64, (i + 1) as f64, (i * 2) as f64]).unwrap();
316        }
317        let mat = cm.matrix().unwrap();
318        assert_eq!(mat.len(), 9);
319        // diagonal should be 1
320        assert_eq!(mat[0], 1.0);
321        assert_eq!(mat[4], 1.0);
322        assert_eq!(mat[8], 1.0);
323    }
324
325    #[test]
326    fn test_invalid_period_error() {
327        assert!(matches!(
328            CorrelationMatrix::new(2, 1, 0.95).unwrap_err(),
329            FinError::InvalidPeriod(_)
330        ));
331    }
332
333    #[test]
334    fn test_invalid_indicator_count_error() {
335        assert!(matches!(
336            CorrelationMatrix::new(1, 5, 0.95).unwrap_err(),
337            FinError::InvalidInput(_)
338        ));
339    }
340
341    #[test]
342    fn test_window_rolls_old_samples() {
343        let mut cm = CorrelationMatrix::new(2, 3, 0.95).unwrap();
344        // Feed 5 samples; only last 3 matter
345        for i in 0..5 {
346            cm.update(&[i as f64, (i * 2) as f64]).unwrap();
347        }
348        assert_eq!(cm.sample_count(), 3);
349    }
350}