1use crate::drift::detector::DriftLevel;
17use crate::error::{RillError, checked_finite_add, ensure_finite};
18
19#[derive(Debug, Clone)]
47#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
48pub struct TimeDecayedMean {
49 decay: f64,
50 weighted_sum: f64,
51 weight_total: f64,
52 last_time: Option<f64>,
53}
54
55impl TimeDecayedMean {
56 pub fn new(decay: f64) -> Result<Self, RillError> {
61 ensure_finite("decay", decay)?;
62 if decay <= 0.0 {
63 return Err(RillError::InvalidParameter {
64 name: "decay",
65 value: decay,
66 });
67 }
68 Ok(Self {
69 decay,
70 weighted_sum: 0.0,
71 weight_total: 0.0,
72 last_time: None,
73 })
74 }
75
76 pub const fn decay(&self) -> f64 {
78 self.decay
79 }
80
81 pub fn update(&mut self, time: f64, value: f64) -> Result<(), RillError> {
86 ensure_finite("time", time)?;
87 ensure_finite("value", value)?;
88 match self.last_time {
89 None => {
90 self.weighted_sum = value;
92 self.weight_total = 1.0;
93 }
94 Some(prev) => {
95 if time < prev {
96 return Err(RillError::InvalidParameter {
97 name: "time",
98 value: time,
99 });
100 }
101 let dt = time - prev;
102 let factor = (-self.decay * dt).exp();
103 self.weighted_sum =
104 checked_finite_add(factor * self.weighted_sum, value, "weighted_sum")?;
105 self.weight_total =
106 checked_finite_add(factor * self.weight_total, 1.0, "weight_total")?;
107 }
108 }
109 self.last_time = Some(time);
110 Ok(())
111 }
112
113 pub fn value(&self) -> Option<f64> {
115 if self.weight_total > 0.0 {
116 Some(self.weighted_sum / self.weight_total)
117 } else {
118 None
119 }
120 }
121
122 pub const fn weight_total(&self) -> f64 {
124 self.weight_total
125 }
126
127 pub const fn last_time(&self) -> Option<f64> {
129 self.last_time
130 }
131
132 pub fn reset(&mut self) {
134 self.weighted_sum = 0.0;
135 self.weight_total = 0.0;
136 self.last_time = None;
137 }
138}
139
140#[derive(Debug, Clone)]
167#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
168pub struct LearningRateScheduler {
169 base_lr: f64,
170 warning_multiplier: f64,
171 drift_multiplier: f64,
172 current_state: DriftLevel,
173}
174
175impl LearningRateScheduler {
176 pub fn new(
183 base_lr: f64,
184 warning_multiplier: f64,
185 drift_multiplier: f64,
186 ) -> Result<Self, RillError> {
187 ensure_finite("base_lr", base_lr)?;
188 ensure_finite("warning_multiplier", warning_multiplier)?;
189 ensure_finite("drift_multiplier", drift_multiplier)?;
190 if base_lr <= 0.0 {
191 return Err(RillError::InvalidLearningRate(base_lr));
192 }
193 if warning_multiplier < 1.0 {
194 return Err(RillError::InvalidParameter {
195 name: "warning_multiplier",
196 value: warning_multiplier,
197 });
198 }
199 if drift_multiplier < warning_multiplier {
200 return Err(RillError::InvalidParameter {
201 name: "drift_multiplier",
202 value: drift_multiplier,
203 });
204 }
205 Ok(Self {
206 base_lr,
207 warning_multiplier,
208 drift_multiplier,
209 current_state: DriftLevel::None,
210 })
211 }
212
213 pub const fn base_lr(&self) -> f64 {
215 self.base_lr
216 }
217
218 pub const fn warning_multiplier(&self) -> f64 {
220 self.warning_multiplier
221 }
222
223 pub const fn drift_multiplier(&self) -> f64 {
225 self.drift_multiplier
226 }
227
228 pub const fn current_state(&self) -> DriftLevel {
230 self.current_state
231 }
232
233 pub fn on_drift_level(&mut self, level: DriftLevel) {
235 self.current_state = level;
236 }
237
238 pub fn current_lr(&self) -> f64 {
240 match self.current_state {
241 DriftLevel::None => self.base_lr,
242 DriftLevel::Warning => self.base_lr * self.warning_multiplier,
243 DriftLevel::Drift => self.base_lr * self.drift_multiplier,
244 }
245 }
246
247 pub fn reset(&mut self) {
249 self.current_state = DriftLevel::None;
250 }
251}
252
253impl Default for LearningRateScheduler {
254 fn default() -> Self {
255 Self::new(0.01, 2.0, 5.0).expect("default config is valid")
256 }
257}
258
259#[derive(Debug, Clone)]
286#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
287#[cfg_attr(feature = "serde", serde(try_from = "PersistedFixedWindowBuffer"))]
288pub struct FixedWindowBuffer {
289 buffer: Vec<f64>,
290 capacity: usize,
291 head: usize,
292 len: usize,
293}
294
295impl FixedWindowBuffer {
296 pub fn new(capacity: usize) -> Result<Self, RillError> {
300 if capacity == 0 {
301 return Err(RillError::InvalidCapacity(capacity));
302 }
303 Ok(Self {
304 buffer: vec![0.0; capacity],
305 capacity,
306 head: 0,
307 len: 0,
308 })
309 }
310
311 pub const fn capacity(&self) -> usize {
313 self.capacity
314 }
315
316 pub const fn len(&self) -> usize {
318 self.len
319 }
320
321 pub const fn is_empty(&self) -> bool {
323 self.len == 0
324 }
325
326 pub const fn is_full(&self) -> bool {
328 self.len == self.capacity
329 }
330
331 pub fn push(&mut self, value: f64) -> Result<(), RillError> {
335 ensure_finite("value", value)?;
336 self.buffer[self.head] = value;
337 self.head = (self.head + 1) % self.capacity;
338 if self.len < self.capacity {
339 self.len += 1;
340 }
341 Ok(())
342 }
343
344 pub fn mean(&self) -> Option<f64> {
346 if self.len == 0 {
347 return None;
348 }
349 let sum: f64 = self.iter().sum();
350 if !sum.is_finite() {
351 return None;
352 }
353 Some(sum / self.len as f64)
354 }
355
356 pub fn iter(&self) -> impl Iterator<Item = &f64> {
358 let start = if self.is_full() { self.head } else { 0 };
359 let len = self.len;
360 let cap = self.capacity;
361 (0..len).map(move |i| &self.buffer[(start + i) % cap])
362 }
363
364 pub fn reset(&mut self) {
366 self.head = 0;
367 self.len = 0;
368 }
369}
370
371#[cfg(feature = "serde")]
372impl FixedWindowBuffer {
373 fn validate_persisted_state(&self) -> Result<(), RillError> {
374 if self.capacity == 0 {
375 return Err(RillError::InvalidState(
376 "fixed-window capacity must be greater than zero".into(),
377 ));
378 }
379 if self.buffer.len() != self.capacity {
380 return Err(RillError::InvalidState(format!(
381 "fixed-window backing length {} does not match capacity {}",
382 self.buffer.len(),
383 self.capacity
384 )));
385 }
386 if self.len > self.capacity {
387 return Err(RillError::InvalidState(format!(
388 "fixed-window length {} exceeds capacity {}",
389 self.len, self.capacity
390 )));
391 }
392 if self.head >= self.capacity {
393 return Err(RillError::InvalidState(format!(
394 "fixed-window head {} is outside capacity {}",
395 self.head, self.capacity
396 )));
397 }
398 if self.len < self.capacity && self.head != self.len {
399 return Err(RillError::InvalidState(format!(
400 "fixed-window head {} must equal length {} before the buffer is full",
401 self.head, self.len
402 )));
403 }
404 for value in self.buffer.iter().take(self.len) {
405 ensure_finite("buffer value", *value)?;
406 }
407 Ok(())
408 }
409}
410
411#[cfg(feature = "serde")]
412#[derive(serde::Deserialize)]
413struct PersistedFixedWindowBuffer {
414 buffer: Vec<f64>,
415 capacity: usize,
416 head: usize,
417 len: usize,
418}
419
420#[cfg(feature = "serde")]
421impl TryFrom<PersistedFixedWindowBuffer> for FixedWindowBuffer {
422 type Error = RillError;
423
424 fn try_from(persisted: PersistedFixedWindowBuffer) -> Result<Self, Self::Error> {
425 let buffer = Self {
426 buffer: persisted.buffer,
427 capacity: persisted.capacity,
428 head: persisted.head,
429 len: persisted.len,
430 };
431 buffer.validate_persisted_state()?;
432 Ok(buffer)
433 }
434}
435
436#[cfg(test)]
441mod tests {
442 use super::*;
443
444 #[test]
447 fn tdm_first_sample_seeds_mean() {
448 let mut m = TimeDecayedMean::new(0.1).unwrap();
449 m.update(0.0, 10.0).unwrap();
450 assert!((m.value().unwrap() - 10.0).abs() < 1e-12);
451 }
452
453 #[test]
454 fn tdm_decay_weights_old_samples() {
455 let mut m = TimeDecayedMean::new(1.0).unwrap();
456 m.update(0.0, 100.0).unwrap();
457 m.update(10.0, 1.0).unwrap();
458 let v = m.value().unwrap();
461 assert!(
462 (v - 1.0).abs() < 0.01,
463 "recent sample should dominate, got {}",
464 v
465 );
466 }
467
468 #[test]
469 fn tdm_value_correct() {
470 let mut m = TimeDecayedMean::new(0.5).unwrap();
471 m.update(0.0, 10.0).unwrap();
472 m.update(1.0, 20.0).unwrap();
473 let v = m.value().unwrap();
478 assert!((v - 16.22).abs() < 0.1, "expected ~16.22, got {}", v);
479 }
480
481 #[test]
482 fn tdm_reset_clears_state() {
483 let mut m = TimeDecayedMean::new(0.1).unwrap();
484 m.update(0.0, 10.0).unwrap();
485 m.update(1.0, 20.0).unwrap();
486 assert!(m.value().is_some());
487 m.reset();
488 assert!(m.value().is_none());
489 assert_eq!(m.weight_total(), 0.0);
490 assert_eq!(m.last_time(), None);
491 }
492
493 #[test]
494 fn tdm_rejects_invalid_decay() {
495 assert!(TimeDecayedMean::new(0.0).is_err());
496 assert!(TimeDecayedMean::new(-1.0).is_err());
497 assert!(TimeDecayedMean::new(f64::NAN).is_err());
498 assert!(TimeDecayedMean::new(f64::INFINITY).is_err());
499 }
500
501 #[test]
502 fn tdm_rejects_non_finite() {
503 let mut m = TimeDecayedMean::new(0.1).unwrap();
504 assert!(m.update(f64::NAN, 1.0).is_err());
505 assert!(m.update(1.0, f64::NAN).is_err());
506 assert!(m.update(f64::INFINITY, 1.0).is_err());
507 assert!(m.update(1.0, f64::INFINITY).is_err());
508 }
509
510 #[test]
511 fn tdm_rejects_negative_dt() {
512 let mut m = TimeDecayedMean::new(0.1).unwrap();
513 m.update(5.0, 10.0).unwrap();
514 assert!(m.update(3.0, 20.0).is_err());
515 }
516
517 #[test]
518 fn tdm_equal_time_no_decay() {
519 let mut m = TimeDecayedMean::new(1.0).unwrap();
520 m.update(0.0, 10.0).unwrap();
521 m.update(0.0, 20.0).unwrap();
522 assert!((m.value().unwrap() - 15.0).abs() < 1e-12);
524 }
525
526 #[cfg(feature = "serde")]
527 #[test]
528 fn tdm_serde_roundtrip() {
529 let mut m = TimeDecayedMean::new(0.5).unwrap();
530 m.update(0.0, 10.0).unwrap();
531 m.update(1.0, 20.0).unwrap();
532 let json = serde_json::to_string(&m).unwrap();
533 let restored: TimeDecayedMean = serde_json::from_str(&json).unwrap();
534 assert!((restored.decay() - 0.5).abs() < 1e-12);
535 assert!((restored.value().unwrap() - m.value().unwrap()).abs() < 1e-12);
536 }
537
538 #[test]
541 fn lrs_default_lr() {
542 let sched = LearningRateScheduler::default();
543 assert!((sched.current_lr() - 0.01).abs() < 1e-12);
544 assert_eq!(sched.current_state(), DriftLevel::None);
545 }
546
547 #[test]
548 fn lrs_warning_increases_lr() {
549 let mut sched = LearningRateScheduler::new(0.05, 2.0, 5.0).unwrap();
550 sched.on_drift_level(DriftLevel::Warning);
551 assert!((sched.current_lr() - 0.10).abs() < 1e-12);
552 }
553
554 #[test]
555 fn lrs_drift_increases_more() {
556 let mut sched = LearningRateScheduler::new(0.05, 2.0, 5.0).unwrap();
557 sched.on_drift_level(DriftLevel::Drift);
558 assert!((sched.current_lr() - 0.25).abs() < 1e-12);
559 }
560
561 #[test]
562 fn lrs_reset_to_base() {
563 let mut sched = LearningRateScheduler::new(0.05, 2.0, 5.0).unwrap();
564 sched.on_drift_level(DriftLevel::Drift);
565 sched.reset();
566 assert_eq!(sched.current_state(), DriftLevel::None);
567 assert!((sched.current_lr() - 0.05).abs() < 1e-12);
568 }
569
570 #[test]
571 fn lrs_rejects_invalid_config() {
572 assert!(LearningRateScheduler::new(0.0, 2.0, 5.0).is_err());
574 assert!(LearningRateScheduler::new(-1.0, 2.0, 5.0).is_err());
575 assert!(LearningRateScheduler::new(0.01, 0.5, 5.0).is_err());
577 assert!(LearningRateScheduler::new(0.01, 3.0, 2.0).is_err());
579 assert!(LearningRateScheduler::new(f64::NAN, 2.0, 5.0).is_err());
581 }
582
583 #[cfg(feature = "serde")]
584 #[test]
585 fn lrs_serde_roundtrip() {
586 let mut sched = LearningRateScheduler::new(0.02, 3.0, 7.0).unwrap();
587 sched.on_drift_level(DriftLevel::Warning);
588 let json = serde_json::to_string(&sched).unwrap();
589 let restored: LearningRateScheduler = serde_json::from_str(&json).unwrap();
590 assert!((restored.base_lr() - 0.02).abs() < 1e-12);
591 assert!((restored.warning_multiplier() - 3.0).abs() < 1e-12);
592 assert!((restored.drift_multiplier() - 7.0).abs() < 1e-12);
593 assert_eq!(restored.current_state(), DriftLevel::Warning);
594 assert!((restored.current_lr() - 0.06).abs() < 1e-12);
595 }
596
597 #[test]
600 fn fwb_push_below_capacity() {
601 let mut buf = FixedWindowBuffer::new(5).unwrap();
602 buf.push(1.0).unwrap();
603 buf.push(2.0).unwrap();
604 buf.push(3.0).unwrap();
605 assert_eq!(buf.len(), 3);
606 assert!(!buf.is_full());
607 assert!(!buf.is_empty());
608 let collected: Vec<f64> = buf.iter().copied().collect();
609 assert_eq!(collected, vec![1.0, 2.0, 3.0]);
610 }
611
612 #[test]
613 fn fwb_push_overwrites_oldest() {
614 let mut buf = FixedWindowBuffer::new(3).unwrap();
615 buf.push(1.0).unwrap();
616 buf.push(2.0).unwrap();
617 buf.push(3.0).unwrap();
618 assert!(buf.is_full());
619 buf.push(4.0).unwrap();
620 let collected: Vec<f64> = buf.iter().copied().collect();
622 assert_eq!(collected, vec![2.0, 3.0, 4.0]);
623 assert_eq!(buf.len(), 3);
624 }
625
626 #[test]
627 fn fwb_mean_correct() {
628 let mut buf = FixedWindowBuffer::new(4).unwrap();
629 buf.push(1.0).unwrap();
630 buf.push(2.0).unwrap();
631 buf.push(3.0).unwrap();
632 buf.push(4.0).unwrap();
633 assert_eq!(buf.mean(), Some(2.5));
634 buf.push(10.0).unwrap(); assert_eq!(buf.mean(), Some((2.0 + 3.0 + 4.0 + 10.0) / 4.0));
636 }
637
638 #[test]
639 fn fwb_iter_returns_in_order() {
640 let mut buf = FixedWindowBuffer::new(3).unwrap();
641 for v in &[10.0, 20.0, 30.0, 40.0, 50.0] {
642 buf.push(*v).unwrap();
643 }
644 let collected: Vec<f64> = buf.iter().copied().collect();
646 assert_eq!(collected, vec![30.0, 40.0, 50.0]);
647 }
648
649 #[test]
650 fn fwb_empty_buffer_mean_none() {
651 let buf = FixedWindowBuffer::new(3).unwrap();
652 assert_eq!(buf.mean(), None);
653 assert!(buf.is_empty());
654 assert!(!buf.is_full());
655 }
656
657 #[test]
658 fn fwb_rejects_zero_capacity() {
659 assert!(FixedWindowBuffer::new(0).is_err());
660 }
661
662 #[test]
663 fn fwb_rejects_non_finite() {
664 let mut buf = FixedWindowBuffer::new(3).unwrap();
665 assert!(buf.push(f64::NAN).is_err());
666 assert!(buf.push(f64::INFINITY).is_err());
667 assert!(buf.push(f64::NEG_INFINITY).is_err());
668 assert_eq!(buf.len(), 0);
669 }
670
671 #[test]
672 fn fwb_reset_clears() {
673 let mut buf = FixedWindowBuffer::new(3).unwrap();
674 buf.push(1.0).unwrap();
675 buf.push(2.0).unwrap();
676 buf.reset();
677 assert_eq!(buf.len(), 0);
678 assert!(buf.is_empty());
679 assert_eq!(buf.mean(), None);
680 }
681
682 #[test]
683 fn fwb_wrap_around_multiple_times() {
684 let mut buf = FixedWindowBuffer::new(2).unwrap();
685 for i in 1..=10 {
686 buf.push(i as f64).unwrap();
687 }
688 assert_eq!(buf.len(), 2);
689 assert!(buf.is_full());
690 let collected: Vec<f64> = buf.iter().copied().collect();
691 assert_eq!(collected, vec![9.0, 10.0]);
692 }
693
694 #[cfg(feature = "serde")]
695 #[test]
696 fn fwb_serde_roundtrip() {
697 let mut buf = FixedWindowBuffer::new(3).unwrap();
698 buf.push(1.0).unwrap();
699 buf.push(2.0).unwrap();
700 buf.push(3.0).unwrap();
701 buf.push(4.0).unwrap(); let json = serde_json::to_string(&buf).unwrap();
703 let restored: FixedWindowBuffer = serde_json::from_str(&json).unwrap();
704 assert_eq!(restored.capacity(), 3);
705 assert_eq!(restored.len(), 3);
706 assert!(restored.is_full());
707 let collected: Vec<f64> = restored.iter().copied().collect();
708 assert_eq!(collected, vec![2.0, 3.0, 4.0]);
709 }
710}