use crate::binning::{histogram_from_edges, BinDefinition, ContinuousBinning, Histogram};
use crate::error::{DriftError, Result};
#[derive(Clone, Copy, Debug)]
pub enum LiveFeature<'a> {
Continuous(&'a [f64]),
Categorical(&'a [&'a str]),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum FeatureKind {
Continuous,
Categorical,
}
enum Reference {
Continuous {
edges: Vec<f64>,
histogram: Histogram,
samples: Vec<f64>,
},
Categorical {
categories: Vec<String>,
counts: Vec<f64>,
histogram: Histogram,
},
}
pub struct ReferenceDistribution {
name: String,
reference: Reference,
}
pub(crate) struct Comparison<'a> {
pub reference_hist: Histogram,
pub live_hist: Histogram,
pub raw_samples: Option<(&'a [f64], &'a [f64])>,
pub kind: FeatureKind,
}
impl ReferenceDistribution {
pub fn fit_continuous(
name: impl Into<String>,
reference: &[f64],
binning: impl ContinuousBinning,
) -> Result<Self> {
let edges = binning.fit_edges(reference)?;
let histogram = histogram_from_edges(&edges, reference)?;
Ok(Self {
name: name.into(),
reference: Reference::Continuous {
edges,
histogram,
samples: reference.to_vec(),
},
})
}
pub fn fit_categorical(name: impl Into<String>, reference: &[&str]) -> Result<Self> {
if reference.is_empty() {
return Err(DriftError::EmptyInput(
"categorical reference needs at least one label".into(),
));
}
let mut categories: Vec<String> = reference.iter().map(|s| s.to_string()).collect();
categories.sort();
categories.dedup();
let counts = count_categories(&categories, reference);
let histogram = Histogram::new(
BinDefinition::Categorical {
categories: categories.clone(),
},
counts.clone(),
)?;
Ok(Self {
name: name.into(),
reference: Reference::Categorical {
categories,
counts,
histogram,
},
})
}
pub fn name(&self) -> &str {
&self.name
}
pub fn kind(&self) -> FeatureKind {
match self.reference {
Reference::Continuous { .. } => FeatureKind::Continuous,
Reference::Categorical { .. } => FeatureKind::Categorical,
}
}
pub fn histogram(&self) -> &Histogram {
match &self.reference {
Reference::Continuous { histogram, .. } => histogram,
Reference::Categorical { histogram, .. } => histogram,
}
}
pub(crate) fn compare<'a>(&'a self, live: LiveFeature<'a>) -> Result<Comparison<'a>> {
match (&self.reference, live) {
(
Reference::Continuous {
edges,
histogram,
samples,
},
LiveFeature::Continuous(live_samples),
) => {
let live_hist = histogram_from_edges(edges, live_samples)?;
Ok(Comparison {
reference_hist: histogram.clone(),
live_hist,
raw_samples: Some((samples.as_slice(), live_samples)),
kind: FeatureKind::Continuous,
})
}
(
Reference::Categorical {
categories, counts, ..
},
LiveFeature::Categorical(live_labels),
) => {
let (reference_hist, live_hist) =
align_categorical(categories, counts, live_labels)?;
Ok(Comparison {
reference_hist,
live_hist,
raw_samples: None,
kind: FeatureKind::Categorical,
})
}
(Reference::Continuous { .. }, LiveFeature::Categorical(_)) => {
Err(DriftError::UnknownFeature(format!(
"feature '{}' is continuous but was given categorical live data",
self.name
)))
}
(Reference::Categorical { .. }, LiveFeature::Continuous(_)) => {
Err(DriftError::UnknownFeature(format!(
"feature '{}' is categorical but was given continuous live data",
self.name
)))
}
}
}
}
fn count_categories(categories: &[String], data: &[&str]) -> Vec<f64> {
categories
.iter()
.map(|c| data.iter().filter(|&&d| d == c.as_str()).count() as f64)
.collect()
}
fn align_categorical(
reference_categories: &[String],
reference_counts: &[f64],
live: &[&str],
) -> Result<(Histogram, Histogram)> {
let mut novel: Vec<String> = live
.iter()
.filter(|&&l| !reference_categories.iter().any(|c| c == l))
.map(|s| s.to_string())
.collect();
novel.sort();
novel.dedup();
let mut union: Vec<String> = reference_categories.to_vec();
union.extend(novel.iter().cloned());
let mut ref_counts = reference_counts.to_vec();
ref_counts.extend(std::iter::repeat(0.0).take(novel.len()));
let live_counts = count_categories(&union, live);
let bins = BinDefinition::Categorical { categories: union };
let reference_hist = Histogram::new(bins.clone(), ref_counts)?;
let live_hist = Histogram::new(bins, live_counts)?;
Ok((reference_hist, live_hist))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::binning::EqualFrequencyBinning;
#[test]
fn continuous_reference_histogram_normalizes() {
let baseline: Vec<f64> = (0..100).map(|i| i as f64).collect();
let reference = ReferenceDistribution::fit_continuous(
"x",
&baseline,
EqualFrequencyBinning::new(10).unwrap(),
)
.unwrap();
let sum: f64 = reference.histogram().frequencies().iter().sum();
assert!((sum - 1.0).abs() < 1e-9);
assert_eq!(reference.kind(), FeatureKind::Continuous);
}
#[test]
fn novel_category_becomes_zero_reference_bin() {
let reference = ReferenceDistribution::fit_categorical("color", &["a", "a", "b"]).unwrap();
let live = ["a", "c", "c"];
let comparison = reference.compare(LiveFeature::Categorical(&live)).unwrap();
if let BinDefinition::Categorical { categories } = comparison.reference_hist.bins() {
assert_eq!(
categories,
&["a".to_string(), "b".to_string(), "c".to_string()]
);
} else {
panic!("expected categorical bins");
}
let ref_counts = comparison.reference_hist.counts();
assert_eq!(ref_counts, &[2.0, 1.0, 0.0]);
assert_eq!(comparison.live_hist.counts(), &[1.0, 0.0, 2.0]);
let psi = crate::psi(&comparison.reference_hist, &comparison.live_hist).unwrap();
assert!(psi.is_finite());
}
#[test]
fn kind_mismatch_is_rejected() {
let reference = ReferenceDistribution::fit_categorical("c", &["a", "b"]).unwrap();
let err = reference.compare(LiveFeature::Continuous(&[1.0, 2.0]));
assert!(err.is_err());
}
}