Skip to main content

wickra_core/indicators/
wma.rs

1//! Weighted Moving Average (linear weights).
2
3use crate::error::{Error, Result};
4use crate::traits::Indicator;
5
6/// Weighted Moving Average with linear weights `1, 2, ..., period`.
7///
8/// Output is `sum(weight_i * price_i) / sum(weights)`. Maintained incrementally in
9/// O(1) by keeping the rolling sum of values and the rolling weighted sum.
10///
11/// Both running sums are recomputed exactly from the live window every
12/// `16 · period` steady-state updates, the same bound the SMA puts on drift.
13/// Without it the weighted sum — updated as `W − S + period · x`, a difference of
14/// quantities `period` times larger than the result — lost about 6e-10 of
15/// relative accuracy over 500 000 updates of a price series.
16///
17/// # Example
18///
19/// ```
20/// use wickra_core::{Indicator, Wma};
21///
22/// let mut indicator = Wma::new(3).unwrap();
23/// let mut last = None;
24/// for i in 0..80 {
25///     last = indicator.update(100.0 + f64::from(i));
26/// }
27/// assert!(last.is_some());
28/// ```
29#[derive(Debug, Clone)]
30pub struct Wma {
31    period: usize,
32    /// Ring buffer of the last `period` finite inputs; `head` is the next slot
33    /// to write and, once full, the oldest value. A flat buffer with a write
34    /// cursor instead of a `VecDeque`: no per-update bookkeeping on the hot path.
35    buf: Box<[f64]>,
36    head: usize,
37    /// Slots filled, saturating at `period`.
38    count: usize,
39    weight_sum: f64, // sum_i (weight_i * value_i)
40    value_sum: f64,  // sum_i (value_i)
41    weights_total: f64,
42    /// Steady-state laps of the ring (`period` updates each) since the running
43    /// sums were last recomputed. Counted where the write cursor wraps, so the
44    /// per-update path carries no bookkeeping of its own; steady state starts
45    /// with the cursor at 0, so a lap count lands on the same update an update
46    /// count would.
47    laps_since_reseed: usize,
48}
49
50/// Recompute the running sums every `RESEED_EVERY * period` steady-state
51/// updates (the SMA's cadence), i.e. every `RESEED_EVERY` laps of the ring.
52const RESEED_EVERY: usize = 16;
53
54impl Wma {
55    /// Construct a new WMA with the given window length.
56    ///
57    /// # Errors
58    ///
59    /// Returns [`Error::PeriodZero`] if `period == 0`.
60    pub fn new(period: usize) -> Result<Self> {
61        if period == 0 {
62            return Err(Error::PeriodZero);
63        }
64        if period > crate::error::MAX_PERIOD {
65            return Err(Error::InvalidPeriod {
66                message: crate::error::PERIOD_ABOVE_MAX,
67            });
68        }
69        let n = period as f64;
70        let weights_total = n * (n + 1.0) / 2.0;
71        Ok(Self {
72            period,
73            buf: vec![0.0; period].into_boxed_slice(),
74            head: 0,
75            count: 0,
76            weight_sum: 0.0,
77            value_sum: 0.0,
78            weights_total,
79            laps_since_reseed: 0,
80        })
81    }
82
83    /// Configured period.
84    pub const fn period(&self) -> usize {
85        self.period
86    }
87
88    /// Whether the WMA has taken no input since construction or reset.
89    pub(crate) fn is_empty(&self) -> bool {
90        self.count == 0
91    }
92
93    /// Current value if available.
94    pub fn value(&self) -> Option<f64> {
95        if self.count == self.period {
96            Some(self.weight_sum / self.weights_total)
97        } else {
98            None
99        }
100    }
101
102    /// The weighted sum `Σ (k + 1) · x_k` over the window in chronological
103    /// order (oldest first, weight 1).
104    /// The steady state of a full window as a [`Steady`] run, for a batch loop.
105    ///
106    /// # Panics
107    ///
108    /// Panics unless the window is full.
109    pub(crate) fn steady(&mut self) -> Steady<'_> {
110        assert!(self.is_ready(), "a steady run needs a full window");
111        Steady {
112            weight_sum: self.weight_sum,
113            value_sum: self.value_sum,
114            head: self.head,
115            laps: self.laps_since_reseed,
116            period_f: self.period as f64,
117            total: self.weights_total,
118            wma: self,
119        }
120    }
121
122    fn weighted_window_sum(&self) -> f64 {
123        self.buf[self.head..]
124            .iter()
125            .chain(&self.buf[..self.head])
126            .enumerate()
127            .map(|(i, v)| (i as f64 + 1.0) * v)
128            .sum()
129    }
130}
131
132/// A full-window [`Wma`] advanced input by input with its sums, cursor and
133/// reseed count in locals: `update`'s steady-state operations in `update`'s
134/// order -- the same bits and reseed cadence -- held in registers rather than
135/// written back through the indicator for every input, which made a replay of
136/// `update` three times slower than streaming at long periods. The state goes
137/// back into the indicator when the run is dropped.
138pub(crate) struct Steady<'a> {
139    wma: &'a mut Wma,
140    weight_sum: f64,
141    value_sum: f64,
142    head: usize,
143    laps: usize,
144    period_f: f64,
145    total: f64,
146}
147
148impl Steady<'_> {
149    /// `update(x)` for a finite `x`: the new average.
150    // Inlined into the batch loop, where the locals become registers.
151    #[allow(clippy::inline_always)]
152    #[inline(always)]
153    pub(crate) fn step(&mut self, x: f64) -> f64 {
154        let buf = &mut self.wma.buf;
155        let oldest = std::mem::replace(&mut buf[self.head], x);
156        self.weight_sum = self.weight_sum - self.value_sum + self.period_f * x;
157        self.value_sum = self.value_sum - oldest + x;
158        self.head += 1;
159        if self.head == buf.len() {
160            self.head = 0;
161            self.laps += 1;
162            if self.laps == RESEED_EVERY {
163                // The cursor just wrapped, so the window runs oldest-first from 0.
164                self.value_sum = buf.iter().sum();
165                self.weight_sum = buf
166                    .iter()
167                    .enumerate()
168                    .map(|(i, v)| (i as f64 + 1.0) * v)
169                    .sum();
170                self.laps = 0;
171            }
172        }
173        self.weight_sum / self.total
174    }
175}
176
177impl Drop for Steady<'_> {
178    fn drop(&mut self) {
179        self.wma.weight_sum = self.weight_sum;
180        self.wma.value_sum = self.value_sum;
181        self.wma.head = self.head;
182        self.wma.laps_since_reseed = self.laps;
183    }
184}
185
186impl Indicator for Wma {
187    type Input = f64;
188    type Output = f64;
189
190    #[inline]
191    fn update(&mut self, input: f64) -> Option<f64> {
192        if !input.is_finite() {
193            return None;
194        }
195        if self.count < self.period {
196            // Warmup. Just accumulate; compute weight_sum once when the window first
197            // becomes full to avoid having to track changing weights during warmup.
198            self.buf[self.head] = input;
199            self.head += 1;
200            if self.head == self.period {
201                self.head = 0;
202            }
203            self.value_sum += input;
204            self.count += 1;
205            if self.count == self.period {
206                self.weight_sum = self.weighted_window_sum();
207            }
208            return self.value();
209        }
210        // Steady state: slide the window. With weights [1, 2, ..., period],
211        //   new_weight_sum = old_weight_sum - old_value_sum + period * new_input
212        // because every retained element's weight drops by one and the newcomer
213        // enters at weight = period. Order matters: subtract `value_sum` BEFORE
214        // updating it.
215        // One indexed slot for both the read of the oldest value and the write.
216        let slot = &mut self.buf[self.head];
217        let oldest = std::mem::replace(slot, input);
218        self.weight_sum = self.weight_sum - self.value_sum + self.period as f64 * input;
219        self.value_sum = self.value_sum - oldest + input;
220        self.head += 1;
221        if self.head == self.period {
222            self.head = 0;
223            self.laps_since_reseed += 1;
224            if self.laps_since_reseed == RESEED_EVERY {
225                // The cursor just wrapped, so the window runs oldest-first from 0.
226                self.value_sum = self.buf.iter().sum();
227                self.weight_sum = self.weighted_window_sum();
228                self.laps_since_reseed = 0;
229            }
230        }
231        self.value()
232    }
233
234    /// The exact batch with the window sums, cursor and reseed count in
235    /// locals: `update`'s operations in `update`'s order -- the same bits,
236    /// the same reseed cadence -- held in registers rather than written back
237    /// through `self` for every input, which made the replay three times slower
238    /// than streaming at long periods. The warmup still goes through `update`.
239    fn batch_nan_into(&mut self, inputs: &[f64], out: &mut [f64]) {
240        assert_eq!(
241            inputs.len(),
242            out.len(),
243            "batch output length must equal input length"
244        );
245        let mut start = 0;
246        while !self.is_ready() && start < inputs.len() {
247            out[start] = self.update(inputs[start]).unwrap_or(f64::NAN);
248            start += 1;
249        }
250        if start == inputs.len() {
251            return;
252        }
253        let mut run = self.steady();
254        for (slot, &x) in out[start..].iter_mut().zip(&inputs[start..]) {
255            *slot = if x.is_finite() { run.step(x) } else { f64::NAN };
256        }
257    }
258
259    fn reset(&mut self) {
260        self.head = 0;
261        self.count = 0;
262        self.weight_sum = 0.0;
263        self.value_sum = 0.0;
264        self.laps_since_reseed = 0;
265    }
266
267    #[inline]
268    fn warmup_period(&self) -> usize {
269        self.period
270    }
271
272    #[inline]
273    fn is_ready(&self) -> bool {
274        self.count == self.period
275    }
276
277    #[inline]
278    fn name(&self) -> &'static str {
279        "WMA"
280    }
281
282    /// SIMD kernel: rolling sums and the weighted-numerator recurrence
283    /// `N = N − S + period · x` as prefix scans, re-anchored on exact window
284    /// values every `16 · period` inputs, scaled by `1 / Σ weights`. Agrees with
285    /// the exact batch to within a few units in the last place; warmup `NaN`s
286    /// and length are identical. Afterwards the window is rebuilt exactly from
287    /// the last `period` inputs.
288    fn batch_fast_into(&mut self, inputs: &[f64], out: &mut [f64]) {
289        assert_eq!(
290            inputs.len(),
291            out.len(),
292            "batch output length must equal input length"
293        );
294        let p = self.period;
295        if self.count != 0 || inputs.len() < p || !crate::fast::in_range(inputs) {
296            self.batch_nan_into(inputs, out);
297            return;
298        }
299        wickra_simd::dispatch(crate::fast::WmaFast {
300            x: inputs,
301            period: p,
302            out,
303            _borrow: std::marker::PhantomData,
304        });
305        crate::fast::replay_tail(self, &inputs[inputs.len() - p..]);
306    }
307}
308
309#[cfg(test)]
310mod tests {
311    use super::*;
312    use crate::traits::BatchExt;
313    use approx::assert_relative_eq;
314
315    /// Over many reseed intervals the WMA must stay at its definition
316    /// (`Σ (k + 1) · x_k / Σ weights` over the live window) to a few ulps; the
317    /// incremental weighted sum alone drifted to ~1e-10 on such a series.
318    #[test]
319    fn long_stream_drift_stays_bounded() {
320        let period = 14;
321        let mut wma = Wma::new(period).unwrap();
322        let xs: Vec<f64> = (0..40_000)
323            .map(|i| {
324                let t = f64::from(i);
325                100.0 + (t * 0.0137).sin() * 5.0 + (t * 0.37).cos() + (t * 0.0011).sin() * 20.0
326            })
327            .collect();
328        let total = (period * (period + 1) / 2) as f64;
329        for (i, &x) in xs.iter().enumerate() {
330            let got = wma.update(x);
331            if i + 1 >= period && i % 509 == 0 {
332                let def: f64 = xs[i + 1 - period..=i]
333                    .iter()
334                    .enumerate()
335                    .map(|(k, v)| (k as f64 + 1.0) * v)
336                    .sum::<f64>()
337                    / total;
338                let got = got.unwrap();
339                assert!(((got - def) / def).abs() < 1e-13, "at {i}: {got} vs {def}");
340            }
341        }
342    }
343
344    /// Until the first reseed the ring buffer computes exactly what the old
345    /// sliding sums did: warmup accumulation, then `W − S + p·x`.
346    #[test]
347    fn matches_the_incremental_form_before_the_first_reseed() {
348        let period = 5;
349        let xs: Vec<f64> = (0..70).map(|i| f64::from(i % 7) * 1.5 + 3.25).collect();
350        let mut wma = Wma::new(period).unwrap();
351        let (mut value_sum, mut weight_sum) = (0.0_f64, 0.0_f64);
352        for (i, &x) in xs.iter().enumerate() {
353            let got = wma.update(x);
354            if i < period {
355                value_sum += x;
356                if i + 1 == period {
357                    weight_sum = xs[..period]
358                        .iter()
359                        .enumerate()
360                        .map(|(k, v)| (k as f64 + 1.0) * v)
361                        .sum();
362                }
363            } else {
364                weight_sum = weight_sum - value_sum + period as f64 * x;
365                value_sum = value_sum - xs[i - period] + x;
366            }
367            if i + 1 >= period {
368                assert_eq!(
369                    got.unwrap().to_bits(),
370                    (weight_sum / 15.0).to_bits(),
371                    "at {i}"
372                );
373            }
374        }
375    }
376
377    /// Reference implementation: explicit weighted average over a window.
378    fn wma_naive(prices: &[f64], period: usize) -> Vec<Option<f64>> {
379        let weights_total = (period as f64) * (period as f64 + 1.0) / 2.0;
380        prices
381            .iter()
382            .enumerate()
383            .map(|(i, _)| {
384                if i + 1 < period {
385                    None
386                } else {
387                    let window = &prices[i + 1 - period..=i];
388                    let s: f64 = window
389                        .iter()
390                        .enumerate()
391                        .map(|(j, p)| (j as f64 + 1.0) * p)
392                        .sum();
393                    Some(s / weights_total)
394                }
395            })
396            .collect()
397    }
398
399    #[test]
400    fn new_rejects_zero_period() {
401        assert!(matches!(Wma::new(0), Err(Error::PeriodZero)));
402    }
403
404    /// Cover the const accessor `period` (56-58) and the Indicator-impl
405    /// `warmup_period` (111-113) + `name` (119-121). Existing tests never
406    /// inspect these metadata methods.
407    #[test]
408    fn accessors_and_metadata() {
409        let wma = Wma::new(7).unwrap();
410        assert_eq!(wma.period(), 7);
411        assert_eq!(wma.warmup_period(), 7);
412        assert_eq!(wma.name(), "WMA");
413    }
414
415    #[test]
416    fn warmup_returns_none() {
417        let mut wma = Wma::new(3).unwrap();
418        assert_eq!(wma.update(1.0), None);
419        assert_eq!(wma.update(2.0), None);
420        // WMA(3) of [1,2,3]: oldest = 1 (weight 1), middle = 2 (weight 2), newest = 3 (weight 3)
421        // -> (1*1 + 2*2 + 3*3) / (1+2+3) = 14/6
422        assert_relative_eq!(wma.update(3.0).unwrap(), 14.0 / 6.0, epsilon = 1e-12);
423    }
424
425    #[test]
426    fn known_values_period_4() {
427        // WMA(4) weights 1,2,3,4 (total 10); inputs [1,2,3,4]:
428        // (1*1 + 2*2 + 3*3 + 4*4) / 10 = (1+4+9+16)/10 = 30/10 = 3.0
429        let mut wma = Wma::new(4).unwrap();
430        let v = wma.batch(&[1.0, 2.0, 3.0, 4.0]);
431        assert_relative_eq!(v[3].unwrap(), 3.0, epsilon = 1e-12);
432    }
433
434    #[test]
435    fn matches_naive_over_random_inputs() {
436        let prices: Vec<f64> = (1..=30).map(|i| f64::from(i) * 1.7 - 5.0).collect();
437        let mut wma = Wma::new(7).unwrap();
438        let got = wma.batch(&prices);
439        let want = wma_naive(&prices, 7);
440        for (i, (g, w)) in got.iter().zip(want.iter()).enumerate() {
441            // Same warmup — emission shape must agree at every index.
442            assert_eq!(g.is_some(), w.is_some(), "warmup mismatch at index {i}");
443            if let (Some(a), Some(b)) = (g, w) {
444                assert_relative_eq!(*a, *b, epsilon = 1e-9);
445            }
446        }
447    }
448
449    #[test]
450    fn period_one_is_pass_through() {
451        let mut wma = Wma::new(1).unwrap();
452        assert_relative_eq!(wma.update(5.5).unwrap(), 5.5, epsilon = 1e-12);
453        assert_relative_eq!(wma.update(7.5).unwrap(), 7.5, epsilon = 1e-12);
454    }
455
456    #[test]
457    fn reset_clears_state() {
458        let mut wma = Wma::new(4).unwrap();
459        wma.batch(&[1.0, 2.0, 3.0, 4.0, 5.0]);
460        assert!(wma.is_ready());
461        wma.reset();
462        assert!(!wma.is_ready());
463        assert_eq!(wma.update(10.0), None);
464    }
465
466    #[test]
467    fn batch_equals_streaming() {
468        let prices: Vec<f64> = (1..=20).map(|i| f64::from(i) * 0.5).collect();
469        let mut a = Wma::new(5).unwrap();
470        let mut b = Wma::new(5).unwrap();
471        assert_eq!(
472            a.batch(&prices),
473            prices.iter().map(|p| b.update(*p)).collect::<Vec<_>>()
474        );
475    }
476
477    /// The fused exact batch is the `update` replay bit for bit: across several
478    /// reseeds, with non-finite inputs in the warmup and in the steady state,
479    /// split into two calls at every point, and at period 1.
480    #[test]
481    fn batch_nan_into_is_the_update_replay_bit_for_bit() {
482        let mut series: Vec<f64> = (0..400)
483            .map(|i| 100.0 + (f64::from(i) * 0.37).sin() * 7.0 + f64::from(i % 11) * 0.3)
484            .collect();
485        series[3] = f64::NAN;
486        series[150] = f64::INFINITY;
487        series[151] = f64::NEG_INFINITY;
488        series[222] = f64::NAN;
489        let bits = |v: &[f64]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
490        for period in [1, 2, 7] {
491            let mut replay = Wma::new(period).unwrap();
492            let want: Vec<f64> = series
493                .iter()
494                .map(|&x| replay.update(x).unwrap_or(f64::NAN))
495                .collect();
496            for split in (0..series.len()).step_by(13).chain([series.len()]) {
497                let mut wma = Wma::new(period).unwrap();
498                let mut got = vec![0.0; series.len()];
499                let (head, tail) = got.split_at_mut(split);
500                wma.batch_nan_into(&series[..split], head);
501                wma.batch_nan_into(&series[split..], tail);
502                assert_eq!(bits(&got), bits(&want), "period {period} split {split}");
503                // And the state carries on as the replay's does.
504                assert_eq!(wma.update(101.5), replay.clone().update(101.5));
505            }
506        }
507    }
508
509    #[test]
510    fn ignores_non_finite_input_but_keeps_state() {
511        let mut wma = Wma::new(3).unwrap();
512        wma.update(1.0);
513        wma.update(2.0);
514        wma.update(3.0).expect("WMA(3) ready after three inputs");
515        // Non-finite inputs return the last value without mutating the window.
516        assert_eq!(wma.update(f64::NAN), None);
517        assert_eq!(wma.update(f64::INFINITY), None);
518        // The window still holds 1, 2, 3 -> next real input slides it to 2, 3, 4.
519        assert_relative_eq!(
520            wma.update(4.0).unwrap(),
521            (2.0 * 1.0 + 3.0 * 2.0 + 4.0 * 3.0) / 6.0,
522            epsilon = 1e-12
523        );
524    }
525
526    proptest::proptest! {
527        #![proptest_config(proptest::test_runner::Config::with_cases(48))]
528        #[test]
529        fn proptest_matches_naive(
530            period in 1usize..15,
531            prices in proptest::collection::vec(-500.0_f64..500.0, 0..120),
532        ) {
533            let mut wma = Wma::new(period).unwrap();
534            let got = wma.batch(&prices);
535            let want = wma_naive(&prices, period);
536            proptest::prop_assert_eq!(got.len(), want.len());
537            for (g, w) in got.iter().zip(want.iter()) {
538                match (g, w) {
539                    (None, None) => {}
540                    (Some(a), Some(b)) => proptest::prop_assert!(
541                        (a - b).abs() < 1e-7,
542                        "got={a} want={b}"
543                    ),
544                    _ => proptest::prop_assert!(false, "warmup mismatch"),
545                }
546            }
547        }
548    }
549}