#[inline]
pub fn coco_match_floor(iou_thr: f64) -> f64 {
iou_thr.min(1.0 - 1e-10)
}
pub fn best_above_floor(sims: &[f64], eligible: &[bool], floor: f64) -> Option<usize> {
let mut best: Option<usize> = None;
let mut best_sim = f64::NEG_INFINITY;
for (i, &s) in sims.iter().enumerate() {
if eligible[i] && s >= floor && s > best_sim {
best = Some(i);
best_sim = s;
}
}
best
}
#[derive(Debug, Clone)]
pub struct ThreshMatrix<T> {
rows: usize,
row_len: usize,
data: Vec<T>,
}
impl<T: Clone> ThreshMatrix<T> {
pub fn new(rows: usize, row_len: usize, fill: T) -> Self {
Self {
rows,
row_len,
data: vec![fill; rows * row_len],
}
}
pub fn repeat_row(rows: usize, row: &[T]) -> Self {
let mut data = Vec::with_capacity(rows * row.len());
for _ in 0..rows {
data.extend_from_slice(row);
}
Self {
rows,
row_len: row.len(),
data,
}
}
}
impl<T> ThreshMatrix<T> {
pub fn num_rows(&self) -> usize {
self.rows
}
pub fn row_len(&self) -> usize {
self.row_len
}
pub fn row(&self, t: usize) -> &[T] {
&self.data[t * self.row_len..(t + 1) * self.row_len]
}
pub fn row_mut(&mut self, t: usize) -> &mut [T] {
&mut self.data[t * self.row_len..(t + 1) * self.row_len]
}
pub fn iter_rows(&self) -> impl Iterator<Item = &[T]> + '_ {
(0..self.rows).map(move |t| self.row(t))
}
}
impl<T> std::ops::Index<(usize, usize)> for ThreshMatrix<T> {
type Output = T;
fn index(&self, (t, i): (usize, usize)) -> &T {
debug_assert!(t < self.rows && i < self.row_len);
&self.data[t * self.row_len + i]
}
}
impl<T> std::ops::IndexMut<(usize, usize)> for ThreshMatrix<T> {
fn index_mut(&mut self, (t, i): (usize, usize)) -> &mut T {
debug_assert!(t < self.rows && i < self.row_len);
&mut self.data[t * self.row_len + i]
}
}
pub struct GreedyMatches {
pub dt_gt: ThreshMatrix<Option<usize>>,
pub gt_matched: ThreshMatrix<bool>,
}
#[derive(Debug, Clone, Copy, Default)]
pub struct GtMasks<'a> {
pub rematchable: Option<&'a [bool]>,
pub phase2_eligible: Option<&'a [bool]>,
}
pub fn greedy_match_masked(
iou_flat: &[f64],
d: usize,
g: usize,
num_gt_not_ignored: usize,
masks: GtMasks<'_>,
iou_thrs: &[f64],
) -> GreedyMatches {
assert_eq!(
iou_flat.len(),
d * g,
"greedy_match_masked: iou_flat must be d*g row-major ({d}x{g} = {}, got {})",
d * g,
iou_flat.len()
);
assert!(
num_gt_not_ignored <= g,
"greedy_match_masked: num_gt_not_ignored ({num_gt_not_ignored}) exceeds g ({g})"
);
for (name, mask) in [
("rematchable", masks.rematchable),
("phase2_eligible", masks.phase2_eligible),
] {
if let Some(m) = mask {
assert_eq!(
m.len(),
g,
"greedy_match_masked: {name} mask must have one entry per GT \
(got {}, expected g = {g})",
m.len()
);
}
}
let t = iou_thrs.len();
let mut dt_gt = ThreshMatrix::new(t, d, None);
let mut gt_matched = ThreshMatrix::new(t, g, false);
let rematchable = masks.rematchable.unwrap_or(&[]);
let phase2_eligible = masks.phase2_eligible.unwrap_or(&[]);
for (ti, &iou_thr) in iou_thrs.iter().enumerate() {
let dt_row = dt_gt.row_mut(ti);
let gt_row = gt_matched.row_mut(ti);
for (di, dt_slot) in dt_row.iter_mut().enumerate() {
let base = di * g;
let mut best_iou = iou_thr;
let mut best_gi: Option<usize> = None;
for gi in 0..num_gt_not_ignored {
if gt_row[gi] && !rematchable.get(gi).copied().unwrap_or(false) {
continue;
}
let iou_val = iou_flat[base + gi];
if iou_val >= best_iou {
best_iou = iou_val;
best_gi = Some(gi);
}
}
if best_gi.is_none() {
for gi in num_gt_not_ignored..g {
if !phase2_eligible.get(gi).copied().unwrap_or(true) {
continue;
}
if gt_row[gi] && !rematchable.get(gi).copied().unwrap_or(false) {
continue;
}
let iou_val = iou_flat[base + gi];
if iou_val >= best_iou {
best_iou = iou_val;
best_gi = Some(gi);
}
}
}
if let Some(gi) = best_gi {
*dt_slot = Some(gi);
gt_row[gi] = true;
}
}
}
GreedyMatches { dt_gt, gt_matched }
}
#[cfg(test)]
mod tests {
use super::*;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
#[test]
fn greedy_match_contract_random() {
let mut rng = StdRng::seed_from_u64(0x6DEED1);
for case in 0..5000 {
let d = rng.random_range(1..=6);
let g = rng.random_range(1..=6);
let num_ni = rng.random_range(0..=g);
let quantized = rng.random_bool(0.5);
let iou_flat: Vec<f64> = (0..d * g)
.map(|_| {
if quantized {
rng.random_range(0..=4) as f64 / 4.0
} else {
rng.random_range(0.0..=1.0)
}
})
.collect();
let rematchable: Vec<bool> = (0..g).map(|_| rng.random_bool(0.25)).collect();
let phase2: Vec<bool> = (0..g).map(|_| rng.random_bool(0.75)).collect();
let mut thrs: Vec<f64> = (0..rng.random_range(1..=4))
.map(|_| rng.random_range(0.0..=1.0))
.collect();
thrs.sort_by(f64::total_cmp);
let m = greedy_match_masked(
&iou_flat,
d,
g,
num_ni,
GtMasks {
rematchable: Some(&rematchable),
phase2_eligible: Some(&phase2),
},
&thrs,
);
let ctx = format!("case {case}: d={d} g={g} num_ni={num_ni} thrs={thrs:?}");
assert_eq!(m.dt_gt.num_rows(), thrs.len(), "{ctx}");
assert_eq!(m.gt_matched.num_rows(), thrs.len(), "{ctx}");
for (ti, &thr) in thrs.iter().enumerate() {
assert_eq!(m.dt_gt.row_len(), d, "{ctx}");
assert_eq!(m.gt_matched.row_len(), g, "{ctx}");
let mut claimed = vec![0usize; g];
for di in 0..d {
let Some(gi) = m.dt_gt[(ti, di)] else {
continue;
};
assert!(gi < g, "{ctx}: gt index {gi} out of range");
assert!(
iou_flat[di * g + gi] >= thr,
"{ctx}: dt {di} matched gt {gi} at IoU {} < {thr}",
iou_flat[di * g + gi]
);
if gi >= num_ni {
assert!(
phase2[gi],
"{ctx}: dt {di} matched phase-2-ineligible gt {gi}"
);
}
claimed[gi] += 1;
}
for gi in 0..g {
if claimed[gi] > 1 {
assert!(
rematchable[gi],
"{ctx}: gt {gi} claimed {} times but is not rematchable",
claimed[gi]
);
}
assert_eq!(
m.gt_matched[(ti, gi)],
claimed[gi] > 0,
"{ctx}: gt_matched[{gi}] disagrees with dt_gt"
);
}
}
}
}
#[test]
fn tp_eligible_matches_are_monotone_in_threshold() {
let mut rng = StdRng::seed_from_u64(0xA11CE);
for case in 0..20000 {
let d = rng.random_range(1..=5);
let g = rng.random_range(1..=5);
let num_ni = rng.random_range(0..=g);
let quantized = rng.random_bool(0.5);
let iou: Vec<f64> = (0..d * g)
.map(|_| {
if quantized {
rng.random_range(0..=4) as f64 / 4.0
} else {
rng.random_range(0.0..=1.0)
}
})
.collect();
let rematchable: Vec<bool> = (0..g).map(|_| rng.random_bool(0.2)).collect();
let phase2: Vec<bool> = (0..g).map(|_| rng.random_bool(0.8)).collect();
let mut thrs: Vec<f64> = (0..2).map(|_| rng.random_range(0.0..=1.0)).collect();
thrs.sort_by(f64::total_cmp);
let m = greedy_match_masked(
&iou,
d,
g,
num_ni,
GtMasks {
rematchable: Some(&rematchable),
phase2_eligible: Some(&phase2),
},
&thrs,
);
let tp_at = |ti: usize| {
m.dt_gt
.row(ti)
.iter()
.flatten()
.filter(|&&gi| gi < num_ni)
.count()
};
assert!(
tp_at(1) <= tp_at(0),
"case {case}: raising the threshold {:?} -> {:?} grew TP-eligible \
matches {} -> {} (d={d} g={g} num_ni={num_ni}) iou={iou:?}",
thrs[0],
thrs[1],
tp_at(0),
tp_at(1),
);
}
}
fn simple(iou_flat: &[f64], d: usize, g: usize, num_ni: usize, thrs: &[f64]) -> GreedyMatches {
greedy_match_masked(iou_flat, d, g, num_ni, GtMasks::default(), thrs)
}
#[test]
fn none_masks_equal_their_explicit_uniform_forms() {
#[rustfmt::skip]
let iou = [
0.9, 0.95, 0.2,
0.5, 0.99, 0.3,
0.1, 0.10, 0.9,
];
let (d, g, num_ni) = (3, 3, 2);
let thrs = [0.5, 0.85];
let implicit = greedy_match_masked(&iou, d, g, num_ni, GtMasks::default(), &thrs);
let explicit = greedy_match_masked(
&iou,
d,
g,
num_ni,
GtMasks {
rematchable: Some(&vec![false; g]),
phase2_eligible: Some(&vec![true; g]),
},
&thrs,
);
for ti in 0..thrs.len() {
assert_eq!(implicit.dt_gt.row(ti), explicit.dt_gt.row(ti));
assert_eq!(implicit.gt_matched.row(ti), explicit.gt_matched.row(ti));
}
}
#[test]
fn matches_highest_iou_above_threshold() {
let m = simple(&[0.6, 0.9], 1, 2, 2, &[0.5]);
assert_eq!(m.dt_gt[(0, 0)], Some(1));
assert_eq!(m.gt_matched.row(0), &[false, true]);
}
#[test]
fn below_threshold_is_no_match() {
let m = simple(&[0.4, 0.49], 1, 2, 2, &[0.5]);
assert_eq!(m.dt_gt[(0, 0)], None);
}
#[test]
fn score_order_gives_earlier_dt_first_pick() {
let m = simple(&[0.9, 0.8], 2, 1, 1, &[0.5]);
assert_eq!(m.dt_gt[(0, 0)], Some(0));
assert_eq!(m.dt_gt[(0, 1)], None);
}
#[test]
fn phase1_preferred_over_better_ignored_gt() {
let m = simple(&[0.6, 0.99], 1, 2, 1, &[0.5]);
assert_eq!(m.dt_gt[(0, 0)], Some(0));
}
#[test]
fn falls_back_to_ignored_gt_when_no_phase1_match() {
let m = simple(&[0.4, 0.8], 1, 2, 1, &[0.5]);
assert_eq!(m.dt_gt[(0, 0)], Some(1));
}
#[test]
fn crowd_gt_rematched_by_multiple_dts() {
let m = greedy_match_masked(
&[0.9, 0.8],
2,
1,
0,
GtMasks {
rematchable: Some(&[true]),
phase2_eligible: None,
},
&[0.5],
);
assert_eq!(m.dt_gt[(0, 0)], Some(0));
assert_eq!(m.dt_gt[(0, 1)], Some(0));
}
#[test]
fn non_rematchable_gt_taken_only_once() {
let m = greedy_match_masked(
&[0.9, 0.8],
2,
1,
0,
GtMasks {
rematchable: Some(&[false]),
phase2_eligible: None,
},
&[0.5],
);
assert_eq!(m.dt_gt[(0, 0)], Some(0));
assert_eq!(m.dt_gt[(0, 1)], None);
}
#[test]
fn phase2_ineligible_gt_is_skipped() {
let m = greedy_match_masked(
&[0.4, 0.9],
1,
2,
1,
GtMasks {
rematchable: None,
phase2_eligible: Some(&[true, false]),
},
&[0.5],
);
assert_eq!(m.dt_gt[(0, 0)], None);
}
#[test]
fn equal_iou_later_index_wins() {
let m = simple(&[0.7, 0.7], 1, 2, 2, &[0.5]);
assert_eq!(m.dt_gt[(0, 0)], Some(1));
}
#[test]
#[should_panic(expected = "iou_flat must be d*g")]
fn wrong_iou_matrix_length_panics() {
greedy_match_masked(&[0.9, 0.8, 0.7], 2, 2, 2, GtMasks::default(), &[0.5]);
}
#[test]
#[should_panic(expected = "num_gt_not_ignored")]
fn num_not_ignored_beyond_g_panics() {
greedy_match_masked(&[0.9], 1, 1, 2, GtMasks::default(), &[0.5]);
}
#[test]
#[should_panic(expected = "rematchable mask")]
fn short_rematchable_mask_panics() {
greedy_match_masked(
&[0.9, 0.8],
1,
2,
2,
GtMasks {
rematchable: Some(&[true]),
phase2_eligible: None,
},
&[0.5],
);
}
#[test]
#[should_panic(expected = "phase2_eligible mask")]
fn overlong_phase2_mask_panics() {
greedy_match_masked(
&[0.9, 0.8],
1,
2,
1,
GtMasks {
rematchable: None,
phase2_eligible: Some(&[true, true, false]),
},
&[0.5],
);
}
}