1use crate::error::{Error, Result};
4use crate::traits::Indicator;
5
6#[derive(Debug, Clone)]
30pub struct Wma {
31 period: usize,
32 buf: Box<[f64]>,
36 head: usize,
37 count: usize,
39 weight_sum: f64, value_sum: f64, weights_total: f64,
42 laps_since_reseed: usize,
48}
49
50const RESEED_EVERY: usize = 16;
53
54impl Wma {
55 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 pub const fn period(&self) -> usize {
85 self.period
86 }
87
88 pub(crate) fn is_empty(&self) -> bool {
90 self.count == 0
91 }
92
93 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 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
132pub(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 #[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 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 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 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 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 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 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 #[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 #[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 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 #[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 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 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 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 #[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 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 assert_eq!(wma.update(f64::NAN), None);
517 assert_eq!(wma.update(f64::INFINITY), None);
518 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}