use av_core::geometry::{Aabb, RotBox};
use av_core::types::Detection;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RotNmsMetric {
Polygon,
ProbIou,
}
pub fn envelope_half_extents(w: f32, h: f32, theta: f32) -> (f32, f32) {
let (sin, cos) = theta.sin_cos();
let hw = (w * cos.abs() + h * sin.abs()) / 2.0;
let hh = (w * sin.abs() + h * cos.abs()) / 2.0;
(hw, hh)
}
fn det_rotbox(d: &Detection) -> RotBox {
let (cx, cy) = d.bbox.center();
RotBox {
cx,
cy,
w: d.bbox.w(),
h: d.bbox.h(),
theta: d.angle.unwrap_or(0.0),
}
}
fn envelope_aabb(r: &RotBox) -> Aabb {
let (hw, hh) = envelope_half_extents(r.w, r.h, r.theta);
Aabb::new(r.cx - hw, r.cy - hh, r.cx + hw, r.cy + hh)
}
pub fn rotate_nms(mut dets: Vec<Detection>, iou_thr: f32, metric: RotNmsMetric) -> Vec<Detection> {
dets.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
let mut kept: Vec<(Detection, RotBox, Aabb)> = Vec::new();
for d in dets {
let rb = det_rotbox(&d);
let env = envelope_aabb(&rb);
let suppressed = kept.iter().any(|(_, rk, renv)| {
if matches!(metric, RotNmsMetric::Polygon) && renv.intersection(&env).area() <= 0.0 {
return false; }
let ov = match metric {
RotNmsMetric::Polygon => rk.iou(&rb),
RotNmsMetric::ProbIou => probiou_scalar(*rk, rb),
};
ov > iou_thr
});
if !suppressed {
kept.push((d, rb, env));
}
}
kept.into_iter().map(|(d, _, _)| d).collect()
}
pub fn probiou_scalar(a: RotBox, b: RotBox) -> f32 {
fn cov(r: &RotBox) -> [f32; 3] {
let (sin, cos) = r.theta.sin_cos();
let m = r.w * r.w / 12.0;
let n = r.h * r.h / 12.0;
[
m * cos * cos + n * sin * sin,
(m - n) * sin * cos,
m * sin * sin + n * cos * cos,
]
}
let [axx, axy, ayy] = cov(&a);
let [bxx, bxy, byy] = cov(&b);
let (dx, dy) = (b.cx - a.cx, b.cy - a.cy);
let sxx = (axx + bxx) * 0.5;
let sxy = (axy + bxy) * 0.5;
let syy = (ayy + byy) * 0.5;
let det_s = (sxx * syy - sxy * sxy).max(1e-12);
let det_a = (axx * ayy - axy * axy).max(1e-12);
let det_b = (bxx * byy - bxy * bxy).max(1e-12);
let q = (syy * dx * dx + sxx * dy * dy - sxy * dx * dy * 2.0) / det_s;
((det_a * det_b).sqrt() / det_s).sqrt() * (-0.125 * q).exp()
}
#[cfg(test)]
mod tests {
use super::*;
use av_core::geometry::Aabb;
use std::f32::consts::PI;
fn det(cx: f32, cy: f32, w: f32, h: f32, theta: f32, score: f32) -> Detection {
Detection {
bbox: Aabb::new(cx - w / 2.0, cy - h / 2.0, cx + w / 2.0, cy + h / 2.0),
score,
class_id: 0,
angle: Some(theta),
keypoints: None,
}
}
#[test]
fn angle_normalize_roundtrip_and_period_pi() {
for d in [
av_core::conventions::AngleDomain::Le90,
av_core::conventions::AngleDomain::Le135,
av_core::conventions::AngleDomain::OpenCv,
] {
for &t in &[-4.0f32, -2.8, -0.9, 0.0, 0.3, 1.2, 2.5] {
let n1 = d.normalize(t);
let n2 = d.normalize(n1); assert!((n1 - n2).abs() < 1e-5, "{d:?}: {t} -> {n1} -> {n2}");
let n3 = d.normalize(t + PI); assert!((n1 - n3).abs() < 1e-4, "{d:?}: period π {t}: {n1} vs {n3}");
let (lo, hi) = d.range();
assert!(n1 >= lo - 1e-4 && n1 < hi, "{d:?}: {t} -> {n1} 越界");
}
}
}
#[test]
fn rotated_boxes_avoid_angle_blind_suppression() {
let s = std::f32::consts::FRAC_1_SQRT_2;
let a = det(50.0, 50.0, 40.0, 8.0, std::f32::consts::FRAC_PI_4, 0.9);
let b = det(
50.0 - 8.0 * s,
50.0 + 8.0 * s,
40.0,
8.0,
std::f32::consts::FRAC_PI_4,
0.7,
);
let (ra, rb) = (det_rotbox(&a), det_rotbox(&b));
let rot_iou = ra.iou(&rb);
assert!(
rot_iou < 0.01,
"前置:平行错位条带旋转 IoU 应≈0,got {rot_iou}"
);
let (hwa, hha) = envelope_half_extents(40.0, 8.0, std::f32::consts::FRAC_PI_4);
let env_a = Aabb::new(50.0 - hwa, 50.0 - hha, 50.0 + hwa, 50.0 + hha);
let (cxb, cyb) = b.bbox.center();
let env_b = Aabb::new(cxb - hwa, cyb - hha, cxb + hwa, cyb + hha);
assert!(env_a.iou(&env_b) > 0.5, "前置:外接框角度盲抑制会误杀");
let kept = rotate_nms(vec![a, b], 0.5, RotNmsMetric::Polygon);
assert_eq!(kept.len(), 2, "平行错位旋转框不得被抑制");
assert!((kept[0].score - 0.9).abs() < 1e-6);
assert!((kept[1].score - 0.7).abs() < 1e-6);
}
#[test]
fn cross_placed_rotated_boxes_not_suppressed() {
let a = det(50.0, 50.0, 40.0, 8.0, 0.0, 0.9);
let b = det(50.0, 50.0, 40.0, 8.0, std::f32::consts::FRAC_PI_2, 0.7);
let kept = rotate_nms(vec![a, b], 0.5, RotNmsMetric::Polygon);
assert_eq!(kept.len(), 2, "交叉放置的旋转框不得被抑制");
assert!((kept[0].score - 0.9).abs() < 1e-6);
assert!((kept[1].score - 0.7).abs() < 1e-6);
}
#[test]
fn rotated_duplicate_is_suppressed() {
let a = det(50.0, 50.0, 40.0, 8.0, 0.0, 0.9);
let dup = det(51.0, 50.0, 40.0, 8.0, 0.0, 0.8);
let far = det(200.0, 200.0, 40.0, 8.0, 0.0, 0.5);
let kept = rotate_nms(vec![a, dup, far], 0.5, RotNmsMetric::Polygon);
assert_eq!(kept.len(), 2);
assert!((kept[0].score - 0.9).abs() < 1e-6);
assert!((kept[1].score - 0.5).abs() < 1e-6);
}
#[test]
fn angle_none_degenerates_to_plain_nms() {
let base = |x1: f32, score: f32| Detection {
bbox: Aabb::new(x1, 0.0, x1 + 10.0, 10.0),
score,
class_id: 0,
angle: None,
keypoints: None,
};
let dets = vec![base(0.0, 0.7), base(1.0, 0.9), base(50.0, 0.5)];
let got = rotate_nms(dets.clone(), 0.5, RotNmsMetric::Polygon);
let want = av_core::types::nms(dets, 0.5);
assert_eq!(got.len(), want.len());
for (g, w) in got.iter().zip(&want) {
assert_eq!(g.score.to_bits(), w.score.to_bits());
}
}
#[test]
fn probiou_metric_hand_check_and_nms() {
let a = RotBox {
cx: 0.0,
cy: 0.0,
w: 40.0,
h: 8.0,
theta: 0.0,
};
let b = RotBox {
cx: 0.0,
cy: 0.0,
w: 8.0,
h: 40.0,
theta: 0.0,
};
let bc_cross = probiou_scalar(a, b);
assert!((bc_cross - 0.384616).abs() < 1e-3, "BC={bc_cross}");
let dup = RotBox {
cx: 1.0,
cy: 0.0,
w: 40.0,
h: 8.0,
theta: 0.0,
};
let bc_dup = probiou_scalar(a, dup);
assert!(bc_dup > 0.99, "BC dup={bc_dup}");
assert!(probiou_scalar(a, a) > 0.9999);
let s = std::f32::consts::FRAC_1_SQRT_2;
let ka = det(50.0, 50.0, 40.0, 8.0, std::f32::consts::FRAC_PI_4, 0.9);
let kdup = det(
50.0 + s,
50.0 + s,
40.0,
8.0,
std::f32::consts::FRAC_PI_4,
0.8,
);
let kb = det(
50.0 - 8.0 * s,
50.0 + 8.0 * s,
40.0,
8.0,
std::f32::consts::FRAC_PI_4,
0.7,
);
let kept = rotate_nms(vec![ka, kdup, kb], 0.5, RotNmsMetric::ProbIou);
assert_eq!(kept.len(), 2, "ProbIou 度量:重复框抑制、错位条带保留");
}
#[test]
fn envelope_matches_rotbox_corners() {
let (w, h, th) = (10.0f32, 6.0f32, std::f32::consts::FRAC_PI_6);
let rb = RotBox {
cx: 0.0,
cy: 0.0,
w,
h,
theta: th,
};
let mut xmax = 0f32;
let mut ymax = 0f32;
for c in rb.corners() {
xmax = xmax.max(c[0].abs());
ymax = ymax.max(c[1].abs());
}
let (hw, hh) = envelope_half_extents(w, h, th);
assert!((hw - xmax).abs() < 1e-5);
assert!((hh - ymax).abs() < 1e-5);
}
}