fin_primitives/correlation/
mod.rs1pub mod stats;
23
24use crate::error::FinError;
25use std::collections::VecDeque;
26
27pub const DEFAULT_REDUNDANCY_THRESHOLD: f64 = 0.95;
29
30#[derive(Debug)]
49pub struct CorrelationMatrix {
50 n: usize,
52 window: usize,
54 threshold: f64,
56 buf: VecDeque<Vec<f64>>,
58}
59
60impl CorrelationMatrix {
61 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 pub fn with_defaults(n_indicators: usize, window: usize) -> Result<Self, FinError> {
98 Self::new(n_indicators, window, DEFAULT_REDUNDANCY_THRESHOLD)
99 }
100
101 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 pub fn is_ready(&self) -> bool {
124 self.buf.len() >= self.window
125 }
126
127 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 Some(r.clamp(-1.0, 1.0))
161 }
162
163 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 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 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 pub fn n_indicators(&self) -> usize {
222 self.n
223 }
224
225 pub fn window(&self) -> usize {
227 self.window
228 }
229
230 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 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 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 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 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 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 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}