#![cfg(feature = "label-drift")]
use driftwatch::LabelDriftMonitor;
use model_selection_rs::scoring::Scorer;
use ndarray::Array1;
struct Accuracy;
impl Scorer for Accuracy {
fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
let n = y_true.len();
if n == 0 {
return 0.0;
}
let correct = y_true
.iter()
.zip(y_pred.iter())
.filter(|(a, b)| (**a - **b).abs() < 1e-9)
.count();
correct as f64 / n as f64
}
fn name(&self) -> &str {
"accuracy"
}
fn greater_is_better(&self) -> bool {
true
}
}
fn pairs_with_errors(n: usize, wrong: usize) -> Vec<(f64, f64)> {
(0..n)
.map(|i| {
let actual = 1.0;
let prediction = if i < wrong { 0.0 } else { 1.0 };
(prediction, actual)
})
.collect()
}
#[test]
fn flags_degrading_accuracy() {
let reference = pairs_with_errors(100, 0);
let monitor = LabelDriftMonitor::from_reference(Accuracy, &reference, 0.1).unwrap();
assert!((monitor.baseline_score() - 1.0).abs() < 1e-9);
let degraded = pairs_with_errors(100, 30);
let report = monitor.check(°raded).unwrap();
assert!((report.rolling_score - 0.7).abs() < 1e-9);
assert!((report.degradation - 0.3).abs() < 1e-9);
assert!(report.drifted);
}
#[test]
fn stable_accuracy_does_not_flag() {
let reference = pairs_with_errors(100, 0);
let monitor = LabelDriftMonitor::from_reference(Accuracy, &reference, 0.1).unwrap();
let ok = pairs_with_errors(100, 5);
let report = monitor.check(&ok).unwrap();
assert!(!report.drifted);
}
#[test]
fn empty_window_is_an_error() {
let monitor = LabelDriftMonitor::new(Accuracy, 1.0, 0.1);
assert!(monitor.check(&[]).is_err());
}