use std::collections::hash_map::Entry;
use std::collections::{BTreeMap, HashMap, HashSet};
use rayon::prelude::*;
use serde::Serialize;
use crate::metrics::counts::{ApScratch, average_precision_ranked_into};
use super::COCOeval;
use super::matching::EvalImg;
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub(super) enum ErrType {
Cls,
Loc,
Both,
Dupe,
Bkg,
}
impl ErrType {
pub(super) fn as_str(self) -> &'static str {
match self {
ErrType::Cls => "Cls",
ErrType::Loc => "Loc",
ErrType::Both => "Both",
ErrType::Dupe => "Dupe",
ErrType::Bkg => "Bkg",
}
}
}
const FP_TYPES: [ErrType; 5] = [
ErrType::Cls,
ErrType::Loc,
ErrType::Both,
ErrType::Dupe,
ErrType::Bkg,
];
pub(super) struct FpEvidence {
pub(super) max_same_iou: f64,
pub(super) max_cross_iou: f64,
pub(super) best_same_gt_matched: bool,
}
pub(super) fn classify_fp(ev: &FpEvidence, pos_thr: f64, bg_thr: f64) -> ErrType {
if ev.max_same_iou >= bg_thr && ev.max_same_iou <= pos_thr {
ErrType::Loc
} else if ev.max_cross_iou >= pos_thr {
ErrType::Cls
} else if ev.best_same_gt_matched {
ErrType::Dupe
} else if ev.max_cross_iou <= bg_thr {
ErrType::Bkg
} else {
ErrType::Both
}
}
type CrossIouMap = HashMap<u64, HashMap<u64, (f64, Option<u64>)>>;
#[derive(Default)]
struct SameClassScan {
max_iou: f64,
argmax_gt_ann_id: Option<u64>,
best_gt_matched: bool,
}
struct CatData {
scores: Vec<f64>,
matched: Vec<bool>,
ignored: Vec<bool>,
fp_types: Vec<Option<ErrType>>,
num_gt: usize,
}
impl CatData {
fn new() -> Self {
CatData {
scores: Vec::new(),
matched: Vec::new(),
ignored: Vec::new(),
fp_types: Vec::new(),
num_gt: 0,
}
}
fn rank_by_score_desc(&mut self) {
let mut order: Vec<usize> = (0..self.scores.len()).collect();
order.sort_by(|&a, &b| {
self.scores[b]
.partial_cmp(&self.scores[a])
.unwrap_or(std::cmp::Ordering::Equal)
});
self.scores = order.iter().map(|&i| self.scores[i]).collect();
self.matched = order.iter().map(|&i| self.matched[i]).collect();
self.ignored = order.iter().map(|&i| self.ignored[i]).collect();
self.fp_types = order.iter().map(|&i| self.fp_types[i]).collect();
}
fn extend(&mut self, mut other: CatData) {
self.scores.append(&mut other.scores);
self.matched.append(&mut other.matched);
self.ignored.append(&mut other.ignored);
self.fp_types.append(&mut other.fp_types);
self.num_gt += other.num_gt;
}
}
#[derive(Default)]
struct Classified {
cat_data: HashMap<u64, CatData>,
fp_counts: [u64; FP_TYPES.len()],
covered_gts: HashSet<u64>,
}
impl Classified {
fn merge(mut self, other: Self) -> Self {
for (cat_id, data) in other.cat_data {
match self.cat_data.entry(cat_id) {
Entry::Vacant(v) => {
v.insert(data);
}
Entry::Occupied(o) => o.into_mut().extend(data),
}
}
for (slot, n) in self.fp_counts.iter_mut().zip(other.fp_counts) {
*slot += n;
}
self.covered_gts.extend(other.covered_gts);
self
}
}
#[derive(Default)]
struct MissCounts {
total: u64,
per_cat: HashMap<u64, usize>,
fn_per_cat: HashMap<u64, usize>,
}
impl MissCounts {
fn merge(mut self, other: Self) -> Self {
self.total += other.total;
for (cat_id, n) in other.per_cat {
*self.per_cat.entry(cat_id).or_insert(0) += n;
}
for (cat_id, n) in other.fn_per_cat {
*self.fn_per_cat.entry(cat_id).or_insert(0) += n;
}
self
}
}
#[derive(Debug, Default)]
struct FixScratch {
scores: Vec<f64>,
matched: Vec<bool>,
ignored: Vec<bool>,
}
struct CatDeltas {
baseline: f64,
fp_types: [f64; FP_TYPES.len()],
miss: f64,
fp: f64,
fn_oracle: f64,
}
impl COCOeval {
pub(super) fn compute_ap_from_matched(
scores: &[f64],
matched: &[bool],
ignored: &[bool],
num_gt: usize,
rec_thrs: &[f64],
) -> f64 {
crate::metrics::counts::average_precision(scores, matched, Some(ignored), num_gt, rec_thrs)
}
pub fn tide_errors(&self, pos_thr: f64, bg_thr: f64) -> crate::error::Result<TideErrors> {
if self.eval_imgs.is_empty() {
return Err("tide_errors() requires evaluate() to be called first".into());
}
let t_idx = self.params.nearest_iou_thr_idx(pos_thr);
let cells: Vec<&EvalImg> = self.default_cells().collect();
let cross_iou_map = self.cross_category_ious();
let classified = self.classify_detections(&cells, &cross_iou_map, t_idx, pos_thr, bg_thr);
let misses = count_misses(&cells, &classified.covered_gts, t_idx);
let per_cat = self.category_deltas(&classified.cat_data, &misses);
Ok(assemble(
&per_cat,
classified.fp_counts,
misses.total,
pos_thr,
bg_thr,
))
}
fn cross_category_ious(&self) -> CrossIouMap {
let cat_slots = Self::cat_slots(&self.params.cat_ids);
let iou_type = self.params.iou_type;
let coco_gt = &self.coco_gt;
let coco_dt = &self.coco_dt;
let segm_rles = self.segm_rles.as_ref();
self.params
.img_ids
.par_iter()
.map(|&img_id| {
let mut dt_max_cross: HashMap<u64, (f64, Option<u64>)> = HashMap::new();
let (gt_pairs, dt_pairs) =
Self::cross_category_pairs(coco_gt, coco_dt, &cat_slots, img_id, None);
if dt_pairs.is_empty() || gt_pairs.is_empty() {
for &(_, ann_id) in &dt_pairs {
dt_max_cross.insert(ann_id, (0.0, None));
}
return (img_id, dt_max_cross);
}
let dt_ids: Vec<u64> = dt_pairs.iter().map(|&(_, ann_id)| ann_id).collect();
let gt_ids: Vec<u64> = gt_pairs.iter().map(|&(_, ann_id)| ann_id).collect();
let iou_matrix = Self::cross_category_iou(
&dt_ids, >_ids, coco_dt, coco_gt, iou_type, segm_rles,
);
for (di, &(dt_cat_idx, dt_ann_id)) in dt_pairs.iter().enumerate() {
let row = &iou_matrix[di * gt_pairs.len()..(di + 1) * gt_pairs.len()];
let mut max_cross = 0.0f64;
let mut argmax_cross_gt = None;
for (gi, &(gt_cat_idx, gt_ann_id)) in gt_pairs.iter().enumerate() {
if gt_cat_idx != dt_cat_idx && row[gi] > max_cross {
max_cross = row[gi];
argmax_cross_gt = Some(gt_ann_id);
}
}
dt_max_cross.insert(dt_ann_id, (max_cross, argmax_cross_gt));
}
(img_id, dt_max_cross)
})
.collect()
}
fn classify_detections(
&self,
cells: &[&EvalImg],
cross_iou_map: &CrossIouMap,
t_idx: usize,
pos_thr: f64,
bg_thr: f64,
) -> Classified {
let mut classified = cells
.par_iter()
.fold(Classified::default, |mut acc, eval_img| {
let img_id = eval_img.image_id;
let cat_id = eval_img.category_id;
let dt_orig_ids = self.coco_dt.get_ann_ids_for_img_cat(img_id, cat_id);
let gt_orig_ids = self.coco_gt.get_ann_ids_for_img_cat(img_id, cat_id);
let orig_pos = |ids: &[u64], id: u64| ids.iter().position(|&x| x == id);
let gt_sorted_to_orig: Vec<Option<usize>> = eval_img
.gt_ids
.iter()
.map(|&id| orig_pos(gt_orig_ids, id))
.collect();
let same_iou_mat = self.cell_ious(img_id, cat_id);
let cross_map = cross_iou_map.get(&img_id);
let entry = acc.cat_data.entry(cat_id).or_insert_with(CatData::new);
entry.num_gt += eval_img.num_gt_in_denominator();
let covered_gts = &mut acc.covered_gts;
let fp_counts = &mut acc.fp_counts;
for (di, &dt_ann_id) in eval_img.dt_ids.iter().enumerate() {
let is_matched = eval_img.dt_matched[(t_idx, di)];
let is_ignored = eval_img.dt_ignore[(t_idx, di)];
let fp_type = (!is_matched && !is_ignored).then(|| {
let (max_cross_iou, argmax_cross_gt) = cross_map
.and_then(|m| m.get(&dt_ann_id))
.copied()
.unwrap_or((0.0, None));
let row = same_iou_mat
.zip(orig_pos(dt_orig_ids, dt_ann_id))
.and_then(|(mat, di_orig)| mat.get(di_orig))
.map_or(&[][..], Vec::as_slice);
let same =
same_class_scan(row, >_sorted_to_orig, eval_img, t_idx, pos_thr);
let err = classify_fp(
&FpEvidence {
max_same_iou: same.max_iou,
max_cross_iou,
best_same_gt_matched: same.best_gt_matched,
},
pos_thr,
bg_thr,
);
let target = match err {
ErrType::Loc => same.argmax_gt_ann_id,
ErrType::Cls => argmax_cross_gt,
ErrType::Both | ErrType::Dupe | ErrType::Bkg => None,
};
if let Some(gt_ann_id) = target {
covered_gts.insert(gt_ann_id);
}
fp_counts[err as usize] += 1;
err
});
entry.scores.push(eval_img.dt_scores[di]);
entry.matched.push(is_matched);
entry.ignored.push(is_ignored);
entry.fp_types.push(fp_type);
}
acc
})
.reduce_with(Classified::merge)
.unwrap_or_default();
for data in classified.cat_data.values_mut() {
data.rank_by_score_desc();
}
classified
}
fn category_deltas(
&self,
cat_data: &HashMap<u64, CatData>,
misses: &MissCounts,
) -> Vec<CatDeltas> {
let rec_thrs = &self.params.rec_thrs;
let per_cat: Vec<Option<CatDeltas>> = self
.params
.cat_ids
.par_iter()
.map(|&cat_id| {
let data = match cat_data.get(&cat_id) {
Some(d) if d.num_gt > 0 => d,
_ => return None,
};
let mut ap_scratch = ApScratch::default();
let mut fix = FixScratch::default();
let baseline = average_precision_ranked_into(
&data.matched,
Some(&data.ignored),
data.num_gt,
rec_thrs,
&mut ap_scratch,
);
let mut fix_fp = |fix_type: ErrType| -> f64 {
fix.matched.clear();
fix.matched.extend_from_slice(&data.matched);
fix.ignored.clear();
fix.ignored.extend_from_slice(&data.ignored);
for (i, fp_type) in data.fp_types.iter().enumerate() {
if *fp_type != Some(fix_type) {
continue;
}
match fix_type {
ErrType::Cls | ErrType::Loc => fix.matched[i] = true,
ErrType::Bkg | ErrType::Both | ErrType::Dupe => fix.ignored[i] = true,
}
}
average_precision_ranked_into(
&fix.matched,
Some(&fix.ignored),
data.num_gt,
rec_thrs,
&mut ap_scratch,
)
};
let mut fp_types = [0.0f64; FP_TYPES.len()];
for (slot, &err) in fp_types.iter_mut().zip(FP_TYPES.iter()) {
*slot = fix_fp(err) - baseline;
}
let fp = {
fix.ignored.clear();
fix.ignored.extend(
data.ignored
.iter()
.zip(&data.fp_types)
.map(|(&ig, fp_type)| ig || fp_type.is_some()),
);
average_precision_ranked_into(
&data.matched,
Some(&fix.ignored),
data.num_gt,
rec_thrs,
&mut ap_scratch,
) - baseline
};
let fn_count = misses.fn_per_cat.get(&cat_id).copied().unwrap_or(0);
debug_assert!(
fn_count <= data.num_gt,
"FN count exceeds the GT denominator it was counted from"
);
let fn_oracle = average_precision_ranked_into(
&data.matched,
Some(&data.ignored),
data.num_gt.saturating_sub(fn_count),
rec_thrs,
&mut ap_scratch,
) - baseline;
let miss_count = misses.per_cat.get(&cat_id).copied().unwrap_or(0);
let miss = if miss_count > 0 {
fix.scores.clear();
fix.matched.clear();
fix.ignored.clear();
fix.scores.resize(miss_count, 2.0);
fix.matched.resize(miss_count, true);
fix.ignored.resize(miss_count, false);
fix.scores.extend_from_slice(&data.scores);
fix.matched.extend_from_slice(&data.matched);
fix.ignored.extend_from_slice(&data.ignored);
Self::compute_ap_from_matched(
&fix.scores,
&fix.matched,
&fix.ignored,
data.num_gt,
rec_thrs,
) - baseline
} else {
0.0
};
Some(CatDeltas {
baseline,
fp_types,
miss,
fp,
fn_oracle,
})
})
.collect();
per_cat.into_iter().flatten().collect()
}
}
fn same_class_scan(
row: &[f64],
gt_sorted_to_orig: &[Option<usize>],
eval_img: &EvalImg,
t_idx: usize,
pos_thr: f64,
) -> SameClassScan {
let mut scan = SameClassScan::default();
for (gi_sorted, &gi_orig) in gt_sorted_to_orig.iter().enumerate() {
let Some(gi_orig) = gi_orig else {
continue;
};
let iou = row.get(gi_orig).copied().unwrap_or(0.0);
if iou > scan.max_iou {
scan.max_iou = iou;
scan.argmax_gt_ann_id = Some(eval_img.gt_ids[gi_sorted]);
}
if iou >= pos_thr && eval_img.gt_matched[(t_idx, gi_sorted)] {
scan.best_gt_matched = true;
}
}
scan
}
fn count_misses(cells: &[&EvalImg], covered_gts: &HashSet<u64>, t_idx: usize) -> MissCounts {
cells
.par_iter()
.fold(MissCounts::default, |mut counts, eval_img| {
let (mut n_miss, mut n_fn) = (0usize, 0usize);
for (gi, >_id) in eval_img.gt_ids.iter().enumerate() {
if eval_img.gt_matched[(t_idx, gi)] || !eval_img.counts_as_miss(gi) {
continue;
}
n_fn += 1;
if !covered_gts.contains(>_id) {
n_miss += 1;
}
}
counts.total += n_miss as u64;
if n_miss > 0 {
*counts.per_cat.entry(eval_img.category_id).or_insert(0) += n_miss;
}
if n_fn > 0 {
*counts.fn_per_cat.entry(eval_img.category_id).or_insert(0) += n_fn;
}
counts
})
.reduce_with(MissCounts::merge)
.unwrap_or_default()
}
fn assemble(
per_cat: &[CatDeltas],
fp_counts: [u64; FP_TYPES.len()],
miss_total: u64,
pos_thr: f64,
bg_thr: f64,
) -> TideErrors {
let mean = |f: &dyn Fn(&CatDeltas) -> f64| -> f64 {
if per_cat.is_empty() {
0.0
} else {
per_cat.iter().map(f).sum::<f64>() / per_cat.len() as f64
}
};
let mut delta_ap: BTreeMap<String, f64> = FP_TYPES
.iter()
.enumerate()
.map(|(i, err)| (err.as_str().to_string(), mean(&|c| c.fp_types[i])))
.collect();
delta_ap.insert("Miss".to_string(), mean(&|c| c.miss));
delta_ap.insert("FP".to_string(), mean(&|c| c.fp));
delta_ap.insert("FN".to_string(), mean(&|c| c.fn_oracle));
let mut counts: BTreeMap<String, u64> = FP_TYPES
.iter()
.zip(fp_counts)
.map(|(err, n)| (err.as_str().to_string(), n))
.collect();
counts.insert("Miss".to_string(), miss_total);
TideErrors {
delta_ap,
counts,
ap_base: mean(&|c| c.baseline),
pos_thr,
bg_thr,
}
}
#[cfg(test)]
mod tests {
use super::*;
const POS: f64 = 0.5;
const BG: f64 = 0.1;
fn ev(max_same_iou: f64, max_cross_iou: f64, best_same_gt_matched: bool) -> FpEvidence {
FpEvidence {
max_same_iou,
max_cross_iou,
best_same_gt_matched,
}
}
fn classify(same: f64, cross: f64, dupe: bool) -> ErrType {
classify_fp(&ev(same, cross, dupe), POS, BG)
}
#[test]
fn each_error_type_is_reachable() {
assert_eq!(classify(0.3, 0.0, false), ErrType::Loc);
assert_eq!(classify(0.0, 0.9, false), ErrType::Cls);
assert_eq!(classify(0.9, 0.0, true), ErrType::Dupe);
assert_eq!(classify(0.0, 0.0, false), ErrType::Bkg);
assert_eq!(classify(0.0, 0.3, false), ErrType::Both);
}
#[test]
fn loc_outranks_cls_and_dupe() {
assert_eq!(classify(0.3, 0.9, false), ErrType::Loc);
assert_eq!(classify(0.3, 0.0, true), ErrType::Loc);
}
#[test]
fn cls_outranks_dupe_and_bkg() {
assert_eq!(classify(0.9, 0.9, true), ErrType::Cls);
assert_eq!(classify(0.0, 0.5, false), ErrType::Cls);
}
#[test]
fn dupe_outranks_bkg() {
assert_eq!(classify(0.9, 0.0, true), ErrType::Dupe);
assert_eq!(classify(0.9, 0.0, false), ErrType::Bkg);
}
#[test]
fn loc_window_is_inclusive_at_both_ends() {
assert_eq!(classify(BG, 0.0, false), ErrType::Loc);
assert_eq!(classify(POS, 0.0, true), ErrType::Loc);
assert_ne!(classify(BG - 1e-9, 0.0, false), ErrType::Loc);
assert_ne!(classify(POS + 1e-9, 0.0, true), ErrType::Loc);
}
#[test]
fn bkg_upper_bound_is_inclusive() {
assert_eq!(classify(0.0, BG, false), ErrType::Bkg);
assert_eq!(classify(0.0, BG + 1e-9, false), ErrType::Both);
}
#[test]
fn as_str_covers_every_variant() {
for (err, key) in [
(ErrType::Cls, "Cls"),
(ErrType::Loc, "Loc"),
(ErrType::Both, "Both"),
(ErrType::Dupe, "Dupe"),
(ErrType::Bkg, "Bkg"),
] {
assert_eq!(err.as_str(), key);
}
}
#[test]
fn fp_types_are_indexed_by_discriminant() {
for (i, &err) in FP_TYPES.iter().enumerate() {
assert_eq!(err as usize, i, "{} is out of order", err.as_str());
}
}
}
#[derive(Debug, Clone, Serialize)]
pub struct TideErrors {
pub delta_ap: BTreeMap<String, f64>,
pub counts: BTreeMap<String, u64>,
pub ap_base: f64,
pub pos_thr: f64,
pub bg_thr: f64,
}