use std::collections::BTreeMap;
use rayon::prelude::*;
use av_core::geometry::Aabb;
use av_core::types::Detection;
pub const IOU_THRESHOLDS: [f64; 10] = [0.50, 0.55, 0.60, 0.65, 0.70, 0.75, 0.80, 0.85, 0.90, 0.95];
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct GtBox {
pub bbox: Aabb,
pub class_id: u32,
}
impl GtBox {
pub fn new(bbox: Aabb, class_id: u32) -> Self {
Self { bbox, class_id }
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct MapReport {
pub map50_95: f32,
pub map50: f32,
pub num_classes: usize,
pub num_images: usize,
pub num_gts: usize,
pub num_dets: usize,
}
#[derive(Debug)]
struct ImageRec {
dets: Vec<Detection>,
gts: Vec<GtBox>,
}
#[derive(Debug, Default)]
pub struct CocoEvaluator {
images: BTreeMap<u32, ImageRec>,
}
impl CocoEvaluator {
pub fn new() -> Self {
Self::default()
}
pub fn update(&mut self, image_id: u32, dets: &[Detection], gts: &[GtBox]) {
self.images.insert(
image_id,
ImageRec {
dets: dets.to_vec(),
gts: gts.to_vec(),
},
);
}
pub fn finalize(&self) -> MapReport {
let num_gts: usize = self.images.values().map(|r| r.gts.len()).sum();
let num_dets: usize = self.images.values().map(|r| r.dets.len()).sum();
let mut classes: Vec<u32> = self
.images
.values()
.flat_map(|r| r.gts.iter().map(|g| g.class_id))
.collect();
classes.sort_unstable();
classes.dedup();
let report = MapReport {
map50_95: 0.0,
map50: 0.0,
num_classes: classes.len(),
num_images: self.images.len(),
num_gts,
num_dets,
};
if classes.is_empty() {
return report;
}
let per_class: Vec<[f64; 10]> = classes
.par_iter()
.map(|&c| {
let ctx = self.class_context(c);
let mut aps = [0.0f64; 10];
for (ti, &thr) in IOU_THRESHOLDS.iter().enumerate() {
aps[ti] = class_ap_at(&ctx, thr);
}
aps
})
.collect();
let mut sum_all = 0.0f64; let mut sum_50 = 0.0f64; for aps in &per_class {
for (ti, &ap) in aps.iter().enumerate() {
sum_all += ap;
if ti == 0 {
sum_50 += ap;
}
}
}
let n = classes.len() as f64;
MapReport {
map50_95: (sum_all / (10.0 * n)) as f32,
map50: (sum_50 / n) as f32,
..report
}
}
fn class_context(&self, cls: u32) -> ClassContext {
let mut gt_boxes: Vec<Vec<Aabb>> = Vec::with_capacity(self.images.len());
let mut dets: Vec<(usize, f64, Aabb)> = Vec::new();
let mut n_gt = 0usize;
for (idx, rec) in self.images.values().enumerate() {
let g: Vec<Aabb> = rec
.gts
.iter()
.filter(|gt| gt.class_id == cls)
.map(|gt| gt.bbox)
.collect();
n_gt += g.len();
gt_boxes.push(g);
for d in &rec.dets {
if d.class_id == cls {
dets.push((idx, d.score as f64, d.bbox));
}
}
}
dets.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let ious: Vec<Vec<f64>> = dets
.iter()
.map(|(img, _, bbox)| gt_boxes[*img].iter().map(|g| g.iou(bbox) as f64).collect())
.collect();
ClassContext {
gt_boxes,
dets_img: dets.iter().map(|(img, _, _)| *img).collect(),
ious,
n_gt,
}
}
}
struct ClassContext {
gt_boxes: Vec<Vec<Aabb>>,
dets_img: Vec<usize>,
ious: Vec<Vec<f64>>,
n_gt: usize,
}
fn class_ap_at(ctx: &ClassContext, thr: f64) -> f64 {
if ctx.n_gt == 0 {
return 0.0;
}
let mut gt_matched: Vec<Vec<bool>> =
ctx.gt_boxes.iter().map(|g| vec![false; g.len()]).collect();
let mut rec = Vec::with_capacity(ctx.dets_img.len());
let mut prec = Vec::with_capacity(ctx.dets_img.len());
let (mut tp, mut fp) = (0usize, 0usize);
for (di, &img) in ctx.dets_img.iter().enumerate() {
let mut best = (0.0f64, usize::MAX); for (gi, _) in ctx.gt_boxes[img].iter().enumerate() {
if gt_matched[img][gi] {
continue;
}
let iou = ctx.ious[di][gi];
if iou >= thr && iou > best.0 {
best = (iou, gi);
}
}
if best.1 != usize::MAX {
gt_matched[img][best.1] = true;
tp += 1;
} else {
fp += 1;
}
rec.push(tp as f64 / ctx.n_gt as f64);
prec.push(tp as f64 / (tp + fp) as f64);
}
ap_101(&rec, &prec)
}
fn ap_101(rec: &[f64], prec: &[f64]) -> f64 {
let mut env = vec![0.0f64; prec.len() + 1];
for i in (0..prec.len()).rev() {
env[i] = env[i + 1].max(prec[i]);
}
let mut ap = 0.0f64;
let mut ri = 0usize; for k in 0..=100 {
let r = k as f64 / 100.0;
while ri < rec.len() && rec[ri] < r {
ri += 1;
}
ap += env[ri];
}
ap / 101.0
}
#[cfg(test)]
mod tests {
use super::*;
fn det(x1: f32, y1: f32, x2: f32, y2: f32, score: f32, class_id: u32) -> Detection {
Detection {
bbox: Aabb::new(x1, y1, x2, y2),
score,
class_id,
angle: None,
keypoints: None,
}
}
#[test]
fn hand_computed_two_images_two_classes() {
let mut ev = CocoEvaluator::new();
ev.update(
1,
&[
det(30.0, 30.0, 40.0, 40.0, 0.95, 1), det(0.0, 0.0, 10.0, 10.0, 0.9, 0), det(0.0, 0.0, 10.0, 20.0, 0.8, 1), ],
&[
GtBox::new(Aabb::new(0.0, 0.0, 10.0, 10.0), 0),
GtBox::new(Aabb::new(0.0, 0.0, 20.0, 20.0), 1),
GtBox::new(Aabb::new(30.0, 30.0, 40.0, 40.0), 1),
],
);
ev.update(
2,
&[
det(5.0, 5.0, 15.0, 15.0, 0.7, 0), det(100.0, 100.0, 110.0, 110.0, 0.6, 0), ],
&[GtBox::new(Aabb::new(5.0, 5.0, 15.0, 15.0), 0)],
);
let rep = ev.finalize();
assert_eq!(rep.num_images, 2);
assert_eq!(rep.num_gts, 4);
assert_eq!(rep.num_dets, 5);
assert_eq!(rep.num_classes, 2);
assert!((rep.map50 - 1.0).abs() < 1e-3, "mAP50={}", rep.map50);
let expect = (11.0 + 9.0 * 51.0 / 101.0) / 20.0;
assert!(
(rep.map50_95 - expect as f32).abs() < 1e-3,
"mAP50:95={} 期望 {expect}",
rep.map50_95
);
}
#[test]
fn empty_or_gtless_is_zero_not_nan() {
let rep = CocoEvaluator::new().finalize();
assert_eq!(rep.map50, 0.0);
assert_eq!(rep.map50_95, 0.0);
assert!(!rep.map50.is_nan() && !rep.map50_95.is_nan());
let mut ev = CocoEvaluator::new();
ev.update(
7,
&[det(0.0, 0.0, 10.0, 10.0, 0.9, 0)],
&[], );
let rep = ev.finalize();
assert_eq!(rep.num_classes, 0);
assert_eq!(rep.num_images, 1);
assert_eq!(rep.num_dets, 1);
assert_eq!(rep.map50, 0.0);
assert!(!rep.map50_95.is_nan());
}
#[test]
fn unmatched_gt_gives_zero_ap() {
let mut ev = CocoEvaluator::new();
ev.update(
1,
&[det(50.0, 50.0, 60.0, 60.0, 0.9, 0)],
&[GtBox::new(Aabb::new(0.0, 0.0, 10.0, 10.0), 0)],
);
let rep = ev.finalize();
assert_eq!(rep.num_classes, 1);
assert_eq!(rep.map50, 0.0);
assert_eq!(rep.map50_95, 0.0);
}
#[test]
fn greedy_takes_best_iou_among_unmatched() {
let mut ev = CocoEvaluator::new();
ev.update(
1,
&[
det(0.0, 0.0, 10.0, 10.0, 0.8, 0),
det(0.0, 0.0, 10.0, 10.0, 0.6, 0),
],
&[GtBox::new(Aabb::new(0.0, 0.0, 10.0, 10.0), 0)],
);
let rep = ev.finalize();
assert!((rep.map50 - 1.0).abs() < 1e-3, "mAP50={}", rep.map50);
}
}