use super::StatsError;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct CalibrationBin {
pub mean_predicted: f64,
pub observed_freq: f64,
pub count: usize,
}
pub fn calibration_curve(
scores: &[f64],
labels: &[bool],
n_bins: usize,
) -> Result<Vec<CalibrationBin>, StatsError> {
if scores.len() != labels.len() {
return Err(StatsError::LengthMismatch {
scores: scores.len(),
labels: labels.len(),
});
}
if scores.is_empty() || n_bins == 0 {
return Err(StatsError::EmptyInput);
}
let mut sum_pred = vec![0.0; n_bins];
let mut sum_pos = vec![0usize; n_bins];
let mut count = vec![0usize; n_bins];
for (&s, &l) in scores.iter().zip(labels) {
let c = s.clamp(0.0, 1.0);
let mut idx = (c * n_bins as f64).floor() as usize;
if idx >= n_bins {
idx = n_bins - 1;
}
sum_pred[idx] += c;
sum_pos[idx] += l as usize;
count[idx] += 1;
}
Ok((0..n_bins)
.filter(|&b| count[b] > 0)
.map(|b| CalibrationBin {
mean_predicted: sum_pred[b] / count[b] as f64,
observed_freq: sum_pos[b] as f64 / count[b] as f64,
count: count[b],
})
.collect())
}