use tch::Tensor;
const COV: f64 = 1.0 / 12.0;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum KfTransform {
Log1p,
Sqrt1p,
Identity,
}
fn covariance(box5: &Tensor) -> (Tensor, Tensor, Tensor) {
debug_assert_eq!(box5.size().get(1).copied(), Some(5), "box5 通道维必须为 5");
let w = box5.select(1, 2);
let h = box5.select(1, 3);
let th = box5.select(1, 4);
let sin = th.sin();
let cos = th.cos();
let m2 = &w * &w * COV; let n2 = &h * &h * COV; let sxx = &m2 * &cos * &cos + &n2 * &sin * &sin;
let syy = &m2 * &sin * &sin + &n2 * &cos * &cos;
let sxy = (&m2 - &n2) * &sin * &cos;
(sxx, sxy, syy)
}
#[allow(clippy::too_many_arguments)]
fn kl_gaussians(
pxx: &Tensor,
pxy: &Tensor,
pyy: &Tensor,
txx: &Tensor,
txy: &Tensor,
tyy: &Tensor,
dx: &Tensor,
dy: &Tensor,
) -> Tensor {
let det_t = (txx * tyy - txy * txy).clamp_min(1e-12);
let det_p = (pxx * pyy - pxy * pxy).clamp_min(1e-12);
let tr = (&(tyy * pxx + txx * pyy) - txy * pxy * 2.0) / &det_t;
let quad = ((tyy * dx) * dx + (txx * dy) * dy - (txy * dx) * dy * 2.0) / &det_t;
let log_term = (&det_t / &det_p).log();
(tr + quad + log_term - 2.0) * 0.5
}
pub fn kfiou_element(pred5: &Tensor, target5: &Tensor, transform: KfTransform) -> Tensor {
let (pxx, pxy, pyy) = covariance(pred5);
let (txx, txy, tyy) = covariance(target5);
let dx = target5.select(1, 0) - pred5.select(1, 0);
let dy = target5.select(1, 1) - pred5.select(1, 1);
let kl = kl_gaussians(&pxx, &pxy, &pyy, &txx, &txy, &tyy, &dx, &dy);
let kl = kl.clamp_min(0.0);
match transform {
KfTransform::Log1p => (&kl + 1.0).log(),
KfTransform::Sqrt1p => (&kl + 1.0).sqrt() - 1.0,
KfTransform::Identity => kl,
}
}
pub fn probiou_element(pred5: &Tensor, target5: &Tensor) -> Tensor {
let (pxx, pxy, pyy) = covariance(pred5);
let (txx, txy, tyy) = covariance(target5);
let dx = target5.select(1, 0) - pred5.select(1, 0);
let dy = target5.select(1, 1) - pred5.select(1, 1);
let sxx = (&pxx + &txx) * 0.5;
let sxy = (&pxy + &txy) * 0.5;
let syy = (&pyy + &tyy) * 0.5;
let det_s = (&sxx * &syy - &sxy * &sxy).clamp_min(1e-12);
let det_p = (&pxx * &pyy - &pxy * &pxy).clamp_min(1e-12);
let det_t = (&txx * &tyy - &txy * &txy).clamp_min(1e-12);
let q = ((&syy * &dx) * &dx + (&sxx * &dy) * &dy - (&sxy * &dx) * &dy * 2.0) / &det_s;
let bc = (((&det_p * &det_t).sqrt() / &det_s).sqrt()) * (&q * -0.125).exp();
bc.clamp_min(1e-7).log().neg()
}
#[cfg(all(test, feature = "torch"))]
mod tests {
use super::*;
use tch::{Device, Kind};
const TOL: f64 = 1e-3;
fn b5(cx: f64, cy: f64, w: f64, h: f64, theta: f64) -> Tensor {
let v = [cx as f32, cy as f32, w as f32, h as f32, theta as f32];
Tensor::from_slice(&v).reshape([1i64, 5, 1, 1])
}
fn scalar(t: &Tensor) -> f64 {
t.double_value(&[])
}
#[test]
fn kfiou_identical_boxes_is_zero() {
let a = b5(3.0, 7.0, 10.0, 6.0, 0.3);
for tr in [
KfTransform::Log1p,
KfTransform::Sqrt1p,
KfTransform::Identity,
] {
let l = scalar(&kfiou_element(&a, &a, tr));
assert!(l.abs() < 1e-6, "{tr:?}: got {l}");
}
let p = scalar(&probiou_element(&a, &a));
assert!(p.abs() < 1e-6, "probiou: got {p}");
}
#[test]
fn kfiou_quarter_turn_with_wh_swap_is_near_zero() {
let a = b5(3.0, 7.0, 10.0, 6.0, 0.3);
let b = b5(3.0, 7.0, 6.0, 10.0, 0.3 + std::f64::consts::FRAC_PI_2);
let l = scalar(&kfiou_element(&a, &b, KfTransform::Identity));
assert!(l.abs() < TOL, "KFIoU 90°+互换应 ≈0,got {l}");
let p = scalar(&probiou_element(&a, &b));
assert!(p.abs() < TOL, "ProbIoU 90°+互换应 ≈0,got {p}");
}
#[test]
fn kfiou_known_angle_hand_computed() {
let a = b5(0.0, 0.0, 10.0, 6.0, 0.0);
let b = b5(0.0, 0.0, 10.0, 6.0, std::f64::consts::PI / 6.0);
let kl = scalar(&kfiou_element(&a, &b, KfTransform::Identity));
assert!((kl - 32.0 / 225.0).abs() < TOL, "KL={kl}");
let log1p = scalar(&kfiou_element(&a, &b, KfTransform::Log1p));
assert!(
(log1p - (257.0f64 / 225.0).ln()).abs() < TOL,
"log1p={log1p}"
);
let sqrt1p = scalar(&kfiou_element(&a, &b, KfTransform::Sqrt1p));
assert!(
(sqrt1p - ((1.0f64 + 32.0f64 / 225.0).sqrt() - 1.0)).abs() < TOL,
"sqrt1p={sqrt1p}"
);
}
#[test]
fn kfiou_center_offset_hand_computed() {
let a = b5(0.0, 0.0, 10.0, 6.0, 0.0);
let b = b5(2.0, 0.0, 10.0, 6.0, 0.0);
let kl = scalar(&kfiou_element(&a, &b, KfTransform::Identity));
assert!((kl - 0.24).abs() < TOL, "KL={kl}");
let log1p = scalar(&kfiou_element(&a, &b, KfTransform::Log1p));
assert!((log1p - 1.24f64.ln()).abs() < TOL, "log1p={log1p}");
}
#[test]
fn probiou_known_angle_hand_computed() {
let a = b5(0.0, 0.0, 10.0, 6.0, 0.0);
let b = b5(0.0, 0.0, 10.0, 6.0, std::f64::consts::PI / 6.0);
let l = scalar(&probiou_element(&a, &b));
assert!((l - 0.0343191).abs() < TOL, "probiou loss={l}");
}
#[test]
fn kfiou_backward_smoke() {
let pred = Tensor::from_slice(&[1.0f32, 2.0, 10.0, 6.0, 0.2])
.reshape([1i64, 5, 1, 1])
.set_requires_grad(true);
let target = b5(2.0, 3.0, 10.0, 6.0, 0.5);
let loss = kfiou_element(&pred, &target, KfTransform::Log1p).sum(Kind::Float);
loss.backward();
let g = pred.grad();
assert!(g.numel() == 5);
let gv = g.to_device(Device::Cpu).to_kind(Kind::Float).reshape([-1]);
for i in 0..5i64 {
let v = gv.double_value(&[i]);
assert!(v.is_finite(), "grad[{i}] = {v}");
}
}
}