pub struct GreedyMatches {
pub dt_gt: Vec<Vec<Option<usize>>>,
pub gt_matched: Vec<Vec<bool>>,
}
pub fn greedy_match(
iou_flat: &[f64],
d: usize,
g: usize,
num_gt_not_ignored: usize,
gt_rematchable: &[bool],
gt_phase2_eligible: &[bool],
iou_thrs: &[f64],
) -> GreedyMatches {
let t = iou_thrs.len();
let mut dt_gt = vec![vec![None; d]; t];
let mut gt_matched = vec![vec![false; g]; t];
for (ti, &iou_thr) in iou_thrs.iter().enumerate() {
for (di, dt_slot) in dt_gt[ti].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_matched[ti][gi] && !gt_rematchable[gi] {
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 !gt_phase2_eligible[gi] {
continue;
}
if gt_matched[ti][gi] && !gt_rematchable[gi] {
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_matched[ti][gi] = true;
}
}
}
GreedyMatches { dt_gt, gt_matched }
}
#[cfg(test)]
mod tests {
use super::*;
fn simple(iou_flat: &[f64], d: usize, g: usize, num_ni: usize, thrs: &[f64]) -> GreedyMatches {
greedy_match(
iou_flat,
d,
g,
num_ni,
&vec![false; g],
&vec![true; g],
thrs,
)
}
#[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[0], vec![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(&[0.9, 0.8], 2, 1, 0, &[true], &[true], &[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(&[0.9, 0.8], 2, 1, 0, &[false], &[true], &[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(
&[0.4, 0.9],
1,
2,
1,
&[false, false],
&[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));
}
}