1use crate::error::{Error, Result};
4use crate::indicators::rolling_moments::ShiftedPairMoments;
5use crate::traits::Indicator;
6
7#[derive(Debug, Clone)]
46pub struct PearsonCorrelation {
47 period: usize,
48 buf: Box<[(f64, f64)]>,
51 head: usize,
52 count: usize,
54 moments: ShiftedPairMoments,
55}
56
57impl PearsonCorrelation {
58 pub fn new(period: usize) -> Result<Self> {
64 if period < 2 {
65 return Err(Error::InvalidPeriod {
66 message: "pearson correlation needs period >= 2",
67 });
68 }
69 if period > crate::error::MAX_PERIOD {
70 return Err(Error::InvalidPeriod {
71 message: crate::error::PERIOD_ABOVE_MAX,
72 });
73 }
74 Ok(Self {
75 period,
76 buf: vec![(0.0, 0.0); period].into_boxed_slice(),
77 head: 0,
78 count: 0,
79 moments: ShiftedPairMoments::new(),
80 })
81 }
82
83 pub const fn period(&self) -> usize {
85 self.period
86 }
87}
88
89impl PearsonCorrelation {
90 pub fn batch_pairs_into(&mut self, a: &[f64], b: &[f64], out: &mut [f64]) {
97 assert!(
98 a.len() == b.len() && out.len() == a.len(),
99 "both series and the output must be equal length"
100 );
101 for ((slot, &x), &y) in out.iter_mut().zip(a).zip(b) {
102 *slot = self.update((x, y)).unwrap_or(f64::NAN);
103 }
104 }
105
106 pub fn batch_pairs_fast_into(&mut self, a: &[f64], b: &[f64], out: &mut [f64]) {
121 assert!(
122 a.len() == b.len() && out.len() == a.len(),
123 "both series and the output must be equal length"
124 );
125 let p = self.period;
126 let n = a.len();
127 if self.count != 0 || n < p || !crate::fast::in_range(a) || !crate::fast::in_range(b) {
128 self.batch_pairs_into(a, b, out);
129 return;
130 }
131 crate::fast::with_scratch(crate::fast::power_scratch_len(5, p), |scratch| {
132 wickra_simd::dispatch(crate::fast::PearsonFast {
133 a,
134 b,
135 period: p,
136 scratch,
137 out,
138 _borrow: std::marker::PhantomData,
139 });
140 });
141 self.reset();
142 for (&x, &y) in a[n - p..].iter().zip(&b[n - p..]) {
143 let _ = self.update((x, y));
144 }
145 }
146}
147
148impl Indicator for PearsonCorrelation {
149 type Input = (f64, f64);
150 type Output = f64;
151
152 #[inline]
153 fn update(&mut self, input: (f64, f64)) -> Option<f64> {
154 let (x, y) = input;
155 if !x.is_finite() || !y.is_finite() {
156 return None;
157 }
158 let slot = &mut self.buf[self.head];
160 if self.count == self.period {
161 let (ox, oy) = std::mem::replace(slot, (x, y));
162 self.moments.evict(ox, oy);
163 } else {
164 *slot = (x, y);
165 self.count += 1;
166 }
167 self.head += 1;
168 if self.head == self.period {
169 self.head = 0;
170 }
171 self.moments.push(x, y);
172 if self.moments.needs_reseed(self.period) {
173 let (older, newer) = if self.count == self.period {
176 (&self.buf[self.head..], &self.buf[..self.head])
177 } else {
178 (&self.buf[..self.count], &self.buf[..0])
179 };
180 self.moments.reseed(older.iter().chain(newer).copied());
181 }
182 if self.count < self.period {
183 return None;
184 }
185 let var_x = self.moments.var_a(self.period);
186 let var_y = self.moments.var_b(self.period);
187 let cov = self.moments.cov(self.period);
188 let denom = (var_x * var_y).sqrt();
189 if denom == 0.0 {
190 return Some(0.0);
192 }
193 Some((cov / denom).clamp(-1.0, 1.0))
194 }
195
196 fn reset(&mut self) {
197 self.head = 0;
198 self.count = 0;
199 self.moments.reset();
200 }
201
202 #[inline]
203 fn warmup_period(&self) -> usize {
204 self.period
205 }
206
207 #[inline]
208 fn is_ready(&self) -> bool {
209 self.count == self.period
210 }
211
212 #[inline]
213 fn name(&self) -> &'static str {
214 "PearsonCorrelation"
215 }
216}
217
218#[cfg(test)]
219mod tests {
220 use super::*;
221 use crate::traits::BatchExt;
222 use approx::assert_relative_eq;
223
224 #[test]
225 fn rejects_period_below_two() {
226 assert!(PearsonCorrelation::new(0).is_err());
227 assert!(PearsonCorrelation::new(1).is_err());
228 assert!(PearsonCorrelation::new(2).is_ok());
229 }
230
231 #[test]
232 fn accessors_and_metadata() {
233 let p = PearsonCorrelation::new(14).unwrap();
234 assert_eq!(p.period(), 14);
235 assert_eq!(p.warmup_period(), 14);
236 assert_eq!(p.name(), "PearsonCorrelation");
237 }
238
239 #[test]
240 fn perfect_positive_is_one() {
241 let pairs: Vec<(f64, f64)> = (0..10)
242 .map(|i| (f64::from(i), 3.0 * f64::from(i) + 1.0))
243 .collect();
244 let last = PearsonCorrelation::new(5)
245 .unwrap()
246 .batch(&pairs)
247 .into_iter()
248 .flatten()
249 .last()
250 .unwrap();
251 assert_relative_eq!(last, 1.0, epsilon = 1e-9);
252 }
253
254 #[test]
255 fn perfect_negative_is_minus_one() {
256 let pairs: Vec<(f64, f64)> = (0..10)
257 .map(|i| (f64::from(i), -2.0 * f64::from(i) + 5.0))
258 .collect();
259 let last = PearsonCorrelation::new(5)
260 .unwrap()
261 .batch(&pairs)
262 .into_iter()
263 .flatten()
264 .last()
265 .unwrap();
266 assert_relative_eq!(last, -1.0, epsilon = 1e-9);
267 }
268
269 #[test]
270 fn constant_channel_yields_zero() {
271 let pairs: Vec<(f64, f64)> = (0..10).map(|i| (f64::from(i), 7.0)).collect();
272 let last = PearsonCorrelation::new(5)
273 .unwrap()
274 .batch(&pairs)
275 .into_iter()
276 .flatten()
277 .last()
278 .unwrap();
279 assert_relative_eq!(last, 0.0, epsilon = 1e-12);
280 }
281
282 #[test]
283 fn output_in_minus_one_to_one_range() {
284 let pairs: Vec<(f64, f64)> = (0..60)
285 .map(|i| {
286 let t = f64::from(i);
287 (100.0 + t.sin() * 5.0, 50.0 + (t * 0.3).cos() * 3.0)
288 })
289 .collect();
290 let mut p = PearsonCorrelation::new(20).unwrap();
291 for v in p.batch(&pairs).into_iter().flatten() {
292 assert!((-1.0..=1.0).contains(&v));
293 }
294 }
295
296 #[test]
297 fn reset_clears_state() {
298 let mut p = PearsonCorrelation::new(5).unwrap();
299 p.batch(&[(1.0, 2.0), (2.0, 4.0), (3.0, 6.0), (4.0, 8.0), (5.0, 10.0)]);
300 assert!(p.is_ready());
301 p.reset();
302 assert!(!p.is_ready());
303 assert_eq!(p.update((1.0, 1.0)), None);
304 }
305
306 #[test]
307 fn batch_equals_streaming() {
308 let pairs: Vec<(f64, f64)> = (0..60)
309 .map(|i| {
310 let t = f64::from(i);
311 (t.sin(), (t * 0.5).cos())
312 })
313 .collect();
314 let batch = PearsonCorrelation::new(14).unwrap().batch(&pairs);
315 let mut b = PearsonCorrelation::new(14).unwrap();
316 let streamed: Vec<_> = pairs.iter().map(|p| b.update(*p)).collect();
317 assert_eq!(batch, streamed);
318 }
319
320 #[test]
321 fn non_finite_input_returns_none() {
322 let mut p = PearsonCorrelation::new(3).unwrap();
323 assert_eq!(p.update((f64::NAN, 1.0)), None);
324 assert_eq!(p.update((1.0, f64::INFINITY)), None);
325 assert_eq!(p.update((1.0, 2.0)), None);
327 assert_eq!(p.update((2.0, 5.0)), None);
328 assert!(p.update((3.0, 7.0)).is_some());
329 }
330}