use super::StatsError;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct RocPoint {
pub fpr: f64,
pub tpr: f64,
pub threshold: f64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct RocCurveData {
pub points: Vec<RocPoint>,
pub auc: f64,
}
pub fn auc_trapezoid(points: &[(f64, f64)]) -> f64 {
points
.windows(2)
.map(|w| {
let (x1, y1) = w[0];
let (x2, y2) = w[1];
(x2 - x1) * (y1 + y2) / 2.0
})
.sum()
}
pub fn roc_curve(scores: &[f64], labels: &[bool]) -> Result<RocCurveData, StatsError> {
if scores.len() != labels.len() {
return Err(StatsError::LengthMismatch {
scores: scores.len(),
labels: labels.len(),
});
}
if scores.is_empty() {
return Err(StatsError::EmptyInput);
}
let total_p = labels.iter().filter(|&&l| l).count();
let total_n = labels.len() - total_p;
if total_p == 0 {
return Err(StatsError::NoPositiveLabels);
}
if total_n == 0 {
return Err(StatsError::NoNegativeLabels);
}
let mut order: Vec<usize> = (0..scores.len()).collect();
order.sort_by(|&a, &b| {
scores[b]
.partial_cmp(&scores[a])
.unwrap_or(std::cmp::Ordering::Equal)
});
let (p, n) = (total_p as f64, total_n as f64);
let mut points = Vec::with_capacity(order.len() + 1);
points.push(RocPoint {
fpr: 0.0,
tpr: 0.0,
threshold: f64::INFINITY,
});
let mut tp = 0usize;
let mut fp = 0usize;
let mut i = 0usize;
while i < order.len() {
let score = scores[order[i]];
while i < order.len() && scores[order[i]] == score {
if labels[order[i]] {
tp += 1;
} else {
fp += 1;
}
i += 1;
}
points.push(RocPoint {
fpr: fp as f64 / n,
tpr: tp as f64 / p,
threshold: score,
});
}
let auc = auc_trapezoid(&points.iter().map(|pt| (pt.fpr, pt.tpr)).collect::<Vec<_>>());
Ok(RocCurveData { points, auc })
}