Skip to main content

fin_primitives/correlation/
stats.rs

1//! # Module: correlation::stats
2//!
3//! Rolling and pairwise correlation measures:
4//! - Pearson correlation
5//! - Spearman rank correlation (with average-rank tie handling)
6//! - Kendall tau-b (O(n²) concordant/discordant counting)
7//! - `SymbolCorrelationMatrix`: full Pearson matrix for N symbols
8//! - `RollingCorrelation`: rolling-window pairwise correlations via `VecDeque`
9
10use std::collections::HashMap;
11use std::collections::VecDeque;
12
13// ─── Pearson ─────────────────────────────────────────────────────────────────
14
15/// Compute the Pearson product-moment correlation between two equal-length slices.
16///
17/// Returns `None` when:
18/// - Either slice has fewer than 2 elements.
19/// - Either series has zero (or near-zero) variance.
20/// - The slices have different lengths.
21pub fn pearson_correlation(x: &[f64], y: &[f64]) -> Option<f64> {
22    if x.len() != y.len() || x.len() < 2 {
23        return None;
24    }
25    let n = x.len() as f64;
26    let sum_x: f64 = x.iter().sum();
27    let sum_y: f64 = y.iter().sum();
28    let sum_xy: f64 = x.iter().zip(y.iter()).map(|(a, b)| a * b).sum();
29    let sum_x2: f64 = x.iter().map(|a| a * a).sum();
30    let sum_y2: f64 = y.iter().map(|b| b * b).sum();
31
32    let num = n * sum_xy - sum_x * sum_y;
33    let den_sq = (n * sum_x2 - sum_x * sum_x) * (n * sum_y2 - sum_y * sum_y);
34    if den_sq <= 0.0 {
35        return None;
36    }
37    Some((num / den_sq.sqrt()).clamp(-1.0, 1.0))
38}
39
40// ─── Spearman ────────────────────────────────────────────────────────────────
41
42/// Compute Spearman rank correlation using the rank transformation.
43///
44/// Ties are broken by average rank.
45/// Returns `None` under the same conditions as [`pearson_correlation`].
46pub fn spearman_correlation(x: &[f64], y: &[f64]) -> Option<f64> {
47    if x.len() != y.len() || x.len() < 2 {
48        return None;
49    }
50    let rx = average_ranks(x);
51    let ry = average_ranks(y);
52    pearson_correlation(&rx, &ry)
53}
54
55/// Assign average ranks to a slice, handling ties with average rank.
56fn average_ranks(data: &[f64]) -> Vec<f64> {
57    let n = data.len();
58    // Create (value, original_index) pairs sorted by value
59    let mut indexed: Vec<(f64, usize)> = data.iter().copied().enumerate().map(|(i, v)| (v, i)).collect();
60    indexed.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(std::cmp::Ordering::Equal));
61
62    let mut ranks = vec![0.0_f64; n];
63    let mut i = 0;
64    while i < n {
65        // Find the run of equal values
66        let mut j = i + 1;
67        while j < n && (indexed[j].0 - indexed[i].0).abs() < f64::EPSILON {
68            j += 1;
69        }
70        // Average rank for positions i..j (1-indexed ranks)
71        let avg_rank = (i + j + 1) as f64 / 2.0; // = ((i+1) + j) / 2 as 1-indexed
72        for k in i..j {
73            ranks[indexed[k].1] = avg_rank;
74        }
75        i = j;
76    }
77    ranks
78}
79
80// ─── Kendall tau-b ───────────────────────────────────────────────────────────
81
82/// Compute Kendall tau-b correlation coefficient.
83///
84/// O(n²) concordant/discordant pair counting with full tie correction.
85/// Returns `None` if either slice has fewer than 2 elements or slices differ in length.
86pub fn kendall_tau(x: &[f64], y: &[f64]) -> Option<f64> {
87    if x.len() != y.len() || x.len() < 2 {
88        return None;
89    }
90    let n = x.len();
91    let mut concordant: i64 = 0;
92    let mut discordant: i64 = 0;
93    let mut ties_x: i64 = 0;
94    let mut ties_y: i64 = 0;
95    let mut ties_xy: i64 = 0;
96
97    for i in 0..n {
98        for j in (i + 1)..n {
99            let dx = x[i] - x[j];
100            let dy = y[i] - y[j];
101            let prod = dx * dy;
102            let x_tied = dx.abs() < f64::EPSILON;
103            let y_tied = dy.abs() < f64::EPSILON;
104
105            if x_tied && y_tied {
106                ties_xy += 1;
107            } else if x_tied {
108                ties_x += 1;
109            } else if y_tied {
110                ties_y += 1;
111            } else if prod > 0.0 {
112                concordant += 1;
113            } else {
114                discordant += 1;
115            }
116        }
117    }
118
119    let total_pairs = (n as i64 * (n as i64 - 1)) / 2;
120    let n0 = total_pairs;
121    let n1 = n0 - ties_x - ties_xy;
122    let n2 = n0 - ties_y - ties_xy;
123
124    let denom = (n1 as f64 * n2 as f64).sqrt();
125    if denom == 0.0 {
126        return None;
127    }
128
129    let tau = (concordant - discordant) as f64 / denom;
130    Some(tau.clamp(-1.0, 1.0))
131}
132
133// ─── Symbol-based correlation matrix ─────────────────────────────────────────
134
135/// Full Pearson correlation matrix for a fixed set of symbols.
136///
137/// Constructed from complete return histories; not updated incrementally.
138/// For rolling / streaming use, see [`RollingCorrelation`].
139#[derive(Debug, Clone)]
140pub struct SymbolCorrelationMatrix {
141    /// Symbol labels in order.
142    pub symbols: Vec<String>,
143    /// n×n correlation matrix (row-major).
144    pub matrix: Vec<Vec<f64>>,
145    /// Dimension (number of symbols).
146    pub n: usize,
147}
148
149impl SymbolCorrelationMatrix {
150    /// Build the Pearson correlation matrix from full return series.
151    ///
152    /// `symbols` and `returns` must have the same length; all return slices must also
153    /// have the same length (the minimum across series is used).
154    pub fn from_returns(symbols: Vec<String>, returns: Vec<Vec<f64>>) -> Self {
155        let n = symbols.len();
156        let mut matrix = vec![vec![1.0_f64; n]; n];
157
158        for i in 0..n {
159            for j in (i + 1)..n {
160                let corr = pearson_correlation(&returns[i], &returns[j]).unwrap_or(0.0);
161                matrix[i][j] = corr;
162                matrix[j][i] = corr;
163            }
164        }
165
166        Self { symbols, matrix, n }
167    }
168
169    /// Get the correlation between symbols at indices `i` and `j`.
170    pub fn get(&self, i: usize, j: usize) -> f64 {
171        self.matrix[i][j]
172    }
173
174    /// Render the correlation matrix as a plain-text ASCII table.
175    pub fn to_table(&self) -> String {
176        // Determine column width
177        let col_w = self.symbols.iter().map(|s| s.len()).max().unwrap_or(6).max(6);
178        let fmt = |v: f64| format!("{:>width$.4}", v, width = col_w);
179        let pad = |s: &str| format!("{:>width$}", s, width = col_w);
180
181        let mut out = String::new();
182        // Header row
183        out.push_str(&" ".repeat(col_w + 1));
184        for sym in &self.symbols {
185            out.push(' ');
186            out.push_str(&pad(sym));
187        }
188        out.push('\n');
189
190        for (i, sym) in self.symbols.iter().enumerate() {
191            out.push_str(&pad(sym));
192            for j in 0..self.n {
193                out.push(' ');
194                out.push_str(&fmt(self.matrix[i][j]));
195            }
196            out.push('\n');
197        }
198        out
199    }
200
201    /// Returns all symbol pairs where `|correlation| > threshold`.
202    ///
203    /// Each entry is `(symbol_a, symbol_b, correlation)` for `i < j`.
204    pub fn highly_correlated(&self, threshold: f64) -> Vec<(String, String, f64)> {
205        let mut result = Vec::new();
206        for i in 0..self.n {
207            for j in (i + 1)..self.n {
208                let c = self.matrix[i][j];
209                if c.abs() > threshold {
210                    result.push((self.symbols[i].clone(), self.symbols[j].clone(), c));
211                }
212            }
213        }
214        result
215    }
216
217    /// Compute eigenvalues of the correlation matrix using the Jacobi sweep algorithm.
218    ///
219    /// Returns eigenvalues in descending order.
220    /// The Jacobi method iteratively zeroes off-diagonal elements via plane rotations.
221    pub fn eigenvalues(&self) -> Vec<f64> {
222        if self.n == 0 {
223            return vec![];
224        }
225        jacobi_eigenvalues(&self.matrix, self.n)
226    }
227}
228
229/// Jacobi eigenvalue algorithm for symmetric matrices.
230///
231/// Performs up to `max_sweeps * n*(n-1)/2` rotations, converging off-diagonal elements
232/// to near zero. Returns eigenvalues in descending order.
233fn jacobi_eigenvalues(matrix: &[Vec<f64>], n: usize) -> Vec<f64> {
234    // Copy into a flat mutable buffer
235    let mut a: Vec<f64> = matrix.iter().flat_map(|row| row.iter().copied()).collect();
236    let idx = |i: usize, j: usize| i * n + j;
237
238    let max_sweeps = 100;
239    let tol = 1e-10_f64;
240
241    for _ in 0..max_sweeps {
242        // Find max off-diagonal element
243        let mut max_val = 0.0_f64;
244        for i in 0..n {
245            for j in (i + 1)..n {
246                let v = a[idx(i, j)].abs();
247                if v > max_val {
248                    max_val = v;
249                }
250            }
251        }
252        if max_val < tol {
253            break;
254        }
255
256        // One Jacobi sweep over all off-diagonal pairs
257        for p in 0..n {
258            for q in (p + 1)..n {
259                let apq = a[idx(p, q)];
260                if apq.abs() < tol {
261                    continue;
262                }
263                let app = a[idx(p, p)];
264                let aqq = a[idx(q, q)];
265                let theta = 0.5 * (aqq - app) / apq;
266                let t = if theta >= 0.0 {
267                    1.0 / (theta + (1.0 + theta * theta).sqrt())
268                } else {
269                    -1.0 / (-theta + (1.0 + theta * theta).sqrt())
270                };
271                let c = 1.0 / (1.0 + t * t).sqrt();
272                let s = t * c;
273
274                // Update diagonal
275                a[idx(p, p)] = app - t * apq;
276                a[idx(q, q)] = aqq + t * apq;
277                a[idx(p, q)] = 0.0;
278                a[idx(q, p)] = 0.0;
279
280                // Update off-diagonal rows/columns
281                for r in 0..n {
282                    if r == p || r == q {
283                        continue;
284                    }
285                    let arp = a[idx(r, p)];
286                    let arq = a[idx(r, q)];
287                    a[idx(r, p)] = c * arp - s * arq;
288                    a[idx(p, r)] = a[idx(r, p)];
289                    a[idx(r, q)] = s * arp + c * arq;
290                    a[idx(q, r)] = a[idx(r, q)];
291                }
292            }
293        }
294    }
295
296    // Diagonal entries are eigenvalues
297    let mut eigs: Vec<f64> = (0..n).map(|i| a[idx(i, i)]).collect();
298    eigs.sort_by(|a, b| b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal));
299    eigs
300}
301
302// ─── Rolling correlation ──────────────────────────────────────────────────────
303
304/// Rolling window Pearson correlation matrix, updated tick-by-tick.
305///
306/// Each series is stored in a fixed-size `VecDeque`; once all series have
307/// accumulated `window` values, `compute_matrix` and `pairwise` become available.
308pub struct RollingCorrelation {
309    /// Rolling window size.
310    window: usize,
311    /// Per-symbol deques.
312    series: HashMap<String, VecDeque<f64>>,
313}
314
315impl RollingCorrelation {
316    /// Create a new rolling correlation tracker with the given window size.
317    pub fn new(window: usize) -> Self {
318        Self { window, series: HashMap::new() }
319    }
320
321    /// Push a new value for the given symbol.
322    ///
323    /// If the symbol is not yet tracked it is initialised automatically.
324    /// Once the deque reaches `window` length, the oldest value is evicted.
325    pub fn push(&mut self, symbol: &str, value: f64) {
326        let dq = self.series.entry(symbol.to_string()).or_insert_with(|| VecDeque::with_capacity(self.window));
327        if dq.len() >= self.window {
328            dq.pop_front();
329        }
330        dq.push_back(value);
331    }
332
333    /// Returns `true` when every tracked series has accumulated at least `window` values.
334    pub fn is_ready(&self) -> bool {
335        !self.series.is_empty() && self.series.values().all(|dq| dq.len() >= self.window)
336    }
337
338    /// Compute the full Pearson correlation matrix over all tracked symbols.
339    ///
340    /// Returns `None` if any series has fewer than `window` values.
341    pub fn compute_matrix(&self) -> Option<SymbolCorrelationMatrix> {
342        if !self.is_ready() {
343            return None;
344        }
345        let mut symbols: Vec<String> = self.series.keys().cloned().collect();
346        symbols.sort();
347        let returns: Vec<Vec<f64>> = symbols
348            .iter()
349            .map(|s| self.series[s].iter().copied().collect())
350            .collect();
351        Some(SymbolCorrelationMatrix::from_returns(symbols, returns))
352    }
353
354    /// Compute the rolling Pearson correlation for a specific pair.
355    ///
356    /// Returns `None` if either symbol is not tracked or has fewer than `window` values.
357    pub fn pairwise(&self, sym_a: &str, sym_b: &str) -> Option<f64> {
358        let a = self.series.get(sym_a)?;
359        let b = self.series.get(sym_b)?;
360        if a.len() < self.window || b.len() < self.window {
361            return None;
362        }
363        let va: Vec<f64> = a.iter().copied().collect();
364        let vb: Vec<f64> = b.iter().copied().collect();
365        pearson_correlation(&va, &vb)
366    }
367}
368
369// ─── tests ───────────────────────────────────────────────────────────────────
370
371#[cfg(test)]
372mod tests {
373    use super::*;
374
375    #[test]
376    fn test_pearson_perfect_positive() {
377        let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
378        let y = vec![2.0, 4.0, 6.0, 8.0, 10.0];
379        let r = pearson_correlation(&x, &y).unwrap();
380        assert!((r - 1.0).abs() < 1e-9, "r={r}");
381    }
382
383    #[test]
384    fn test_pearson_perfect_negative() {
385        let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
386        let y = vec![10.0, 8.0, 6.0, 4.0, 2.0];
387        let r = pearson_correlation(&x, &y).unwrap();
388        assert!((r + 1.0).abs() < 1e-9, "r={r}");
389    }
390
391    #[test]
392    fn test_pearson_insufficient_data() {
393        assert!(pearson_correlation(&[1.0], &[1.0]).is_none());
394        assert!(pearson_correlation(&[], &[]).is_none());
395    }
396
397    #[test]
398    fn test_pearson_zero_variance() {
399        let x = vec![5.0, 5.0, 5.0];
400        let y = vec![1.0, 2.0, 3.0];
401        assert!(pearson_correlation(&x, &y).is_none());
402    }
403
404    #[test]
405    fn test_spearman_perfect_positive() {
406        let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
407        let y = vec![10.0, 20.0, 30.0, 40.0, 50.0];
408        let r = spearman_correlation(&x, &y).unwrap();
409        assert!((r - 1.0).abs() < 1e-9, "r={r}");
410    }
411
412    #[test]
413    fn test_spearman_anti_correlation() {
414        let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
415        let y = vec![5.0, 4.0, 3.0, 2.0, 1.0];
416        let r = spearman_correlation(&x, &y).unwrap();
417        assert!((r + 1.0).abs() < 1e-9, "r={r}");
418    }
419
420    #[test]
421    fn test_spearman_rank_transform_with_ties() {
422        // With ties: ranks of [1,1,2] should be [1.5, 1.5, 3]
423        let ranks = average_ranks(&[1.0, 1.0, 2.0]);
424        assert!((ranks[0] - 1.5).abs() < 1e-9);
425        assert!((ranks[1] - 1.5).abs() < 1e-9);
426        assert!((ranks[2] - 3.0).abs() < 1e-9);
427    }
428
429    #[test]
430    fn test_kendall_perfect_concordant() {
431        let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
432        let y = vec![1.0, 2.0, 3.0, 4.0, 5.0];
433        let tau = kendall_tau(&x, &y).unwrap();
434        assert!((tau - 1.0).abs() < 1e-9, "tau={tau}");
435    }
436
437    #[test]
438    fn test_kendall_perfect_discordant() {
439        let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
440        let y = vec![5.0, 4.0, 3.0, 2.0, 1.0];
441        let tau = kendall_tau(&x, &y).unwrap();
442        assert!((tau + 1.0).abs() < 1e-9, "tau={tau}");
443    }
444
445    #[test]
446    fn test_symbol_correlation_matrix_from_returns() {
447        let symbols = vec!["A".to_string(), "B".to_string(), "C".to_string()];
448        let returns = vec![
449            vec![1.0, 2.0, 3.0, 4.0, 5.0],
450            vec![2.0, 4.0, 6.0, 8.0, 10.0], // perfectly correlated with A
451            vec![5.0, 4.0, 3.0, 2.0, 1.0],  // anti-correlated with A
452        ];
453        let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
454        assert!((mat.get(0, 1) - 1.0).abs() < 1e-9);
455        assert!((mat.get(0, 2) + 1.0).abs() < 1e-9);
456        assert_eq!(mat.get(0, 0), 1.0);
457    }
458
459    #[test]
460    fn test_highly_correlated_filter() {
461        let symbols = vec!["A".to_string(), "B".to_string(), "C".to_string()];
462        let returns = vec![
463            vec![1.0, 2.0, 3.0, 4.0, 5.0],
464            vec![2.0, 4.0, 6.0, 8.0, 10.0],
465            vec![5.0, 4.0, 3.0, 2.0, 1.0],
466        ];
467        let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
468        let high = mat.highly_correlated(0.9);
469        // All three pairs exceed |0.9| (|r|=1.0)
470        assert_eq!(high.len(), 3);
471    }
472
473    #[test]
474    fn test_highly_correlated_excludes_below_threshold() {
475        let symbols = vec!["A".to_string(), "B".to_string()];
476        let returns = vec![
477            vec![1.0, 2.0, 3.0, 4.0, 5.0],
478            vec![1.0, 1.5, 1.0, 1.5, 1.0], // low correlation
479        ];
480        let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
481        let high = mat.highly_correlated(0.99);
482        assert!(high.is_empty());
483    }
484
485    #[test]
486    fn test_to_table_contains_symbols() {
487        let symbols = vec!["BTC".to_string(), "ETH".to_string()];
488        let returns = vec![
489            vec![1.0, 2.0, 3.0],
490            vec![1.0, 2.0, 3.0],
491        ];
492        let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
493        let table = mat.to_table();
494        assert!(table.contains("BTC"));
495        assert!(table.contains("ETH"));
496    }
497
498    #[test]
499    fn test_rolling_correlation_not_ready_until_window() {
500        let mut rc = RollingCorrelation::new(5);
501        for i in 0..4 {
502            rc.push("A", i as f64);
503            rc.push("B", i as f64 * 2.0);
504        }
505        assert!(!rc.is_ready());
506        assert!(rc.compute_matrix().is_none());
507        assert!(rc.pairwise("A", "B").is_none());
508    }
509
510    #[test]
511    fn test_rolling_correlation_ready_after_window() {
512        let mut rc = RollingCorrelation::new(5);
513        for i in 0..5 {
514            rc.push("A", i as f64);
515            rc.push("B", i as f64 * 2.0);
516        }
517        assert!(rc.is_ready());
518        let r = rc.pairwise("A", "B").unwrap();
519        assert!((r - 1.0).abs() < 1e-9, "r={r}");
520    }
521
522    #[test]
523    fn test_rolling_window_evicts_old_values() {
524        let mut rc = RollingCorrelation::new(3);
525        // Push 5 values; only last 3 count
526        for i in 0..5 {
527            rc.push("A", i as f64);
528        }
529        let dq = &rc.series["A"];
530        assert_eq!(dq.len(), 3);
531        assert_eq!(dq[0], 2.0);
532        assert_eq!(dq[2], 4.0);
533    }
534
535    #[test]
536    fn test_rolling_correlation_matrix() {
537        let mut rc = RollingCorrelation::new(5);
538        for i in 0..5 {
539            let v = i as f64;
540            rc.push("X", v);
541            rc.push("Y", -v);
542        }
543        let mat = rc.compute_matrix().unwrap();
544        // X and Y are anti-correlated
545        let x_idx = mat.symbols.iter().position(|s| s == "X").unwrap();
546        let y_idx = mat.symbols.iter().position(|s| s == "Y").unwrap();
547        assert!((mat.get(x_idx, y_idx) + 1.0).abs() < 1e-9);
548    }
549
550    #[test]
551    fn test_eigenvalues_length() {
552        let symbols = vec!["A".to_string(), "B".to_string(), "C".to_string()];
553        let returns = vec![
554            vec![1.0, 2.0, 3.0, 4.0, 5.0],
555            vec![2.0, 4.0, 6.0, 8.0, 10.0],
556            vec![5.0, 4.0, 3.0, 2.0, 1.0],
557        ];
558        let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
559        let eigs = mat.eigenvalues();
560        assert_eq!(eigs.len(), 3);
561    }
562
563    #[test]
564    fn test_eigenvalues_descending() {
565        let symbols = vec!["A".to_string(), "B".to_string(), "C".to_string()];
566        let returns = vec![
567            vec![1.0, 2.0, 3.0, 4.0, 5.0],
568            vec![5.0, 3.0, 1.0, 4.0, 2.0],
569            vec![2.0, 5.0, 1.0, 3.0, 4.0],
570        ];
571        let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
572        let eigs = mat.eigenvalues();
573        for i in 0..eigs.len() - 1 {
574            assert!(eigs[i] >= eigs[i + 1] - 1e-9, "eigs not descending: {:?}", eigs);
575        }
576    }
577}