Skip to main content

wickra_core/indicators/
ema.rs

1//! Exponential Moving Average.
2
3use crate::error::{Error, Result};
4use crate::traits::Indicator;
5
6/// Exponential Moving Average with smoothing factor `alpha = 2 / (period + 1)`.
7///
8/// The first value is seeded with the simple mean of the first `period` inputs
9/// (the classical TA-Lib convention). From then on each new input contributes
10/// `alpha * input + (1 - alpha) * previous`.
11///
12/// # Example
13///
14/// ```
15/// use wickra_core::{Indicator, Ema};
16///
17/// let mut indicator = Ema::new(3).unwrap();
18/// let mut last = None;
19/// for i in 0..80 {
20///     last = indicator.update(100.0 + f64::from(i));
21/// }
22/// assert!(last.is_some());
23/// ```
24#[derive(Debug, Clone)]
25pub struct Ema {
26    period: usize,
27    alpha: f64,
28    /// `1 - alpha`, precomputed so the recurrence avoids a subtraction per tick.
29    /// Cached value, so the steady-state output is bit-for-bit unchanged.
30    one_minus_alpha: f64,
31    /// Latest EMA value, valid only once `seeded` is true. Stored as a bare `f64`
32    /// (plus the `seeded` flag) rather than `Option<f64>` so the steady-state
33    /// recurrence reads and writes 8 bytes with no enum-tag handling per tick.
34    current: f64,
35    /// Whether `current` holds a real value yet (warmup complete).
36    seeded: bool,
37    /// Running sum of the warmup inputs, the numerator of the seed mean. It
38    /// starts at `-0.0`, the neutral element `f64::sum` folds from, and adds in
39    /// input order, so the seed is bit-for-bit the mean a buffered
40    /// `iter().sum()` would give — without a heap buffer that would keep the
41    /// state out of registers on the per-tick path.
42    warmup_sum: f64,
43    /// Number of warmup inputs taken so far (saturates at `period` on seeding).
44    warmup_count: usize,
45}
46
47impl Ema {
48    /// Construct an EMA with the given period.
49    ///
50    /// # Errors
51    ///
52    /// Returns [`Error::PeriodZero`] if `period == 0`.
53    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    /// Construct an EMA with a custom smoothing factor `alpha in (0, 1]`.
75    ///
76    /// The reported `period` is derived from `alpha` via `2/alpha - 1` and rounded;
77    /// `warmup_period()` falls back to `1` because the implementation seeds from the
78    /// very first input.
79    ///
80    /// # Errors
81    ///
82    /// Returns [`Error::InvalidPeriod`] if `alpha` is not in `(0.0, 1.0]` or non-finite.
83    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    /// Configured period.
101    pub const fn period(&self) -> usize {
102        self.period
103    }
104
105    /// Smoothing factor.
106    pub const fn alpha(&self) -> f64 {
107        self.alpha
108    }
109
110    /// The cached `1 - alpha` the recurrence multiplies the previous value by.
111    pub(crate) const fn one_minus_alpha(&self) -> f64 {
112        self.one_minus_alpha
113    }
114
115    /// Current value if available.
116    pub const fn value(&self) -> Option<f64> {
117        if self.seeded {
118            Some(self.current)
119        } else {
120            None
121        }
122    }
123
124    /// Whether the EMA has seen no input yet (neither seeded nor mid-warmup).
125    /// Lets composite indicators (e.g. MACD) decide if a fast batch path is safe.
126    pub(crate) fn is_fresh(&self) -> bool {
127        !self.seeded && self.warmup_count == 0
128    }
129
130    /// Force the EMA into its seeded steady state with `current` as the latest
131    /// value. Used by composite fused batch paths (MACD) to leave each sub-EMA
132    /// where a per-tick `update` replay would, so a later `update` continues
133    /// correctly. The post-seed recurrence never re-reads the warmup sum, so it
134    /// is left as-is.
135    pub(crate) fn seed_to(&mut self, current: f64) {
136        self.current = current;
137        self.seeded = true;
138    }
139
140    /// Vectorized batch returning one `f64` per input (`NaN` during warmup).
141    ///
142    /// Kept as an inherent method so existing callers need no trait import; it
143    /// allocates the result and fills it through
144    /// [`batch_nan_into`](Indicator::batch_nan_into), which carries the fast path.
145    pub fn batch_nan(&mut self, inputs: &[f64]) -> Vec<f64> {
146        crate::traits::BatchNanExt::batch_nan(self, inputs)
147    }
148
149    /// Internal helper that feeds a value without finiteness validation. The caller
150    /// guarantees `input.is_finite()`. Used by MACD which has already validated.
151    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    /// For a fresh indicator over an all-finite slice this runs the seed (mean
206    /// of the first `period`) once and then the bare
207    /// `alpha * x + (1 - alpha) * prev` recurrence in a tight loop with no
208    /// per-element `is_finite`/`seeded` branch and no `Option` — yet uses the
209    /// identical `mul_add`, so every value is *bit-for-bit* equal to replaying
210    /// `update`. Any other state, or a non-finite element, defers to the exact
211    /// `update` replay.
212    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            // Not enough to seed; mirror `update` accumulating the warmup.
229            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        // Warmup `[0, p-1)` is `NaN`; the seed lands on `p-1`, the recurrence after.
238        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        // Leave state exactly where `update` would: seeded on `current`, the
250        // warmup sum and count covering the first `period` inputs.
251        self.current = cur;
252        self.seeded = true;
253        self.warmup_sum = seed_sum;
254        self.warmup_count = p;
255    }
256
257    /// SIMD kernel: after the same seed (the mean of the first `period`
258    /// inputs) the recurrence `y = alpha * x + (1 - alpha) * y` runs as a
259    /// linear-recurrence scan, eight values per step. Agrees with the exact
260    /// batch to within a few units in the last place (the scan reassociates
261    /// the recurrence through powers of `1 - alpha`); warmup `NaN`s and length
262    /// are identical. Afterwards the EMA continues streaming from the kernel's
263    /// last value.
264    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    /// Constructors size their buffers from the period, so an absurd value used
299    /// to abort the process before the caller saw anything: `usize::MAX` hit a
300    /// capacity overflow inside `Vec`, and a billion quietly reserved eight
301    /// gigabytes. Both are now rejected as an ordinary error.
302    #[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        // A sane period is untouched.
317        assert!(Ema::new(14).is_ok());
318    }
319    use crate::traits::BatchExt;
320    use approx::assert_relative_eq;
321
322    /// Independent reference: SMA-seeded EMA computed straight from the definition.
323    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    /// Cover the const accessor `period` (74-77) and the Indicator-impl
349    /// `warmup_period` (123-125) + `name` (131-133). `alpha` and `value`
350    /// are exercised by other tests and downstream consumers; only the
351    /// three metadata methods were dead.
352    #[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)); // seed = SMA([1,2,3]) = 2
366    }
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        // period=3 => alpha = 0.5; seed = mean([1,2,3]) = 2; next input 10
388        // expected = 0.5*10 + 0.5*2 = 6
389        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        // The rejected input must not have disturbed the state.
450        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    /// The seed comes from a running sum rather than a buffered
484    /// `iter().sum()`; the two must agree to the bit on every toolchain the
485    /// crate supports, the sign of an all-negative-zero window included.
486    #[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    /// Into a caller buffer that already holds values, both the seeded path and
507    /// the sub-period path must write every cell exactly as the replay would.
508    #[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); // mid-warmup: one input taken, not seeded
532        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        // Warmup state was stashed: feeding the rest seeds exactly as a full stream.
550        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}