pub const COCO_SIGMAS: [f32; 17] = [
0.026, 0.025, 0.025, 0.035, 0.035, 0.079, 0.079, 0.072, 0.072, 0.062, 0.062, 0.107, 0.107,
0.087, 0.087, 0.089, 0.089,
];
pub fn sigma_for(j: usize) -> f32 {
COCO_SIGMAS[j.min(COCO_SIGMAS.len() - 1)]
}
pub fn sigma_table(k: usize) -> Vec<f32> {
(0..k).map(sigma_for).collect()
}
pub fn oks_scalar(pred: &[[f32; 3]], gt: &[[f32; 3]], scale: f32) -> f32 {
let s = scale.max(1e-3) as f64;
let mut num = 0f64;
let mut den = 0usize;
for (j, (p, g)) in pred.iter().zip(gt.iter()).enumerate() {
if g[2] <= 0.0 {
continue; }
let dx = (p[0] - g[0]) as f64;
let dy = (p[1] - g[1]) as f64;
let d2 = dx * dx + dy * dy;
let k = sigma_for(j) as f64;
num += (-d2 / (2.0 * s * s * k * k)).exp();
den += 1;
}
if den == 0 {
return 0.0;
}
(num / den as f64) as f32
}
#[cfg(feature = "torch")]
pub fn oks_loss(
pred: &tch::Tensor,
gt: &tch::Tensor,
vis: &tch::Tensor,
scale: &tch::Tensor,
) -> tch::Tensor {
use tch::Kind;
let size = pred.size();
let (g, k) = (size[0], size[1]);
let device = pred.device();
let sig: Vec<f32> = sigma_table(k as usize);
let sigma = tch::Tensor::from_slice(&sig)
.to_device(device)
.to_kind(Kind::Float)
.reshape([1i64, k, 1]);
let s2 = (scale * scale).reshape([g, 1i64, 1]); let denom = &(&s2 * 2.0) * &sigma * σ let d = pred - gt; let d2 = (d.select(2, 0) * d.select(2, 0) + d.select(2, 1) * d.select(2, 1)).reshape([g, k, 1]);
let e = (&d2 / &denom).neg().exp().reshape([g, k]); let num = (e * vis).sum_dim_intlist(&[1i64][..], false, Kind::Float); let den = vis
.sum_dim_intlist(&[1i64][..], false, Kind::Float)
.clamp_min(1.0);
let oks_mean = (&num / &den).mean(Kind::Float);
oks_mean * -1.0 + 1.0
}
#[cfg(all(test, feature = "torch"))]
mod tests {
use super::*;
use tch::{Device, Kind, Tensor};
#[test]
fn oks_same_point_is_one_and_known_offset_matches_hand_computed() {
let gt = [[10.0f32, 10.0, 2.0]];
let same = [[10.0f32, 10.0, 2.0]];
let oks = oks_scalar(&same, >, 32.0);
assert!((oks - 1.0).abs() < 1e-6, "同点 OKS 应为 1,got {oks}");
let off = [[12.0f32, 10.0, 2.0]];
let oks = oks_scalar(&off, >, 32.0);
let expected = (-(4.0f64) / (2.0 * 32.0 * 32.0 * 0.026 * 0.026)).exp() as f32;
assert!(
(oks - expected).abs() < 1e-3,
"已知偏移 OKS 应为 {expected},got {oks}"
);
assert!((expected - 0.0556).abs() < 1e-3, "手算校验: {expected}");
}
#[test]
fn oks_visibility_gating() {
let gt = [[10.0f32, 10.0, 2.0], [50.0, 50.0, 0.0]];
let pred = [[12.0f32, 10.0, 2.0], [91.0, 91.0, 2.0]];
let oks = oks_scalar(&pred, >, 32.0);
let expected = (-(4.0f64) / (2.0 * 32.0 * 32.0 * 0.026 * 0.026)).exp() as f32;
assert!(
(oks - expected).abs() < 1e-3,
"不可见点应被剔除:期望 {expected},got {oks}"
);
let gt0 = [[10.0f32, 10.0, 0.0]];
assert_eq!(oks_scalar(&pred, >0, 32.0), 0.0);
}
#[test]
fn oks_multi_point_mean() {
let gt = [[10.0f32, 10.0, 2.0], [50.0, 50.0, 2.0]];
let pred = [[12.0f32, 10.0, 2.0], [50.0, 50.0, 2.0]];
let oks = oks_scalar(&pred, >, 32.0);
let e0 = (-(4.0f64) / (2.0 * 32.0 * 32.0 * 0.026 * 0.026)).exp();
let expected = ((e0 + 1.0) / 2.0) as f32;
assert!((oks - expected).abs() < 1e-3, "期望 {expected},got {oks}");
}
#[test]
fn oks_loss_tensor_matches_scalar_and_backprops() {
let pred = Tensor::from_slice(&[
10.0f32, 10.0, 50.0, 50.0, 12.0, 10.0, 50.0, 50.0, ])
.to_device(Device::Cpu)
.to_kind(Kind::Float)
.reshape([2i64, 2, 2])
.set_requires_grad(true);
let gt = Tensor::from_slice(&[10.0f32, 10.0, 50.0, 50.0, 10.0, 10.0, 50.0, 50.0])
.to_device(Device::Cpu)
.to_kind(Kind::Float)
.reshape([2i64, 2, 2]);
let vis = Tensor::from_slice(&[1.0f32, 1.0, 1.0, 1.0])
.to_device(Device::Cpu)
.reshape([2i64, 2]);
let scale = Tensor::from_slice(&[32.0f32, 32.0])
.to_device(Device::Cpu)
.reshape([2i64]);
let loss = oks_loss(&pred, >, &vis, &scale);
let e0 = (-(4.0f64) / (2.0 * 32.0 * 32.0 * 0.026 * 0.026)).exp();
let expected = 1.0 - (1.0 + (e0 + 1.0) / 2.0) / 2.0;
let got = loss.double_value(&[]);
assert!(
(got - expected).abs() < 1e-3,
"张量 OKS 损失应 = {expected},got {got}"
);
loss.backward(); let g = pred.grad();
let gnorm = g.abs().sum(Kind::Float).double_value(&[]);
assert!(gnorm > 0.0, "pred 坐标应收到非零梯度(norm={gnorm})");
}
}