use std::sync::Mutex;
use driftwatch::{EqualWidthBinning, MetricKind, PredictionDriftMonitor};
use crate::error::{Error, Result};
fn drift_err(e: impl std::fmt::Display) -> Error {
Error::Backend(format!("driftwatch: {e}"))
}
#[derive(Clone, Copy, Debug, Default)]
pub struct DriftStatus {
pub drifted: bool,
pub psi: f64,
pub observed: usize,
}
pub struct DriftMonitor {
inner: PredictionDriftMonitor,
live: Mutex<Vec<f64>>,
}
impl DriftMonitor {
pub fn psi(reference_predictions: &[f64]) -> Result<Self> {
let inner =
PredictionDriftMonitor::new(reference_predictions, EqualWidthBinning::default())
.map_err(drift_err)?;
Ok(DriftMonitor {
inner,
live: Mutex::new(Vec::new()),
})
}
pub fn observe(&self, predictions: &[f64]) {
self.live
.lock()
.expect("monitor lock")
.extend_from_slice(predictions);
}
pub fn report(&self) -> Result<DriftStatus> {
let live = self.live.lock().expect("monitor lock").clone();
if live.is_empty() {
return Ok(DriftStatus::default());
}
let report = self.inner.check(&live).map_err(drift_err)?;
let psi = report
.features
.first()
.and_then(|f| f.score(MetricKind::Psi))
.map(|s| s.statistic)
.unwrap_or(0.0);
let drifted = report.drifted_features().next().is_some();
Ok(DriftStatus {
drifted,
psi,
observed: live.len(),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stable_stream_does_not_drift_but_shifted_stream_does() {
let reference: Vec<f64> = (0..200).map(|i| (i % 10) as f64 * 0.1).collect();
let monitor = DriftMonitor::psi(&reference).unwrap();
let stable: Vec<f64> = (0..200).map(|i| (i % 10) as f64 * 0.1).collect();
monitor.observe(&stable);
let stable_report = monitor.report().unwrap();
assert!(!stable_report.drifted, "psi = {}", stable_report.psi);
let shifted = DriftMonitor::psi(&reference).unwrap();
let far: Vec<f64> = (0..200).map(|_| 100.0).collect();
shifted.observe(&far);
let shifted_report = shifted.report().unwrap();
assert!(shifted_report.drifted, "psi = {}", shifted_report.psi);
assert!(shifted_report.psi > stable_report.psi);
}
}