use std::collections::HashMap;
use rayon::prelude::*;
use crate::coco::COCO;
use crate::params::Params;
use crate::primitives::sim::{self, SimKind};
use crate::types::Rle;
use super::{COCOeval, EvalMode};
pub(super) struct SegmRles {
gt: HashMap<u64, Rle>,
dt: HashMap<u64, Rle>,
}
impl SegmRles {
pub(super) fn gt_rle_or_convert(cache: Option<&Self>, coco_gt: &COCO, id: u64) -> Option<Rle> {
if let Some(rle) = cache.and_then(|c| c.gt.get(&id)) {
return Some(rle.clone());
}
coco_gt.ann_to_rle(coco_gt.get_ann(id)?)
}
pub(super) fn dt_rle_or_convert(cache: Option<&Self>, coco_dt: &COCO, id: u64) -> Option<Rle> {
if let Some(rle) = cache.and_then(|c| c.dt.get(&id)) {
return Some(rle.clone());
}
coco_dt.ann_to_rle(coco_dt.get_ann(id)?)
}
pub(super) fn prepare(coco_gt: &COCO, coco_dt: &COCO, params: &Params) -> Self {
let cat_ids: &[u64] = if params.use_cats {
¶ms.cat_ids
} else {
&[]
};
let convert = |coco: &COCO| -> HashMap<u64, Rle> {
coco.get_ann_ids(¶ms.img_ids, cat_ids, None, None)
.into_par_iter()
.filter_map(|id| Some((id, coco.ann_to_rle(coco.get_ann(id)?)?)))
.collect()
};
SegmRles {
gt: convert(coco_gt),
dt: convert(coco_dt),
}
}
}
fn uses_ioa(ann: &crate::types::Annotation, eval_mode: EvalMode) -> bool {
if eval_mode == EvalMode::OpenImages {
ann.is_group_of.unwrap_or(false)
} else {
ann.iscrowd
}
}
fn scatter_full(
valid: Vec<Vec<f64>>,
dt_rows: &[usize],
gt_cols: &[usize],
d: usize,
g: usize,
) -> Vec<Vec<f64>> {
if dt_rows.len() == d && gt_cols.len() == g {
return valid;
}
let mut full = vec![vec![0.0_f64; g]; d];
for (vi, &di) in dt_rows.iter().enumerate() {
for (vj, &gj) in gt_cols.iter().enumerate() {
full[di][gj] = valid[vi][vj];
}
}
full
}
fn iou_scaffold<D, G>(
coco_gt: &COCO,
dt_ids: &[u64],
gt_ids: &[u64],
eval_mode: EvalMode,
dt_geom: impl Fn(u64) -> Option<D>,
gt_geom: impl Fn(&crate::types::Annotation, u64) -> Option<G>,
kernel: impl FnOnce(&[D], &[G], &[bool]) -> Vec<Vec<f64>>,
) -> Vec<Vec<f64>> {
let (dt_rows, dt_geoms): (Vec<usize>, Vec<D>) = dt_ids
.iter()
.enumerate()
.filter_map(|(idx, &id)| Some((idx, dt_geom(id)?)))
.unzip();
let mut gt_cols = Vec::with_capacity(gt_ids.len());
let mut gt_geoms = Vec::with_capacity(gt_ids.len());
let mut iscrowd = Vec::with_capacity(gt_ids.len());
for (idx, &id) in gt_ids.iter().enumerate() {
let Some(ann) = coco_gt.get_ann(id) else {
continue;
};
let Some(geom) = gt_geom(ann, id) else {
continue;
};
gt_cols.push(idx);
gt_geoms.push(geom);
iscrowd.push(uses_ioa(ann, eval_mode));
}
let valid = kernel(&dt_geoms, >_geoms, &iscrowd);
scatter_full(valid, &dt_rows, >_cols, dt_ids.len(), gt_ids.len())
}
impl COCOeval {
pub(super) fn compute_iou_static(
coco_gt: &COCO,
coco_dt: &COCO,
params: &Params,
img_id: u64,
cat_id: u64,
eval_mode: EvalMode,
segm_rles: Option<&SegmRles>,
) -> Vec<Vec<f64>> {
let gt_anns = Self::get_anns_static(coco_gt, params, img_id, cat_id);
let dt_anns = Self::get_anns_static(coco_dt, params, img_id, cat_id);
if gt_anns.is_empty() || dt_anns.is_empty() {
return Vec::new();
}
match SimKind::from(params.iou_type) {
SimKind::Mask => Self::compute_segm_iou_static(
coco_gt, coco_dt, dt_anns, gt_anns, eval_mode, segm_rles,
),
SimKind::Bbox => {
Self::compute_bbox_iou_static(coco_gt, coco_dt, dt_anns, gt_anns, eval_mode)
}
SimKind::Oks => Self::compute_oks_static(coco_gt, coco_dt, params, dt_anns, gt_anns),
SimKind::Obb => {
Self::compute_obb_iou_static(coco_gt, coco_dt, dt_anns, gt_anns, eval_mode)
}
}
}
pub(super) fn get_anns_static<'a>(
coco: &'a COCO,
params: &Params,
img_id: u64,
cat_id: u64,
) -> &'a [u64] {
if params.use_cats {
coco.get_ann_ids_for_img_cat(img_id, cat_id)
} else {
coco.get_ann_ids_for_img(img_id)
}
}
pub(super) fn compute_segm_iou_static(
coco_gt: &COCO,
coco_dt: &COCO,
dt_ids: &[u64],
gt_ids: &[u64],
eval_mode: EvalMode,
segm_rles: Option<&SegmRles>,
) -> Vec<Vec<f64>> {
iou_scaffold(
coco_gt,
dt_ids,
gt_ids,
eval_mode,
|id| SegmRles::dt_rle_or_convert(segm_rles, coco_dt, id),
|_ann, id| SegmRles::gt_rle_or_convert(segm_rles, coco_gt, id),
sim::mask_iou,
)
}
pub(super) fn compute_bbox_iou_static(
coco_gt: &COCO,
coco_dt: &COCO,
dt_ids: &[u64],
gt_ids: &[u64],
eval_mode: EvalMode,
) -> Vec<Vec<f64>> {
iou_scaffold(
coco_gt,
dt_ids,
gt_ids,
eval_mode,
|id| coco_dt.get_ann(id)?.bbox,
|ann, _id| ann.bbox,
sim::bbox_iou,
)
}
pub(super) fn compute_oks_static(
coco_gt: &COCO,
coco_dt: &COCO,
params: &Params,
dt_ids: &[u64],
gt_ids: &[u64],
) -> Vec<Vec<f64>> {
let (gt_cols, gt_anns): (Vec<usize>, Vec<_>) = gt_ids
.iter()
.enumerate()
.filter_map(|(idx, &id)| Some((idx, coco_gt.get_ann(id)?)))
.unzip();
let (dt_rows, dt_anns): (Vec<usize>, Vec<_>) = dt_ids
.iter()
.enumerate()
.filter_map(|(idx, &id)| Some((idx, coco_dt.get_ann(id)?)))
.unzip();
let gt: Vec<crate::primitives::sim::GtPose<'_>> = gt_anns
.iter()
.map(|a| crate::primitives::sim::GtPose {
keypoints: a.keypoints.as_deref().unwrap_or(&[]),
area: a.area.unwrap_or(0.0),
bbox: a.bbox.unwrap_or([0.0; 4]),
})
.collect();
let dt_keypoints: Vec<&[f64]> = dt_anns
.iter()
.map(|a| a.keypoints.as_deref().unwrap_or(&[]))
.collect();
let valid = crate::primitives::sim::oks_matrix(&dt_keypoints, >, ¶ms.kpt_oks_sigmas);
scatter_full(valid, &dt_rows, >_cols, dt_ids.len(), gt_ids.len())
}
pub(super) fn compute_obb_iou_static(
coco_gt: &COCO,
coco_dt: &COCO,
dt_ids: &[u64],
gt_ids: &[u64],
eval_mode: EvalMode,
) -> Vec<Vec<f64>> {
iou_scaffold(
coco_gt,
dt_ids,
gt_ids,
eval_mode,
|id| coco_dt.get_ann(id)?.obb,
|ann, _id| ann.obb,
sim::obb_iou,
)
}
}