use crate::error::{DriftError, Result};
use model_selection_rs::scoring::Scorer;
use ndarray::Array1;
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct LabelDriftReport {
pub rolling_score: f64,
pub baseline_score: f64,
pub degradation: f64,
pub drifted: bool,
}
pub struct LabelDriftMonitor<S: Scorer> {
scorer: S,
baseline_score: f64,
degradation_threshold: f64,
}
impl<S: Scorer> LabelDriftMonitor<S> {
pub fn new(scorer: S, baseline_score: f64, degradation_threshold: f64) -> Self {
Self {
scorer,
baseline_score,
degradation_threshold,
}
}
pub fn from_reference(
scorer: S,
reference_pairs: &[(f64, f64)],
degradation_threshold: f64,
) -> Result<Self> {
let baseline_score = score_pairs(&scorer, reference_pairs)?;
Ok(Self::new(scorer, baseline_score, degradation_threshold))
}
pub fn baseline_score(&self) -> f64 {
self.baseline_score
}
pub fn check(&self, pairs: &[(f64, f64)]) -> Result<LabelDriftReport> {
let rolling_score = score_pairs(&self.scorer, pairs)?;
let degradation = if self.scorer.greater_is_better() {
self.baseline_score - rolling_score
} else {
rolling_score - self.baseline_score
};
Ok(LabelDriftReport {
rolling_score,
baseline_score: self.baseline_score,
degradation,
drifted: degradation > self.degradation_threshold,
})
}
}
fn score_pairs<S: Scorer>(scorer: &S, pairs: &[(f64, f64)]) -> Result<f64> {
if pairs.is_empty() {
return Err(DriftError::EmptyInput(
"label-drift scoring needs at least one (prediction, actual) pair".into(),
));
}
let y_pred: Array1<f64> = pairs.iter().map(|&(p, _)| p).collect();
let y_true: Array1<f64> = pairs.iter().map(|&(_, a)| a).collect();
Ok(scorer.score(&y_true, &y_pred))
}