use std::collections::{BTreeMap, HashMap};
use serde::Serialize;
use super::COCOeval;
use crate::metrics::calibration::{calibration_curve, calibration_error, scores_in_unit_interval};
use crate::metrics::calibration::CalibrationBin;
#[derive(Debug, Clone, Serialize)]
pub struct CalibrationResult {
pub ece: f64,
pub mce: f64,
pub bins: Vec<CalibrationBin>,
pub per_category: BTreeMap<u64, f64>,
pub iou_threshold: f64,
pub n_bins: usize,
pub num_detections: usize,
}
type ScoredOutcomes = (Vec<f64>, Vec<bool>);
impl COCOeval {
pub fn calibration(
&self,
n_bins: usize,
iou_threshold: f64,
) -> crate::error::Result<CalibrationResult> {
if self.eval_imgs.is_empty() {
return Err("calibration() requires evaluate() to be called first".into());
}
let t_idx = self.params.iou_thr_idx(iou_threshold).ok_or_else(|| {
format!(
"iou_threshold={iou_threshold} not found in params.iou_thrs={:?}",
self.params.iou_thrs
)
})?;
let mut all: ScoredOutcomes = (Vec::new(), Vec::new());
let mut per_cat: HashMap<u64, ScoredOutcomes> = HashMap::new();
for eval_img in self.default_cells() {
let matched = eval_img.dt_matched.row(t_idx);
let ignored = eval_img.dt_ignore.row(t_idx);
debug_assert_eq!(matched.len(), eval_img.dt_scores.len());
debug_assert_eq!(ignored.len(), eval_img.dt_scores.len());
let n = matched
.len()
.min(ignored.len())
.min(eval_img.dt_scores.len());
let cat = per_cat.entry(eval_img.category_id).or_default();
for d in 0..n {
if ignored[d] {
continue;
}
let (score, correct) = (eval_img.dt_scores[d], matched[d]);
cat.0.push(score);
cat.1.push(correct);
all.0.push(score);
all.1.push(correct);
}
}
scores_in_unit_interval(&all.0).map_err(|e| format!("calibration(): {e}"))?;
let bins = calibration_curve(&all.0, &all.1, n_bins);
let (ece, mce) = calibration_error(&bins);
let per_category: BTreeMap<u64, f64> = per_cat
.iter()
.map(|(&cat_id, (scores, matched))| {
let cat_bins = calibration_curve(scores, matched, n_bins);
(cat_id, calibration_error(&cat_bins).0)
})
.collect();
Ok(CalibrationResult {
ece,
mce,
bins,
per_category,
iou_threshold,
n_bins,
num_detections: all.0.len(),
})
}
}