Skip to main content

rill_ml/drift/
adwin.rs

1//! ADWIN (Adaptive Windowing) drift detector.
2//!
3//! ADWIN maintains a variable-length window of recent observations and
4//! detects when the distribution of the window's two halves differs
5//! significantly. When drift is detected, the older portion is dropped.
6//!
7//! ## Algorithm
8//!
9//! Based on Bifet & Gavaldà (2007). For each new observation:
10//!
11//! 1. Add the value to the window.
12//! 2. For each possible split point `k` (dividing the window into `W0`
13//!    and `W1`), compute the means `μ0` and `μ1`.
14//! 3. Compute the Hoeffding bound:
15//!    `ε = √(1/(2·m) · ln(4/δ'))` where `m = n0·n1/(n0+n1)` and
16//!    `δ' = δ / ln(n)` (Bonferroni correction for repeated testing).
17//! 4. If `|μ0 − μ1| > ε`, drop `W0` and signal drift.
18//!
19//! ## Space complexity
20//!
21//! `O(max_window)` — the window stores individual values up to
22//! [`AdwinConfig::max_window`] elements. Prefix sums are maintained
23//! incrementally so each update is `O(max_window)` in the worst case.
24
25use crate::drift::detector::{DriftDetector, DriftLevel};
26use crate::error::{RillError, ensure_finite};
27
28/// Configuration for [`Adwin`].
29#[derive(Debug, Clone)]
30#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
31#[non_exhaustive]
32pub struct AdwinConfig {
33    /// Significance level for drift detection. Must be in `(0, 1)`.
34    /// Smaller values reduce false positives. Defaults to `0.002`.
35    pub delta: f64,
36
37    /// Significance level for warnings. Must be in `(0, 1)` and
38    /// greater than or equal to `delta`. Defaults to `0.01`.
39    pub warning_delta: f64,
40
41    /// Maximum window size (number of stored observations). Must be > 0.
42    /// Larger values improve detection sensitivity but increase memory
43    /// and computation. Defaults to `1000`.
44    pub max_window: usize,
45
46    /// Minimum number of samples before any detection is attempted.
47    /// Must be greater than zero. Defaults to `10`.
48    pub min_samples: u64,
49}
50
51impl Default for AdwinConfig {
52    fn default() -> Self {
53        Self {
54            delta: 0.002,
55            warning_delta: 0.01,
56            max_window: 1000,
57            min_samples: 10,
58        }
59    }
60}
61
62/// ADWIN (Adaptive Windowing) drift detector.
63///
64/// Maintains a variable-length window and detects distribution changes by
65/// comparing the means of the window's two halves. See the module
66/// documentation for the algorithm.
67///
68/// # Examples
69///
70/// ```
71/// use rill_ml::drift::{Adwin, DriftDetector, DriftLevel};
72///
73/// let mut adwin = Adwin::default();
74///
75/// // Stable stream.
76/// for _ in 0..100 {
77///     adwin.update(0.0).unwrap();
78/// }
79/// assert_eq!(adwin.level(), DriftLevel::None);
80///
81/// // Sudden shift. ADWIN's level is transient: after detecting drift and
82/// // trimming the window, subsequent stable updates reset the level to
83/// // None, so we check the level returned by each update.
84/// let mut detected = false;
85/// for _ in 0..100 {
86///     let level = adwin.update(5.0).unwrap();
87///     if level == DriftLevel::Drift {
88///         detected = true;
89///         break;
90///     }
91/// }
92/// assert!(detected);
93/// ```
94#[derive(Debug, Clone)]
95#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
96pub struct Adwin {
97    config: AdwinConfig,
98    window: std::collections::VecDeque<f64>,
99    total: f64,
100    samples: u64,
101    current_level: DriftLevel,
102}
103
104impl Adwin {
105    /// Create a new ADWIN detector with the given configuration.
106    ///
107    /// Returns an error if:
108    /// - `delta` is not in `(0, 1)`.
109    /// - `warning_delta` is not in `(0, 1)` or is less than `delta`.
110    /// - `max_window` is zero.
111    /// - `min_samples` is zero.
112    pub fn new(config: AdwinConfig) -> Result<Self, RillError> {
113        ensure_finite("delta", config.delta)?;
114        if config.delta <= 0.0 || config.delta >= 1.0 {
115            return Err(RillError::InvalidSignificanceLevel(config.delta));
116        }
117        ensure_finite("warning_delta", config.warning_delta)?;
118        if config.warning_delta <= 0.0 || config.warning_delta >= 1.0 {
119            return Err(RillError::InvalidSignificanceLevel(config.warning_delta));
120        }
121        if config.warning_delta < config.delta {
122            return Err(RillError::InvalidParameter {
123                name: "warning_delta",
124                value: config.warning_delta,
125            });
126        }
127        if config.max_window == 0 {
128            return Err(RillError::InvalidCapacity(config.max_window));
129        }
130        if config.min_samples == 0 {
131            return Err(RillError::InvalidParameter {
132                name: "min_samples",
133                value: 0.0,
134            });
135        }
136        Ok(Self {
137            window: std::collections::VecDeque::with_capacity(config.max_window),
138            config,
139            total: 0.0,
140            samples: 0,
141            current_level: DriftLevel::None,
142        })
143    }
144
145    /// The number of values currently in the window.
146    pub fn window_size(&self) -> usize {
147        self.window.len()
148    }
149
150    /// The mean of all values currently in the window, or `0.0` if empty.
151    pub fn window_mean(&self) -> f64 {
152        if self.window.is_empty() {
153            0.0
154        } else {
155            self.total / self.window.len() as f64
156        }
157    }
158
159    /// The configuration of this detector.
160    pub const fn config(&self) -> &AdwinConfig {
161        &self.config
162    }
163
164    /// Compute the Hoeffding bound for a split with `n0` and `n1` elements
165    /// at significance level `delta` with total stream length `n`.
166    fn hoeffding_bound(n0: f64, n1: f64, n: u64, delta: f64) -> f64 {
167        let m = n0 * n1 / (n0 + n1);
168        let ln_n = (n as f64).ln().max(1.0);
169        let delta_eff = delta / ln_n;
170        (1.0 / (2.0 * m) * (4.0 / delta_eff).ln()).sqrt()
171    }
172
173    /// Check all split points and return the index where drift is detected,
174    /// along with the level (Warning or Drift). Returns `None` if no split
175    /// exceeds the threshold.
176    fn check_splits(&self) -> Option<(usize, DriftLevel, f64)> {
177        let n = self.window.len();
178        if n < 2 {
179            return None;
180        }
181        // Build prefix sums for O(1) mean computation per split.
182        let mut prefix = Vec::with_capacity(n + 1);
183        prefix.push(0.0_f64);
184        let mut acc = 0.0;
185        for &v in &self.window {
186            acc += v;
187            prefix.push(acc);
188        }
189        let total = prefix[n];
190        let n_total = n as u64;
191
192        let mut best_split: Option<(usize, DriftLevel, f64)> = None;
193        // Check split points from the oldest end.
194        for (k, &sum0) in prefix.iter().enumerate().take(n).skip(1) {
195            let n0 = k as f64;
196            let n1 = (n - k) as f64;
197            let sum1 = total - sum0;
198            let mean0 = sum0 / n0;
199            let mean1 = sum1 / n1;
200            let diff = (mean0 - mean1).abs();
201
202            // Check drift threshold.
203            let eps_drift = Self::hoeffding_bound(n0, n1, n_total, self.config.delta);
204            if diff > eps_drift {
205                return Some((k, DriftLevel::Drift, diff));
206            }
207            // Check warning threshold.
208            let eps_warn = Self::hoeffding_bound(n0, n1, n_total, self.config.warning_delta);
209            if diff > eps_warn && best_split.is_none() {
210                best_split = Some((k, DriftLevel::Warning, diff));
211            }
212        }
213        best_split
214    }
215
216    /// Trim the window by removing the oldest `count` elements.
217    fn trim_front(&mut self, count: usize) {
218        for _ in 0..count {
219            if let Some(v) = self.window.pop_front() {
220                self.total -= v;
221            }
222        }
223    }
224}
225
226impl Default for Adwin {
227    fn default() -> Self {
228        Self::new(AdwinConfig::default()).expect("default config is valid")
229    }
230}
231
232impl DriftDetector for Adwin {
233    fn update(&mut self, value: f64) -> Result<DriftLevel, RillError> {
234        ensure_finite("value", value)?;
235        self.samples += 1;
236        // Add the new value to the window.
237        self.window.push_back(value);
238        self.total += value;
239        // Enforce the max window size by dropping the oldest element.
240        if self.window.len() > self.config.max_window
241            && let Some(v) = self.window.pop_front()
242        {
243            self.total -= v;
244        }
245        // Gate detection by minimum samples.
246        if self.samples < self.config.min_samples || self.window.len() < 2 {
247            self.current_level = DriftLevel::None;
248            return Ok(DriftLevel::None);
249        }
250        // Check for drift.
251        if let Some((split, level, _diff)) = self.check_splits() {
252            // Trim the older portion when drift or warning is detected.
253            // For Drift, trim aggressively. For Warning, keep the window intact.
254            if level == DriftLevel::Drift {
255                self.trim_front(split);
256            }
257            self.current_level = level;
258        } else {
259            self.current_level = DriftLevel::None;
260        }
261        Ok(self.current_level)
262    }
263
264    fn detected(&self) -> bool {
265        self.current_level == DriftLevel::Drift
266    }
267
268    fn warning(&self) -> bool {
269        self.current_level == DriftLevel::Warning
270    }
271
272    fn level(&self) -> DriftLevel {
273        self.current_level
274    }
275
276    fn samples_seen(&self) -> u64 {
277        self.samples
278    }
279
280    fn reset(&mut self) {
281        self.window.clear();
282        self.total = 0.0;
283        self.samples = 0;
284        self.current_level = DriftLevel::None;
285    }
286}
287
288#[cfg(test)]
289mod tests {
290    use super::*;
291
292    /// Deterministic pseudo-random number in `[0, 1)` using a simple LCG.
293    fn next_unit(seed: &mut u64) -> f64 {
294        *seed = seed
295            .wrapping_mul(6364136223846793005)
296            .wrapping_add(1442695040888963407);
297        ((*seed >> 11) as f64) / ((1u64 << 53) as f64)
298    }
299
300    #[test]
301    fn default_config_is_valid() {
302        let adwin = Adwin::default();
303        assert_eq!(adwin.samples_seen(), 0);
304        assert_eq!(adwin.level(), DriftLevel::None);
305        assert_eq!(adwin.window_size(), 0);
306    }
307
308    #[test]
309    fn detects_sudden_mean_shift() {
310        let mut adwin = Adwin::new(AdwinConfig {
311            delta: 0.05,
312            warning_delta: 0.1,
313            max_window: 500,
314            min_samples: 5,
315        })
316        .unwrap();
317        // Stable stream around 0.
318        let mut seed = 42u64;
319        for _ in 0..100 {
320            let noise = 0.1 * (next_unit(&mut seed) - 0.5);
321            adwin.update(noise).unwrap();
322        }
323        assert_eq!(adwin.level(), DriftLevel::None);
324        // Sudden shift to mean 5.
325        let mut detected = false;
326        for _ in 0..200 {
327            let noise = 0.1 * (next_unit(&mut seed) - 0.5);
328            let level = adwin.update(5.0 + noise).unwrap();
329            if level == DriftLevel::Drift {
330                detected = true;
331                break;
332            }
333        }
334        assert!(detected, "ADWIN should detect the sudden mean shift");
335    }
336
337    #[test]
338    fn no_false_positive_on_stable_stream() {
339        let mut adwin = Adwin::new(AdwinConfig {
340            delta: 0.002,
341            warning_delta: 0.01,
342            max_window: 500,
343            min_samples: 10,
344        })
345        .unwrap();
346        let mut seed = 7u64;
347        for _ in 0..2000 {
348            let noise = 0.5 * (next_unit(&mut seed) - 0.5);
349            adwin.update(noise).unwrap();
350        }
351        assert!(
352            !adwin.detected(),
353            "false positive: drift reported on stable stream"
354        );
355    }
356
357    #[test]
358    fn detects_gradual_drift() {
359        let mut adwin = Adwin::new(AdwinConfig {
360            delta: 0.05,
361            warning_delta: 0.1,
362            max_window: 300,
363            min_samples: 5,
364        })
365        .unwrap();
366        // Start at mean 0, gradually increase to mean 5.
367        let mut seed = 99u64;
368        let mut detected = false;
369        for i in 0..500 {
370            let mean = (i as f64 / 100.0).min(5.0);
371            let noise = 0.1 * (next_unit(&mut seed) - 0.5);
372            let level = adwin.update(mean + noise).unwrap();
373            if level == DriftLevel::Drift {
374                detected = true;
375                break;
376            }
377        }
378        assert!(detected, "ADWIN should detect gradual drift");
379    }
380
381    #[test]
382    fn window_trims_after_drift() {
383        let mut adwin = Adwin::new(AdwinConfig {
384            delta: 0.05,
385            warning_delta: 0.1,
386            max_window: 500,
387            min_samples: 5,
388        })
389        .unwrap();
390        // Build up a window of 100 samples at mean 0.
391        for _ in 0..100 {
392            adwin.update(0.0).unwrap();
393        }
394        let size_before = adwin.window_size();
395        assert!(size_before > 0);
396        // Shift to mean 10 to trigger drift.
397        let mut trimmed = false;
398        for _ in 0..200 {
399            adwin.update(10.0).unwrap();
400            if adwin.detected() {
401                // After drift detection and trimming, the window should be
402                // smaller than its peak (all 10.0 values are kept, old 0.0
403                // values are trimmed).
404                if adwin.window_size() < size_before + 200 {
405                    trimmed = true;
406                    break;
407                }
408            }
409        }
410        assert!(trimmed, "window should be trimmed after drift");
411    }
412
413    #[test]
414    fn max_window_enforced() {
415        let mut adwin = Adwin::new(AdwinConfig {
416            max_window: 50,
417            ..Default::default()
418        })
419        .unwrap();
420        for i in 0..200u64 {
421            adwin.update(i as f64).unwrap();
422        }
423        // Drift detection may trim the window below max_window; the invariant
424        // is that the window never exceeds max_window.
425        assert!(
426            adwin.window_size() <= 50,
427            "window should not exceed max_window, got {}",
428            adwin.window_size()
429        );
430    }
431
432    #[test]
433    fn min_samples_gates_detection() {
434        let mut adwin = Adwin::new(AdwinConfig {
435            delta: 0.5,
436            warning_delta: 0.5,
437            max_window: 100,
438            min_samples: 50,
439        })
440        .unwrap();
441        // 48 zeros: samples_seen = 48 < min_samples = 50.
442        for _ in 0..48 {
443            adwin.update(0.0).unwrap();
444        }
445        // Sample 49: extreme value, but 49 < min_samples = 50 → no detection.
446        adwin.update(100.0).unwrap();
447        assert_eq!(adwin.level(), DriftLevel::None);
448        // After min_samples, detection can trigger. The level is transient:
449        // after drift is detected and the window trimmed, subsequent stable
450        // updates reset the level to None. Track detection across the loop.
451        let mut detected = false;
452        for _ in 0..50 {
453            let level = adwin.update(100.0).unwrap();
454            if level.is_change() {
455                detected = true;
456            }
457        }
458        assert!(detected, "should have detected drift after min_samples");
459    }
460
461    #[test]
462    fn reset_clears_state() {
463        let mut adwin = Adwin::default();
464        for _ in 0..50 {
465            adwin.update(1.0).unwrap();
466        }
467        assert!(adwin.window_size() > 0);
468        adwin.reset();
469        assert_eq!(adwin.window_size(), 0);
470        assert_eq!(adwin.samples_seen(), 0);
471        assert_eq!(adwin.level(), DriftLevel::None);
472        assert_eq!(adwin.window_mean(), 0.0);
473    }
474
475    #[test]
476    fn rejects_non_finite_input() {
477        let mut adwin = Adwin::default();
478        assert!(adwin.update(f64::NAN).is_err());
479        assert!(adwin.update(f64::INFINITY).is_err());
480        assert!(adwin.update(f64::NEG_INFINITY).is_err());
481        assert_eq!(adwin.samples_seen(), 0);
482        assert_eq!(adwin.window_size(), 0);
483    }
484
485    #[test]
486    fn rejects_invalid_config() {
487        // delta out of range
488        assert!(
489            Adwin::new(AdwinConfig {
490                delta: 0.0,
491                ..Default::default()
492            })
493            .is_err()
494        );
495        assert!(
496            Adwin::new(AdwinConfig {
497                delta: 1.0,
498                ..Default::default()
499            })
500            .is_err()
501        );
502        // warning_delta < delta
503        assert!(
504            Adwin::new(AdwinConfig {
505                delta: 0.05,
506                warning_delta: 0.01,
507                ..Default::default()
508            })
509            .is_err()
510        );
511        // max_window == 0
512        assert!(
513            Adwin::new(AdwinConfig {
514                max_window: 0,
515                ..Default::default()
516            })
517            .is_err()
518        );
519        // min_samples == 0
520        assert!(
521            Adwin::new(AdwinConfig {
522                min_samples: 0,
523                ..Default::default()
524            })
525            .is_err()
526        );
527    }
528
529    #[test]
530    fn window_mean_correct() {
531        let mut adwin = Adwin::new(AdwinConfig {
532            max_window: 100,
533            min_samples: 11, // prevent drift detection with only 10 samples
534            ..Default::default()
535        })
536        .unwrap();
537        for i in 1..=10 {
538            adwin.update(i as f64).unwrap();
539        }
540        // mean of 1..=10 is 5.5
541        assert!((adwin.window_mean() - 5.5).abs() < 1e-9);
542    }
543
544    #[test]
545    fn hoeffding_bound_decreases_with_more_data() {
546        // With more data, the bound should be tighter (smaller).
547        let b1 = Adwin::hoeffding_bound(5.0, 5.0, 10, 0.01);
548        let b2 = Adwin::hoeffding_bound(50.0, 50.0, 100, 0.01);
549        assert!(
550            b2 < b1,
551            "bound should decrease with more data: {} vs {}",
552            b2,
553            b1
554        );
555    }
556
557    #[cfg(feature = "serde")]
558    #[test]
559    fn serde_roundtrip() {
560        let mut adwin = Adwin::new(AdwinConfig {
561            delta: 0.01,
562            warning_delta: 0.05,
563            max_window: 200,
564            min_samples: 5,
565        })
566        .unwrap();
567        for i in 0..50 {
568            adwin.update(i as f64 * 0.1).unwrap();
569        }
570        let json = serde_json::to_string(&adwin).unwrap();
571        let restored: Adwin = serde_json::from_str(&json).unwrap();
572        assert_eq!(restored.samples_seen(), 50);
573        assert_eq!(restored.window_size(), adwin.window_size());
574        assert!((restored.window_mean() - adwin.window_mean()).abs() < 1e-12);
575        assert_eq!(restored.level(), adwin.level());
576    }
577}