pub const DEFAULT_TOPK: usize = 10;
pub const DEFAULT_ALPHA: f32 = 0.5;
pub const DEFAULT_BETA: f32 = 6.0;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct TalConfig {
pub topk: usize,
pub alpha: f32,
pub beta: f32,
}
impl Default for TalConfig {
fn default() -> Self {
Self {
topk: DEFAULT_TOPK,
alpha: DEFAULT_ALPHA,
beta: DEFAULT_BETA,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct PosCell {
pub cell: usize,
pub gt: usize,
pub metric: f32,
pub weight: f32,
}
pub fn alignment_metric(cls_score: f32, iou: f32, alpha: f32, beta: f32) -> f32 {
cls_score.clamp(0.0, 1.0).powf(alpha) * iou.clamp(0.0, 1.0).powf(beta)
}
pub fn iou_xyxy(a: [f32; 4], b: [f32; 4]) -> f32 {
let iw = (a[2].min(b[2]) - a[0].max(b[0])).max(0.0);
let ih = (a[3].min(b[3]) - a[1].max(b[1])).max(0.0);
let inter = iw * ih;
let union = (a[2] - a[0]) * (a[3] - a[1]) + (b[2] - b[0]) * (b[3] - b[1]) - inter;
if union <= 0.0 {
0.0
} else {
(inter / union).clamp(0.0, 1.0)
}
}
fn center_in_box(c: [f32; 2], b: [f32; 4]) -> bool {
c[0] >= b[0] && c[0] <= b[2] && c[1] >= b[1] && c[1] <= b[3]
}
pub fn assign_single_image(
pred_boxes: &[[f32; 4]],
cell_centers: &[[f32; 2]],
gt_scores: &[&[f32]],
gt_boxes: &[[f32; 4]],
cfg: &TalConfig,
) -> Vec<Option<PosCell>> {
let n_cells = pred_boxes.len();
let mut out: Vec<Option<PosCell>> = vec![None; n_cells];
if n_cells == 0 || gt_boxes.is_empty() {
return out;
}
let k = cfg.topk.clamp(1, n_cells);
let mut best: Vec<Option<PosCell>> = vec![None; n_cells];
let mut cand: Vec<(f32, usize)> = Vec::new();
for (g, gt) in gt_boxes.iter().enumerate() {
let scores = gt_scores.get(g).copied().unwrap_or(&[]);
let score = |c: usize| scores.get(c).copied().unwrap_or(0.0);
cand.clear();
for (c, pb) in pred_boxes.iter().enumerate() {
if !center_in_box(cell_centers[c], *gt) {
continue;
}
let iou = iou_xyxy(*pb, *gt);
if iou <= 0.0 {
continue;
}
let m = alignment_metric(score(c), iou, cfg.alpha, cfg.beta);
cand.push((m, c));
}
if cand.is_empty() {
let gcx = (gt[0] + gt[2]) / 2.0;
let gcy = (gt[1] + gt[3]) / 2.0;
let mut best_c = 0usize;
let mut best_d = f32::INFINITY;
for (c, cc) in cell_centers.iter().enumerate() {
let d = (cc[0] - gcx).powi(2) + (cc[1] - gcy).powi(2);
if d < best_d {
best_d = d;
best_c = c;
}
}
let iou = iou_xyxy(pred_boxes[best_c], *gt);
let m = alignment_metric(score(best_c), iou, cfg.alpha, cfg.beta);
cand.push((m, best_c));
}
cand.sort_by(|a, b| {
b.0.partial_cmp(&a.0)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.1.cmp(&b.1))
});
cand.truncate(k);
let max_m = cand.first().map(|(m, _)| *m).unwrap_or(0.0);
for &(m, c) in &cand {
let w = if max_m > 1e-12 {
(m / max_m).clamp(0.0, 1.0)
} else {
1.0
};
let take = match &best[c] {
Some(prev) => m > prev.metric,
None => true,
};
if take {
best[c] = Some(PosCell {
cell: c,
gt: g,
metric: m,
weight: w,
});
}
}
}
for (c, b) in best.into_iter().enumerate() {
out[c] = b;
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn approx(a: f32, b: f32, tol: f32) -> bool {
(a - b).abs() <= tol
}
fn grid8() -> Vec<[f32; 2]> {
let mut v = Vec::new();
for hi in 0..2usize {
for wi in 0..4usize {
v.push([((2 * wi + 1) * 4) as f32, ((2 * hi + 1) * 4) as f32]);
}
}
v
}
#[test]
fn alignment_metric_hand_computed() {
let m = alignment_metric(0.8, 0.5, 0.5, 6.0);
assert!(approx(m, 0.8f32.sqrt() * 0.5f32.powi(6), 1e-6), "m={m}");
assert!(approx(m, 0.013975424, 1e-6), "m={m}");
}
#[test]
fn iou_hand_computed() {
assert!(approx(
iou_xyxy([0., 0., 2., 2.], [1., 1., 3., 3.]),
1.0 / 7.0,
1e-6
));
assert!(approx(
iou_xyxy([0., 0., 4., 4.], [0., 0., 4., 4.]),
1.0,
1e-6
));
assert_eq!(iou_xyxy([0., 0., 1., 1.], [2., 2., 3., 3.]), 0.0);
assert_eq!(iou_xyxy([0., 0., 0., 0.], [0., 0., 4., 4.]), 0.0);
}
#[test]
fn topk_weights_and_center_filter_hand_computed() {
let centers = grid8();
let gt = [[8.0f32, 0.0, 24.0, 16.0]];
let pred = [
[8., 0., 24., 16.], [8., 0., 24., 16.], [8., 0., 24., 16.], [8., 0., 24., 16.], [8., 0., 24., 16.], [12., 4., 20., 12.], [12., 4., 20., 12.], [8., 0., 24., 16.], ];
let scores: &[&[f32]] = &[&[0.99, 0.9, 0.8, 0.0, 0.0, 0.9, 0.8, 0.0]];
let out = assign_single_image(&pred, ¢ers, scores, >, &TalConfig::default());
for c in [0usize, 3, 4, 7] {
assert!(out[c].is_none(), "c{c} 不应入选");
}
let t1 = 0.9f32.sqrt() * 1.0f32.powi(6); let t2 = 0.8f32.sqrt() * 1.0f32.powi(6); let t5 = 0.9f32.sqrt() * 0.25f32.powi(6); let t6 = 0.8f32.sqrt() * 0.25f32.powi(6); let a = out[1].expect("c1 应入选");
assert_eq!(a.gt, 0);
assert!(approx(a.metric, t1, 1e-6));
assert!(approx(a.weight, 1.0, 1e-6), "最大 t 的权重应为 1");
let b = out[2].expect("c2 应入选");
assert!(approx(b.metric, t2, 1e-6));
assert!(approx(b.weight, t2 / t1, 1e-5));
let e = out[5].expect("c5 应入选");
assert!(approx(e.metric, t5, 1e-10));
assert!(approx(e.weight, t5 / t1, 1e-6));
let f = out[6].expect("c6 应入选");
assert!(approx(f.metric, t6, 1e-10));
assert!(approx(f.weight, t6 / t1, 1e-6));
}
#[test]
fn topk_limits_positive_count() {
let centers = grid8();
let gt = [[8.0f32, 0.0, 24.0, 16.0]];
let pred = [[8., 0., 24., 16.]; 8];
let scores: &[&[f32]] = &[&[0.0, 0.9, 0.8, 0.0, 0.0, 0.9, 0.8, 0.0]];
let cfg = TalConfig {
topk: 2,
..TalConfig::default()
};
let out = assign_single_image(&pred, ¢ers, scores, >, &cfg);
let picked: Vec<usize> = out
.iter()
.enumerate()
.filter(|(_, p)| p.is_some())
.map(|(c, _)| c)
.collect();
assert_eq!(picked, vec![1, 5], "topk=2 应按 t 降序取前两个");
}
#[test]
fn conflict_larger_metric_wins() {
let centers = grid8();
let gts = [[8.0f32, 0.0, 24.0, 16.0], [16.0, 8.0, 32.0, 24.0]];
let pred = [
[0., 0., 1., 1.],
[8., 0., 24., 16.], [8., 0., 24., 16.], [0., 0., 1., 1.],
[0., 0., 1., 1.],
[12., 4., 20., 12.], [16., 8., 24., 16.], [16., 8., 32., 24.], ];
let row_a = [0.0f32, 0.9, 0.8, 0.0, 0.0, 0.9, 0.9, 0.0];
let row_b = [0.0f32, 0.0, 0.0, 0.0, 0.0, 0.0, 0.8, 0.95];
let scores: &[&[f32]] = &[&row_a, &row_b];
let out = assign_single_image(&pred, ¢ers, scores, >s, &TalConfig::default());
assert_eq!(out[1].unwrap().gt, 0);
assert_eq!(out[2].unwrap().gt, 0);
assert_eq!(out[5].unwrap().gt, 0);
assert_eq!(out[6].unwrap().gt, 0, "t_A>t_B,争夺 cell 应归 A");
assert_eq!(out[7].unwrap().gt, 1);
for c in [0usize, 3, 4] {
assert!(out[c].is_none());
}
}
#[test]
fn conflict_tie_breaks_to_lower_gt_index() {
let centers = grid8();
let gts = [[8.0f32, 0.0, 24.0, 16.0], [16.0, 8.0, 32.0, 24.0]];
let mut pred = [[0.0f32, 0., 1., 1.]; 8];
pred[6] = [16., 8., 24., 16.];
let row_a = [0.0f32, 0.0, 0.0, 0.0, 0.0, 0.0, 0.8, 0.0];
let row_b = [0.0f32, 0.0, 0.0, 0.0, 0.0, 0.0, 0.8, 0.0];
let scores: &[&[f32]] = &[&row_a, &row_b];
let out = assign_single_image(&pred, ¢ers, scores, >s, &TalConfig::default());
assert_eq!(out[6].unwrap().gt, 0, "平局应归 gt 序号小者");
}
#[test]
fn tiny_gt_falls_back_to_nearest_center_cell() {
let centers = grid8();
let gt = [[25.0f32, 9.0, 27.0, 11.0]];
let pred = [[0.0f32, 0.0, 1.0, 1.0]; 8];
let scores: &[&[f32]] = &[&[0.5; 8]];
let out = assign_single_image(&pred, ¢ers, scores, >, &TalConfig::default());
let picked: Vec<usize> = out
.iter()
.enumerate()
.filter(|(_, p)| p.is_some())
.map(|(c, _)| c)
.collect();
assert_eq!(picked, vec![7], "应恰好回退到最近中心 cell c7");
assert_eq!(out[7].unwrap().gt, 0);
}
#[test]
fn empty_gts_yields_no_positives() {
let centers = grid8();
let pred = [[8.0f32, 0.0, 24.0, 16.0]; 8];
let scores: &[&[f32]] = &[&[0.9; 8]];
let out = assign_single_image(&pred, ¢ers, scores, &[], &TalConfig::default());
assert!(out.iter().all(|p| p.is_none()));
}
}