1use crate::error::{Error, Result};
4use crate::traits::Indicator;
5
6#[derive(Debug, Clone)]
25pub struct Ema {
26 period: usize,
27 alpha: f64,
28 one_minus_alpha: f64,
31 current: f64,
35 seeded: bool,
37 warmup_sum: f64,
43 warmup_count: usize,
45}
46
47impl Ema {
48 pub fn new(period: usize) -> Result<Self> {
54 if period == 0 {
55 return Err(Error::PeriodZero);
56 }
57 if period > crate::error::MAX_PERIOD {
58 return Err(Error::InvalidPeriod {
59 message: crate::error::PERIOD_ABOVE_MAX,
60 });
61 }
62 let alpha = 2.0 / (period as f64 + 1.0);
63 Ok(Self {
64 period,
65 alpha,
66 one_minus_alpha: 1.0 - alpha,
67 current: 0.0,
68 seeded: false,
69 warmup_sum: -0.0,
70 warmup_count: 0,
71 })
72 }
73
74 pub(crate) fn with_period_and_alpha(period: usize, alpha: f64) -> Self {
79 Self {
80 period,
81 alpha,
82 one_minus_alpha: 1.0 - alpha,
83 current: 0.0,
84 seeded: false,
85 warmup_sum: -0.0,
86 warmup_count: 0,
87 }
88 }
89
90 pub fn with_alpha(alpha: f64) -> Result<Self> {
100 if !alpha.is_finite() || alpha <= 0.0 || alpha > 1.0 {
101 return Err(Error::InvalidPeriod {
102 message: "alpha must be in (0.0, 1.0]",
103 });
104 }
105 Ok(Self {
106 period: 1,
107 alpha,
108 one_minus_alpha: 1.0 - alpha,
109 current: 0.0,
110 seeded: false,
111 warmup_sum: -0.0,
112 warmup_count: 0,
113 })
114 }
115
116 pub const fn period(&self) -> usize {
118 self.period
119 }
120
121 pub const fn alpha(&self) -> f64 {
123 self.alpha
124 }
125
126 pub(crate) const fn one_minus_alpha(&self) -> f64 {
128 self.one_minus_alpha
129 }
130
131 pub const fn value(&self) -> Option<f64> {
133 if self.seeded {
134 Some(self.current)
135 } else {
136 None
137 }
138 }
139
140 pub(crate) fn is_fresh(&self) -> bool {
143 !self.seeded && self.warmup_count == 0
144 }
145
146 pub(crate) fn seed_to(&mut self, current: f64) {
152 self.current = current;
153 self.seeded = true;
154 }
155
156 pub fn batch_nan(&mut self, inputs: &[f64]) -> Vec<f64> {
162 crate::traits::BatchNanExt::batch_nan(self, inputs)
163 }
164
165 pub(crate) fn step_unchecked(&mut self, input: f64) -> Option<f64> {
168 if self.seeded {
169 let new = self
170 .alpha
171 .mul_add(input, self.one_minus_alpha * self.current);
172 self.current = new;
173 return Some(new);
174 }
175 self.warmup_sum += input;
176 self.warmup_count += 1;
177 if self.warmup_count == self.period {
178 let seed = self.warmup_sum / self.period as f64;
179 self.current = seed;
180 self.seeded = true;
181 return Some(seed);
182 }
183 None
184 }
185}
186
187impl Indicator for Ema {
188 type Input = f64;
189 type Output = f64;
190
191 #[inline]
192 fn update(&mut self, input: f64) -> Option<f64> {
193 if !input.is_finite() {
194 return None;
195 }
196 self.step_unchecked(input)
197 }
198
199 fn reset(&mut self) {
200 self.current = 0.0;
201 self.seeded = false;
202 self.warmup_sum = -0.0;
203 self.warmup_count = 0;
204 }
205
206 #[inline]
207 fn warmup_period(&self) -> usize {
208 self.period
209 }
210
211 #[inline]
212 fn is_ready(&self) -> bool {
213 self.seeded
214 }
215
216 #[inline]
217 fn name(&self) -> &'static str {
218 "EMA"
219 }
220
221 fn batch_nan_into(&mut self, inputs: &[f64], out: &mut [f64]) {
229 assert_eq!(
230 inputs.len(),
231 out.len(),
232 "batch output length must equal input length"
233 );
234 let p = self.period;
235 if self.seeded || self.warmup_count != 0 || !inputs.iter().all(|x| x.is_finite()) {
236 for (slot, &x) in out.iter_mut().zip(inputs) {
237 *slot = self.update(x).unwrap_or(f64::NAN);
238 }
239 return;
240 }
241
242 let n = inputs.len();
243 if n < p {
244 for &x in inputs {
246 self.warmup_sum += x;
247 }
248 self.warmup_count = n;
249 out.fill(f64::NAN);
250 return;
251 }
252
253 out[..p - 1].fill(f64::NAN);
255 let seed_sum = inputs[..p].iter().copied().sum::<f64>();
256 let seed = seed_sum / p as f64;
257 out[p - 1] = seed;
258 let mut cur = seed;
259 let (alpha, oma) = (self.alpha, self.one_minus_alpha);
260 for (slot, &x) in out[p..].iter_mut().zip(&inputs[p..]) {
261 cur = alpha.mul_add(x, oma * cur);
262 *slot = cur;
263 }
264
265 self.current = cur;
268 self.seeded = true;
269 self.warmup_sum = seed_sum;
270 self.warmup_count = p;
271 }
272
273 fn batch_fast_into(&mut self, inputs: &[f64], out: &mut [f64]) {
281 assert_eq!(
282 inputs.len(),
283 out.len(),
284 "batch output length must equal input length"
285 );
286 let p = self.period;
287 if self.seeded
288 || self.warmup_count != 0
289 || inputs.len() < p
290 || !crate::fast::in_range(inputs)
291 {
292 self.batch_nan_into(inputs, out);
293 return;
294 }
295 let (last, seed_sum) = wickra_simd::dispatch(crate::fast::EmaFast {
296 x: inputs,
297 period: p,
298 alpha: self.alpha,
299 one_minus_alpha: self.one_minus_alpha,
300 out,
301 _borrow: std::marker::PhantomData,
302 });
303 self.current = last;
304 self.seeded = true;
305 self.warmup_sum = seed_sum;
306 self.warmup_count = p;
307 }
308}
309
310#[cfg(test)]
311mod tests {
312 use super::*;
313
314 #[test]
319 fn rejects_a_period_above_the_maximum() {
320 assert!(matches!(
321 Ema::new(usize::MAX),
322 Err(Error::InvalidPeriod { .. })
323 ));
324 assert!(matches!(
325 Ema::new(crate::error::MAX_PERIOD + 1),
326 Err(Error::InvalidPeriod { .. })
327 ));
328 assert!(matches!(
329 Ema::new(1_000_000_000),
330 Err(Error::InvalidPeriod { .. })
331 ));
332 assert!(Ema::new(14).is_ok());
334 }
335 use crate::traits::BatchExt;
336 use approx::assert_relative_eq;
337
338 fn ema_naive(prices: &[f64], period: usize) -> Vec<Option<f64>> {
340 let alpha = 2.0 / (period as f64 + 1.0);
341 let mut out = Vec::with_capacity(prices.len());
342 let mut state: Option<f64> = None;
343 for (i, &p) in prices.iter().enumerate() {
344 if let Some(prev) = state {
345 let v = alpha * p + (1.0 - alpha) * prev;
346 state = Some(v);
347 out.push(Some(v));
348 } else if i + 1 == period {
349 let seed = prices[..period].iter().sum::<f64>() / period as f64;
350 state = Some(seed);
351 out.push(Some(seed));
352 } else {
353 out.push(None);
354 }
355 }
356 out
357 }
358
359 #[test]
360 fn new_rejects_zero_period() {
361 assert!(matches!(Ema::new(0), Err(Error::PeriodZero)));
362 }
363
364 #[test]
369 fn accessors_and_metadata() {
370 let ema = Ema::new(14).unwrap();
371 assert_eq!(ema.period(), 14);
372 assert_eq!(ema.warmup_period(), 14);
373 assert_eq!(ema.name(), "EMA");
374 }
375
376 #[test]
377 fn warmup_returns_none_until_seed() {
378 let mut ema = Ema::new(3).unwrap();
379 assert_eq!(ema.update(1.0), None);
380 assert_eq!(ema.update(2.0), None);
381 assert_eq!(ema.update(3.0), Some(2.0)); }
383
384 #[test]
385 fn first_value_equals_sma_seed() {
386 let mut ema = Ema::new(5).unwrap();
387 let inputs = [10.0, 20.0, 30.0, 40.0, 50.0];
388 let mut last = None;
389 for v in inputs {
390 last = ema.update(v);
391 }
392 assert_relative_eq!(last.unwrap(), 30.0, epsilon = 1e-12);
393 }
394
395 #[test]
396 fn alpha_matches_period_formula() {
397 let ema = Ema::new(10).unwrap();
398 assert_relative_eq!(ema.alpha(), 2.0 / 11.0, epsilon = 1e-15);
399 }
400
401 #[test]
402 fn step_after_seed_uses_alpha_formula() {
403 let mut ema = Ema::new(3).unwrap();
406 ema.batch(&[1.0, 2.0, 3.0]);
407 assert_relative_eq!(ema.update(10.0).unwrap(), 6.0, epsilon = 1e-12);
408 }
409
410 #[test]
411 fn constant_series_converges_to_constant() {
412 let mut ema = Ema::new(10).unwrap();
413 let out = ema.batch(&[42.0_f64; 100]);
414 for x in out.iter().skip(9) {
415 assert_relative_eq!(x.unwrap(), 42.0, epsilon = 1e-9);
416 }
417 }
418
419 #[test]
420 fn with_alpha_validates_range() {
421 assert!(Ema::with_alpha(0.5).is_ok());
422 assert!(Ema::with_alpha(1.0).is_ok());
423 assert!(matches!(
424 Ema::with_alpha(0.0),
425 Err(Error::InvalidPeriod { .. })
426 ));
427 assert!(matches!(
428 Ema::with_alpha(1.5),
429 Err(Error::InvalidPeriod { .. })
430 ));
431 assert!(matches!(
432 Ema::with_alpha(f64::NAN),
433 Err(Error::InvalidPeriod { .. })
434 ));
435 }
436
437 #[test]
438 fn reset_clears_state() {
439 let mut ema = Ema::new(3).unwrap();
440 ema.batch(&[1.0, 2.0, 3.0]);
441 assert!(ema.is_ready());
442 ema.reset();
443 assert!(!ema.is_ready());
444 assert_eq!(ema.update(1.0), None);
445 }
446
447 #[test]
448 fn batch_equals_streaming() {
449 let prices: Vec<f64> = (1..=30).map(f64::from).collect();
450 let mut a = Ema::new(5).unwrap();
451 let mut b = Ema::new(5).unwrap();
452 assert_eq!(
453 a.batch(&prices),
454 prices.iter().map(|p| b.update(*p)).collect::<Vec<_>>()
455 );
456 }
457
458 #[test]
459 fn ignores_non_finite_input() {
460 let mut ema = Ema::new(3).unwrap();
461 ema.batch(&[1.0, 2.0, 3.0]);
462 let before = ema.value();
463 assert_eq!(ema.update(f64::NAN), None);
464 assert_eq!(ema.update(f64::INFINITY), None);
465 assert_eq!(ema.value(), before);
467 }
468
469 fn bits_eq(a: &[f64], b: &[f64]) -> bool {
470 a.len() == b.len()
471 && a.iter()
472 .zip(b)
473 .all(|(x, y)| x == y || (x.is_nan() && y.is_nan()))
474 }
475
476 fn ema_replay(period: usize, series: &[f64]) -> Vec<f64> {
477 let mut e = Ema::new(period).unwrap();
478 series
479 .iter()
480 .map(|&x| e.update(x).unwrap_or(f64::NAN))
481 .collect()
482 }
483
484 #[test]
485 fn batch_nan_fast_path_is_bit_identical() {
486 let series: Vec<f64> = (0..300)
487 .map(|i| (f64::from(i) * 0.25).cos() * 8.0 + 40.0)
488 .collect();
489 let mut ema = Ema::new(14).unwrap();
490 let got = ema.batch_nan(&series);
491 assert!(bits_eq(&got, &ema_replay(14, &series)));
492 let mut ref_ema = Ema::new(14).unwrap();
493 for &x in &series {
494 ref_ema.update(x);
495 }
496 assert_eq!(ema.update(7.5), ref_ema.update(7.5));
497 }
498
499 #[test]
503 fn seed_matches_a_buffered_sum_bit_for_bit() {
504 let windows: [&[f64]; 3] = [
505 &[-0.0, -0.0, -0.0],
506 &[1e16, 1.0, -1e16, 3.5, 0.1],
507 &[0.3, 0.1, 0.7, 0.2, 0.9, 0.4, 0.6],
508 ];
509 for window in windows {
510 let buffered = window.iter().copied().sum::<f64>() / window.len() as f64;
511 let mut ema = Ema::new(window.len()).unwrap();
512 let seed = window.iter().filter_map(|&x| ema.update(x)).last().unwrap();
513 assert_eq!(seed.to_bits(), buffered.to_bits());
514 let mut out = vec![0.0; window.len()];
515 Ema::new(window.len())
516 .unwrap()
517 .batch_nan_into(window, &mut out);
518 assert_eq!(out[window.len() - 1].to_bits(), buffered.to_bits());
519 }
520 }
521
522 #[test]
525 fn batch_nan_into_overwrites_a_dirty_buffer() {
526 let series: Vec<f64> = (0..120).map(|i| f64::from(i % 11) * 0.75 + 20.0).collect();
527 let mut out = vec![-1.0; series.len()];
528 Ema::new(10).unwrap().batch_nan_into(&series, &mut out);
529 assert!(bits_eq(&out, &ema_replay(10, &series)));
530 let mut short = [-1.0; 4];
531 Ema::new(10)
532 .unwrap()
533 .batch_nan_into(&series[..4], &mut short);
534 assert!(short.iter().all(|x| x.is_nan()));
535 }
536
537 #[test]
538 fn batch_nan_falls_back_on_non_finite() {
539 let series = [1.0, 2.0, 3.0, f64::INFINITY, 5.0, 6.0, 7.0];
540 let mut ema = Ema::new(3).unwrap();
541 assert!(bits_eq(&ema.batch_nan(&series), &ema_replay(3, &series)));
542 }
543
544 #[test]
545 fn batch_nan_falls_back_when_warming() {
546 let mut ema = Ema::new(3).unwrap();
547 ema.update(10.0); let series = [1.0, 2.0, 3.0, 4.0];
549 let mut ref_ema = Ema::new(3).unwrap();
550 ref_ema.update(10.0);
551 let want: Vec<f64> = series
552 .iter()
553 .map(|&x| ref_ema.update(x).unwrap_or(f64::NAN))
554 .collect();
555 assert!(bits_eq(&ema.batch_nan(&series), &want));
556 }
557
558 #[test]
559 fn batch_nan_sub_period_slice_stays_unseeded() {
560 let series = [1.0, 2.0];
561 let mut ema = Ema::new(5).unwrap();
562 let got = ema.batch_nan(&series);
563 assert!(got.iter().all(|x| x.is_nan()) && got.len() == 2);
564 assert!(!ema.is_ready());
565 assert!(bits_eq(
567 &[ema.update(3.0).unwrap_or(f64::NAN)],
568 &[ema_replay(5, &[1.0, 2.0, 3.0])[2]]
569 ));
570 }
571
572 #[test]
573 fn with_period_and_alpha_sets_fields_and_warmup() {
574 let ema = Ema::with_period_and_alpha(12, 0.15);
575 assert_eq!(ema.period(), 12);
576 assert_eq!(ema.alpha().to_bits(), 0.15f64.to_bits());
577 assert_eq!(ema.one_minus_alpha().to_bits(), (1.0f64 - 0.15).to_bits());
578 assert_eq!(ema.warmup_period(), 12);
579 assert!(ema.is_fresh());
580 assert!(!ema.is_ready());
581 assert_eq!(ema.value(), None);
582 }
583
584 #[test]
585 fn with_period_and_alpha_hand_computed() {
586 let mut ema = Ema::with_period_and_alpha(3, 0.15);
590 assert_eq!(ema.update(1.0), None);
591 assert_eq!(ema.update(2.0), None);
592 assert_eq!(ema.update(3.0), Some(2.0));
593 assert_relative_eq!(ema.update(10.0).unwrap(), 3.2, epsilon = 1e-12);
594 assert_relative_eq!(ema.update(0.0).unwrap(), 2.72, epsilon = 1e-12);
595 }
596
597 #[test]
598 fn with_period_and_alpha_matches_new_when_alpha_is_standard() {
599 let series: Vec<f64> = (0..80)
601 .map(|i| (f64::from(i) * 0.3).sin() * 5.0 + 50.0)
602 .collect();
603 let got = Ema::with_period_and_alpha(9, 2.0 / 10.0).batch_nan(&series);
604 assert!(bits_eq(&got, &ema_replay(9, &series)));
605 }
606
607 #[test]
608 fn with_period_and_alpha_batch_equals_streaming_and_reset() {
609 let series: Vec<f64> = (0..120)
610 .map(|i| (f64::from(i) * 0.21).cos() * 7.0 + 30.0)
611 .collect();
612 let mut stream = Ema::with_period_and_alpha(26, 0.075);
613 let streamed: Vec<f64> = series
614 .iter()
615 .map(|&x| stream.update(x).unwrap_or(f64::NAN))
616 .collect();
617 let mut batch = Ema::with_period_and_alpha(26, 0.075);
618 assert!(bits_eq(&batch.batch_nan(&series), &streamed));
619 assert_eq!(batch.update(12.5), stream.update(12.5));
620 let opt = Ema::with_period_and_alpha(26, 0.075).batch(&series);
621 assert!(opt
622 .iter()
623 .zip(&streamed)
624 .all(|(o, s)| o.unwrap_or(f64::NAN).to_bits() == s.to_bits()));
625 assert!(opt[..25].iter().all(Option::is_none));
626 assert!(opt[25].is_some());
627 batch.reset();
628 assert!(bits_eq(&batch.batch_nan(&series), &streamed));
629 }
630
631 proptest::proptest! {
632 #![proptest_config(proptest::test_runner::Config::with_cases(48))]
633 #[test]
634 fn ema_matches_naive(
635 period in 1usize..20,
636 prices in proptest::collection::vec(-1000.0_f64..1000.0, 0..150),
637 ) {
638 let mut ema = Ema::new(period).unwrap();
639 let got = ema.batch(&prices);
640 let want = ema_naive(&prices, period);
641 proptest::prop_assert_eq!(got.len(), want.len());
642 for (g, w) in got.iter().zip(want.iter()) {
643 match (g, w) {
644 (None, None) => {}
645 (Some(a), Some(b)) => proptest::prop_assert!(
646 (a - b).abs() <= 1e-9 * a.abs().max(1.0),
647 "got={a} want={b}"
648 ),
649 _ => proptest::prop_assert!(false, "warmup mismatch"),
650 }
651 }
652 }
653 }
654}