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 fn with_alpha(alpha: f64) -> Result<Self> {
84 if !alpha.is_finite() || alpha <= 0.0 || alpha > 1.0 {
85 return Err(Error::InvalidPeriod {
86 message: "alpha must be in (0.0, 1.0]",
87 });
88 }
89 Ok(Self {
90 period: 1,
91 alpha,
92 one_minus_alpha: 1.0 - alpha,
93 current: 0.0,
94 seeded: false,
95 warmup_sum: -0.0,
96 warmup_count: 0,
97 })
98 }
99
100 pub const fn period(&self) -> usize {
102 self.period
103 }
104
105 pub const fn alpha(&self) -> f64 {
107 self.alpha
108 }
109
110 pub(crate) const fn one_minus_alpha(&self) -> f64 {
112 self.one_minus_alpha
113 }
114
115 pub const fn value(&self) -> Option<f64> {
117 if self.seeded {
118 Some(self.current)
119 } else {
120 None
121 }
122 }
123
124 pub(crate) fn is_fresh(&self) -> bool {
127 !self.seeded && self.warmup_count == 0
128 }
129
130 pub(crate) fn seed_to(&mut self, current: f64) {
136 self.current = current;
137 self.seeded = true;
138 }
139
140 pub fn batch_nan(&mut self, inputs: &[f64]) -> Vec<f64> {
146 crate::traits::BatchNanExt::batch_nan(self, inputs)
147 }
148
149 pub(crate) fn step_unchecked(&mut self, input: f64) -> Option<f64> {
152 if self.seeded {
153 let new = self
154 .alpha
155 .mul_add(input, self.one_minus_alpha * self.current);
156 self.current = new;
157 return Some(new);
158 }
159 self.warmup_sum += input;
160 self.warmup_count += 1;
161 if self.warmup_count == self.period {
162 let seed = self.warmup_sum / self.period as f64;
163 self.current = seed;
164 self.seeded = true;
165 return Some(seed);
166 }
167 None
168 }
169}
170
171impl Indicator for Ema {
172 type Input = f64;
173 type Output = f64;
174
175 #[inline]
176 fn update(&mut self, input: f64) -> Option<f64> {
177 if !input.is_finite() {
178 return None;
179 }
180 self.step_unchecked(input)
181 }
182
183 fn reset(&mut self) {
184 self.current = 0.0;
185 self.seeded = false;
186 self.warmup_sum = -0.0;
187 self.warmup_count = 0;
188 }
189
190 #[inline]
191 fn warmup_period(&self) -> usize {
192 self.period
193 }
194
195 #[inline]
196 fn is_ready(&self) -> bool {
197 self.seeded
198 }
199
200 #[inline]
201 fn name(&self) -> &'static str {
202 "EMA"
203 }
204
205 fn batch_nan_into(&mut self, inputs: &[f64], out: &mut [f64]) {
213 assert_eq!(
214 inputs.len(),
215 out.len(),
216 "batch output length must equal input length"
217 );
218 let p = self.period;
219 if self.seeded || self.warmup_count != 0 || !inputs.iter().all(|x| x.is_finite()) {
220 for (slot, &x) in out.iter_mut().zip(inputs) {
221 *slot = self.update(x).unwrap_or(f64::NAN);
222 }
223 return;
224 }
225
226 let n = inputs.len();
227 if n < p {
228 for &x in inputs {
230 self.warmup_sum += x;
231 }
232 self.warmup_count = n;
233 out.fill(f64::NAN);
234 return;
235 }
236
237 out[..p - 1].fill(f64::NAN);
239 let seed_sum = inputs[..p].iter().copied().sum::<f64>();
240 let seed = seed_sum / p as f64;
241 out[p - 1] = seed;
242 let mut cur = seed;
243 let (alpha, oma) = (self.alpha, self.one_minus_alpha);
244 for (slot, &x) in out[p..].iter_mut().zip(&inputs[p..]) {
245 cur = alpha.mul_add(x, oma * cur);
246 *slot = cur;
247 }
248
249 self.current = cur;
252 self.seeded = true;
253 self.warmup_sum = seed_sum;
254 self.warmup_count = p;
255 }
256
257 fn batch_fast_into(&mut self, inputs: &[f64], out: &mut [f64]) {
265 assert_eq!(
266 inputs.len(),
267 out.len(),
268 "batch output length must equal input length"
269 );
270 let p = self.period;
271 if self.seeded
272 || self.warmup_count != 0
273 || inputs.len() < p
274 || !crate::fast::in_range(inputs)
275 {
276 self.batch_nan_into(inputs, out);
277 return;
278 }
279 let (last, seed_sum) = wickra_simd::dispatch(crate::fast::EmaFast {
280 x: inputs,
281 period: p,
282 alpha: self.alpha,
283 one_minus_alpha: self.one_minus_alpha,
284 out,
285 _borrow: std::marker::PhantomData,
286 });
287 self.current = last;
288 self.seeded = true;
289 self.warmup_sum = seed_sum;
290 self.warmup_count = p;
291 }
292}
293
294#[cfg(test)]
295mod tests {
296 use super::*;
297
298 #[test]
303 fn rejects_a_period_above_the_maximum() {
304 assert!(matches!(
305 Ema::new(usize::MAX),
306 Err(Error::InvalidPeriod { .. })
307 ));
308 assert!(matches!(
309 Ema::new(crate::error::MAX_PERIOD + 1),
310 Err(Error::InvalidPeriod { .. })
311 ));
312 assert!(matches!(
313 Ema::new(1_000_000_000),
314 Err(Error::InvalidPeriod { .. })
315 ));
316 assert!(Ema::new(14).is_ok());
318 }
319 use crate::traits::BatchExt;
320 use approx::assert_relative_eq;
321
322 fn ema_naive(prices: &[f64], period: usize) -> Vec<Option<f64>> {
324 let alpha = 2.0 / (period as f64 + 1.0);
325 let mut out = Vec::with_capacity(prices.len());
326 let mut state: Option<f64> = None;
327 for (i, &p) in prices.iter().enumerate() {
328 if let Some(prev) = state {
329 let v = alpha * p + (1.0 - alpha) * prev;
330 state = Some(v);
331 out.push(Some(v));
332 } else if i + 1 == period {
333 let seed = prices[..period].iter().sum::<f64>() / period as f64;
334 state = Some(seed);
335 out.push(Some(seed));
336 } else {
337 out.push(None);
338 }
339 }
340 out
341 }
342
343 #[test]
344 fn new_rejects_zero_period() {
345 assert!(matches!(Ema::new(0), Err(Error::PeriodZero)));
346 }
347
348 #[test]
353 fn accessors_and_metadata() {
354 let ema = Ema::new(14).unwrap();
355 assert_eq!(ema.period(), 14);
356 assert_eq!(ema.warmup_period(), 14);
357 assert_eq!(ema.name(), "EMA");
358 }
359
360 #[test]
361 fn warmup_returns_none_until_seed() {
362 let mut ema = Ema::new(3).unwrap();
363 assert_eq!(ema.update(1.0), None);
364 assert_eq!(ema.update(2.0), None);
365 assert_eq!(ema.update(3.0), Some(2.0)); }
367
368 #[test]
369 fn first_value_equals_sma_seed() {
370 let mut ema = Ema::new(5).unwrap();
371 let inputs = [10.0, 20.0, 30.0, 40.0, 50.0];
372 let mut last = None;
373 for v in inputs {
374 last = ema.update(v);
375 }
376 assert_relative_eq!(last.unwrap(), 30.0, epsilon = 1e-12);
377 }
378
379 #[test]
380 fn alpha_matches_period_formula() {
381 let ema = Ema::new(10).unwrap();
382 assert_relative_eq!(ema.alpha(), 2.0 / 11.0, epsilon = 1e-15);
383 }
384
385 #[test]
386 fn step_after_seed_uses_alpha_formula() {
387 let mut ema = Ema::new(3).unwrap();
390 ema.batch(&[1.0, 2.0, 3.0]);
391 assert_relative_eq!(ema.update(10.0).unwrap(), 6.0, epsilon = 1e-12);
392 }
393
394 #[test]
395 fn constant_series_converges_to_constant() {
396 let mut ema = Ema::new(10).unwrap();
397 let out = ema.batch(&[42.0_f64; 100]);
398 for x in out.iter().skip(9) {
399 assert_relative_eq!(x.unwrap(), 42.0, epsilon = 1e-9);
400 }
401 }
402
403 #[test]
404 fn with_alpha_validates_range() {
405 assert!(Ema::with_alpha(0.5).is_ok());
406 assert!(Ema::with_alpha(1.0).is_ok());
407 assert!(matches!(
408 Ema::with_alpha(0.0),
409 Err(Error::InvalidPeriod { .. })
410 ));
411 assert!(matches!(
412 Ema::with_alpha(1.5),
413 Err(Error::InvalidPeriod { .. })
414 ));
415 assert!(matches!(
416 Ema::with_alpha(f64::NAN),
417 Err(Error::InvalidPeriod { .. })
418 ));
419 }
420
421 #[test]
422 fn reset_clears_state() {
423 let mut ema = Ema::new(3).unwrap();
424 ema.batch(&[1.0, 2.0, 3.0]);
425 assert!(ema.is_ready());
426 ema.reset();
427 assert!(!ema.is_ready());
428 assert_eq!(ema.update(1.0), None);
429 }
430
431 #[test]
432 fn batch_equals_streaming() {
433 let prices: Vec<f64> = (1..=30).map(f64::from).collect();
434 let mut a = Ema::new(5).unwrap();
435 let mut b = Ema::new(5).unwrap();
436 assert_eq!(
437 a.batch(&prices),
438 prices.iter().map(|p| b.update(*p)).collect::<Vec<_>>()
439 );
440 }
441
442 #[test]
443 fn ignores_non_finite_input() {
444 let mut ema = Ema::new(3).unwrap();
445 ema.batch(&[1.0, 2.0, 3.0]);
446 let before = ema.value();
447 assert_eq!(ema.update(f64::NAN), None);
448 assert_eq!(ema.update(f64::INFINITY), None);
449 assert_eq!(ema.value(), before);
451 }
452
453 fn bits_eq(a: &[f64], b: &[f64]) -> bool {
454 a.len() == b.len()
455 && a.iter()
456 .zip(b)
457 .all(|(x, y)| x == y || (x.is_nan() && y.is_nan()))
458 }
459
460 fn ema_replay(period: usize, series: &[f64]) -> Vec<f64> {
461 let mut e = Ema::new(period).unwrap();
462 series
463 .iter()
464 .map(|&x| e.update(x).unwrap_or(f64::NAN))
465 .collect()
466 }
467
468 #[test]
469 fn batch_nan_fast_path_is_bit_identical() {
470 let series: Vec<f64> = (0..300)
471 .map(|i| (f64::from(i) * 0.25).cos() * 8.0 + 40.0)
472 .collect();
473 let mut ema = Ema::new(14).unwrap();
474 let got = ema.batch_nan(&series);
475 assert!(bits_eq(&got, &ema_replay(14, &series)));
476 let mut ref_ema = Ema::new(14).unwrap();
477 for &x in &series {
478 ref_ema.update(x);
479 }
480 assert_eq!(ema.update(7.5), ref_ema.update(7.5));
481 }
482
483 #[test]
487 fn seed_matches_a_buffered_sum_bit_for_bit() {
488 let windows: [&[f64]; 3] = [
489 &[-0.0, -0.0, -0.0],
490 &[1e16, 1.0, -1e16, 3.5, 0.1],
491 &[0.3, 0.1, 0.7, 0.2, 0.9, 0.4, 0.6],
492 ];
493 for window in windows {
494 let buffered = window.iter().copied().sum::<f64>() / window.len() as f64;
495 let mut ema = Ema::new(window.len()).unwrap();
496 let seed = window.iter().filter_map(|&x| ema.update(x)).last().unwrap();
497 assert_eq!(seed.to_bits(), buffered.to_bits());
498 let mut out = vec![0.0; window.len()];
499 Ema::new(window.len())
500 .unwrap()
501 .batch_nan_into(window, &mut out);
502 assert_eq!(out[window.len() - 1].to_bits(), buffered.to_bits());
503 }
504 }
505
506 #[test]
509 fn batch_nan_into_overwrites_a_dirty_buffer() {
510 let series: Vec<f64> = (0..120).map(|i| f64::from(i % 11) * 0.75 + 20.0).collect();
511 let mut out = vec![-1.0; series.len()];
512 Ema::new(10).unwrap().batch_nan_into(&series, &mut out);
513 assert!(bits_eq(&out, &ema_replay(10, &series)));
514 let mut short = [-1.0; 4];
515 Ema::new(10)
516 .unwrap()
517 .batch_nan_into(&series[..4], &mut short);
518 assert!(short.iter().all(|x| x.is_nan()));
519 }
520
521 #[test]
522 fn batch_nan_falls_back_on_non_finite() {
523 let series = [1.0, 2.0, 3.0, f64::INFINITY, 5.0, 6.0, 7.0];
524 let mut ema = Ema::new(3).unwrap();
525 assert!(bits_eq(&ema.batch_nan(&series), &ema_replay(3, &series)));
526 }
527
528 #[test]
529 fn batch_nan_falls_back_when_warming() {
530 let mut ema = Ema::new(3).unwrap();
531 ema.update(10.0); let series = [1.0, 2.0, 3.0, 4.0];
533 let mut ref_ema = Ema::new(3).unwrap();
534 ref_ema.update(10.0);
535 let want: Vec<f64> = series
536 .iter()
537 .map(|&x| ref_ema.update(x).unwrap_or(f64::NAN))
538 .collect();
539 assert!(bits_eq(&ema.batch_nan(&series), &want));
540 }
541
542 #[test]
543 fn batch_nan_sub_period_slice_stays_unseeded() {
544 let series = [1.0, 2.0];
545 let mut ema = Ema::new(5).unwrap();
546 let got = ema.batch_nan(&series);
547 assert!(got.iter().all(|x| x.is_nan()) && got.len() == 2);
548 assert!(!ema.is_ready());
549 assert!(bits_eq(
551 &[ema.update(3.0).unwrap_or(f64::NAN)],
552 &[ema_replay(5, &[1.0, 2.0, 3.0])[2]]
553 ));
554 }
555
556 proptest::proptest! {
557 #![proptest_config(proptest::test_runner::Config::with_cases(48))]
558 #[test]
559 fn ema_matches_naive(
560 period in 1usize..20,
561 prices in proptest::collection::vec(-1000.0_f64..1000.0, 0..150),
562 ) {
563 let mut ema = Ema::new(period).unwrap();
564 let got = ema.batch(&prices);
565 let want = ema_naive(&prices, period);
566 proptest::prop_assert_eq!(got.len(), want.len());
567 for (g, w) in got.iter().zip(want.iter()) {
568 match (g, w) {
569 (None, None) => {}
570 (Some(a), Some(b)) => proptest::prop_assert!(
571 (a - b).abs() <= 1e-9 * a.abs().max(1.0),
572 "got={a} want={b}"
573 ),
574 _ => proptest::prop_assert!(false, "warmup mismatch"),
575 }
576 }
577 }
578 }
579}