use std::cell::RefCell;
use rayon::prelude::*;
use tch::nn;
use tch::{Device, Kind, Tensor};
use av_core::config::{RunConfig, TaskCfg};
use av_core::conventions::AngleDomain;
use av_core::error::{AvError, AvResult};
use av_core::geometry::Aabb;
use av_core::traits::BaseBackbone;
use av_core::types::{nms, Detection};
use crate::assigner::{self, TalConfig};
use crate::backbone::SimpleCnnBackbone;
use crate::backbone_cspelan::{CspElanBackbone, FAMILY_NAME as CSP_FAMILY};
use crate::backbone_dino::{DinoV2Backbone, FAMILY_NAME as DINO_FAMILY};
use crate::backbone_resnet::{ResNetBackbone, FAMILY_NAME as RESNET_FAMILY};
use crate::heads::{ClassifyHead, DetectHead, REG_MAX};
use crate::keypoint::KeypointHead;
use crate::kfiou::{kfiou_element, KfTransform};
#[cfg(test)]
use crate::mask::mask_iou;
use crate::mask::{dice_loss, MaskBranch, MaskSummary};
use crate::oks::sigma_table;
use crate::rot_nms::{envelope_half_extents, rotate_nms, RotNmsMetric};
pub enum TrainBatch {
Classify {
labels: Tensor,
},
Detect {
boxes: Vec<Vec<[f32; 4]>>,
labels: Vec<Vec<u32>>,
},
Obb {
boxes: Vec<Vec<[f32; 5]>>,
labels: Vec<Vec<u32>>,
},
Seg {
masks: Vec<Vec<Vec<u8>>>,
labels: Vec<Vec<u32>>,
},
Keypoint {
boxes: Vec<Vec<[f32; 4]>>,
kpts: Vec<Vec<Vec<[f32; 3]>>>,
labels: Vec<Vec<u32>>,
},
}
#[derive(Debug, Clone)]
pub struct SegInstance {
pub label: u32,
pub score: f32,
pub mask: Vec<u8>,
}
#[allow(clippy::large_enum_variant)]
pub enum PredictOutput {
Classify {
labels: Vec<u32>,
confs: Vec<f32>,
},
Detect {
per_image: Vec<Vec<Detection>>,
},
Seg {
per_image: Vec<Vec<SegInstance>>,
},
Keypoint {
per_image: Vec<Vec<Detection>>,
},
}
#[allow(clippy::large_enum_variant)]
pub enum TaskModel {
Classify(ClassifyModel),
Detect(DetectModel),
Seg(SegModel),
Keypoint(KeypointModel),
}
pub struct ClassifyModel {
backbone: ClassifyBackbone,
head: ClassifyHead,
img_size: u32,
}
enum ClassifyBackbone {
SimpleCnn(SimpleCnnBackbone),
ResNet18(ResNetBackbone),
DinoV2(DinoV2Backbone),
CspElan(CspElanBackbone),
}
impl ClassifyBackbone {
fn pooled_channels(&self) -> i64 {
match self {
Self::SimpleCnn(b) => b.pooled_channels(),
Self::ResNet18(b) => b.pooled_channels(),
Self::DinoV2(b) => b.pooled_channels(),
Self::CspElan(b) => b.pooled_channels(),
}
}
fn forward_pooled(&self, x: &Tensor) -> AvResult<Tensor> {
match self {
Self::SimpleCnn(b) => b.forward_pooled(x),
Self::ResNet18(b) => b.forward_pooled(x),
Self::DinoV2(b) => b.forward_pooled(x),
Self::CspElan(b) => b.forward_pooled(x),
}
}
fn set_train(&self, train: bool) {
match self {
Self::SimpleCnn(_) | Self::DinoV2(_) => {}
Self::ResNet18(b) => b.set_train(train),
Self::CspElan(b) => b.set_train(train),
}
}
}
fn build_classify_backbone(
p: &nn::Path,
cfg: &av_core::config::BackboneCfg,
img_size: u32,
) -> AvResult<ClassifyBackbone> {
match cfg.family.as_str() {
DINO_FAMILY => {
let mut b = DinoV2Backbone::new(p, cfg, img_size)?;
if !matches!(cfg.pretrained.as_str(), "" | "none") {
let path = std::path::Path::new(&cfg.pretrained);
if !path.exists() {
return Err(AvError::config(format!(
"backbone.pretrained = {:?} 指向的权重文件不存在(下载见 \
tools/export/export_dinov2.py)",
cfg.pretrained
)));
}
let stats = b.load_dinov2_weights(path)?;
println!("[dinov2] 预训练导入 {}", stats.summary());
}
Ok(ClassifyBackbone::DinoV2(b))
}
RESNET_FAMILY => Ok(ClassifyBackbone::ResNet18(ResNetBackbone::new(p, cfg)?)),
CSP_FAMILY => Ok(ClassifyBackbone::CspElan(CspElanBackbone::new(p, cfg)?)),
_ => Ok(ClassifyBackbone::SimpleCnn(SimpleCnnBackbone::new(p, cfg))),
}
}
#[allow(clippy::large_enum_variant)]
enum DetectBackbone {
SimpleCnn(SimpleCnnBackbone),
ResNet18(ResNetBackbone),
CspElan(CspElanBackbone),
DinoV2(DinoV2Backbone),
}
impl DetectBackbone {
fn forward_features(&self, x: &Tensor) -> AvResult<av_core::traits::FeaturePyramid> {
match self {
Self::SimpleCnn(b) => b.forward_features(x),
Self::ResNet18(b) => b.forward_features(x),
Self::DinoV2(b) => b.forward_features(x),
Self::CspElan(b) => b.forward_features(x),
}
}
fn stride_channels(&self, stride: u32) -> AvResult<i64> {
match self {
Self::SimpleCnn(b) => b.stride_channels(stride),
Self::ResNet18(b) => b.stride_channels(stride),
Self::DinoV2(b) => b.stride_channels(stride),
Self::CspElan(b) => b.stride_channels(stride),
}
}
fn set_train(&self, train: bool) {
match self {
Self::SimpleCnn(_) => {}
Self::ResNet18(b) => b.set_train(train),
Self::DinoV2(_) => {}
Self::CspElan(b) => b.set_train(train),
}
}
}
fn build_detect_backbone(
p: &nn::Path,
cfg: &av_core::config::BackboneCfg,
img_size: u32,
) -> AvResult<DetectBackbone> {
match cfg.family.as_str() {
RESNET_FAMILY => Ok(DetectBackbone::ResNet18(ResNetBackbone::new(p, cfg)?)),
CSP_FAMILY => Ok(DetectBackbone::CspElan(CspElanBackbone::new(p, cfg)?)),
DINO_FAMILY => Ok(DetectBackbone::DinoV2(DinoV2Backbone::new(
p, cfg, img_size,
)?)),
_ => Ok(DetectBackbone::SimpleCnn(SimpleCnnBackbone::new(p, cfg))),
}
}
enum SegBackbone {
SimpleCnn(SimpleCnnBackbone),
ResNet18(ResNetBackbone),
CspElan(CspElanBackbone),
DinoV2(DinoV2Backbone),
}
impl SegBackbone {
fn forward_features(&self, x: &Tensor) -> AvResult<av_core::traits::FeaturePyramid> {
match self {
Self::SimpleCnn(b) => b.forward_features(x),
Self::ResNet18(b) => b.forward_features(x),
Self::DinoV2(b) => b.forward_features(x),
Self::CspElan(b) => b.forward_features(x),
}
}
fn stride_channels(&self, stride: u32) -> AvResult<i64> {
match self {
Self::SimpleCnn(b) => b.stride_channels(stride),
Self::ResNet18(b) => b.stride_channels(stride),
Self::DinoV2(b) => b.stride_channels(stride),
Self::CspElan(b) => b.stride_channels(stride),
}
}
fn set_train(&self, train: bool) {
match self {
Self::SimpleCnn(_) | Self::DinoV2(_) => {}
Self::ResNet18(b) => b.set_train(train),
Self::CspElan(b) => b.set_train(train),
}
}
}
fn build_seg_backbone(
p: &nn::Path,
cfg: &av_core::config::BackboneCfg,
img_size: u32,
) -> AvResult<SegBackbone> {
match cfg.family.as_str() {
DINO_FAMILY => {
let mut b = DinoV2Backbone::new(p, cfg, img_size)?;
if !matches!(cfg.pretrained.as_str(), "" | "none") {
let path = std::path::Path::new(&cfg.pretrained);
if !path.exists() {
return Err(AvError::config(format!(
"backbone.pretrained = {:?} 指向的权重文件不存在(下载见 tools/export/export_dinov2.py)",
cfg.pretrained
)));
}
let stats = b.load_dinov2_weights(path)?;
println!("[dinov2] 预训练导入 {}", stats.summary());
}
Ok(SegBackbone::DinoV2(b))
}
RESNET_FAMILY => Ok(SegBackbone::ResNet18(ResNetBackbone::new(p, cfg)?)),
CSP_FAMILY => Ok(SegBackbone::CspElan(CspElanBackbone::new(p, cfg)?)),
_ => Ok(SegBackbone::SimpleCnn(SimpleCnnBackbone::new(p, cfg))),
}
}
#[derive(Debug, Clone, Copy)]
pub struct ObbParams {
pub angle_domain: AngleDomain,
}
pub struct DetectModel {
backbone: DetectBackbone,
head: DetectHead,
img_size: u32,
head_levels: Vec<u32>,
assigner: String,
loss_w_cls: f64,
loss_w_ciou: f64,
loss_w_dfl: f64,
obb: Option<ObbParams>,
}
pub struct SegModel {
backbone: SegBackbone,
mask_branch: MaskBranch,
num_classes: i64,
img_size: u32,
loss_weight: f64,
loss_w_bce: f64,
loss_w_dice: f64,
}
pub struct KeypointModel {
backbone: KeypointBackbone,
head: KeypointHead,
img_size: u32,
loss_weight: f64,
loss_w_oks: f64,
}
enum KeypointBackbone {
SimpleCnn(SimpleCnnBackbone),
ResNet18(ResNetBackbone),
CspElan(CspElanBackbone),
DinoV2(DinoV2Backbone),
}
impl KeypointBackbone {
fn forward_features(&self, x: &Tensor) -> AvResult<av_core::traits::FeaturePyramid> {
match self {
Self::SimpleCnn(b) => b.forward_features(x),
Self::ResNet18(b) => b.forward_features(x),
Self::DinoV2(b) => b.forward_features(x),
Self::CspElan(b) => b.forward_features(x),
}
}
fn stride_channels(&self, stride: u32) -> AvResult<i64> {
match self {
Self::SimpleCnn(b) => b.stride_channels(stride),
Self::ResNet18(b) => b.stride_channels(stride),
Self::DinoV2(b) => b.stride_channels(stride),
Self::CspElan(b) => b.stride_channels(stride),
}
}
fn set_train(&self, train: bool) {
match self {
Self::SimpleCnn(_) => {}
Self::ResNet18(b) => b.set_train(train),
Self::DinoV2(_) => {}
Self::CspElan(b) => b.set_train(train),
}
}
}
fn build_keypoint_backbone(
p: &nn::Path,
cfg: &av_core::config::BackboneCfg,
img_size: u32,
) -> AvResult<KeypointBackbone> {
match cfg.family.as_str() {
RESNET_FAMILY => Ok(KeypointBackbone::ResNet18(ResNetBackbone::new(p, cfg)?)),
CSP_FAMILY => Ok(KeypointBackbone::CspElan(CspElanBackbone::new(p, cfg)?)),
DINO_FAMILY => Ok(KeypointBackbone::DinoV2(DinoV2Backbone::new(
p, cfg, img_size,
)?)),
_ => Ok(KeypointBackbone::SimpleCnn(SimpleCnnBackbone::new(p, cfg))),
}
}
pub const MAX_KPS_PER_IMAGE: usize = 100;
pub const LOSS_W_KP_CLS: f64 = 1.0;
pub const LOSS_W_KP_BOX: f64 = 5.0;
pub const LOSS_W_KP_OFF: f64 = 2.0;
pub const LOSS_W_KP_VIS: f64 = 0.5;
pub const MAX_SEGS_PER_IMAGE: usize = 100;
pub fn mask_nms(mut insts: Vec<SegInstance>, nms_iou: f32) -> Vec<SegInstance> {
insts.sort_by(|a, b| b.score.total_cmp(&a.score));
insts.truncate(MAX_SEGS_PER_IMAGE);
let mut kept: Vec<SegInstance> = Vec::new();
let mut kept_sum: Vec<MaskSummary> = Vec::new();
for cand in insts {
let cand_sum = MaskSummary::of(&cand.mask);
let suppressed = kept.iter().zip(&kept_sum).any(|(kp, ks)| {
kp.label == cand.label && ks.iou(&kp.mask, &cand_sum, &cand.mask) >= nms_iou
});
if !suppressed {
kept_sum.push(cand_sum);
kept.push(cand);
}
}
kept
}
pub fn build_model(p: &nn::Path, cfg: &RunConfig) -> AvResult<TaskModel> {
let task = cfg
.model
.tasks
.first()
.ok_or_else(|| AvError::config("model.tasks 不能为空"))?;
let backbone_cfg = cfg.model.backbone.clone();
match task {
TaskCfg::Classify(c) => {
let backbone = build_classify_backbone(&(p / "backbone"), &backbone_cfg, c.img_size)?;
let head = ClassifyHead::new(
&(p / "head"),
backbone.pooled_channels(),
c.num_classes as i64,
);
Ok(TaskModel::Classify(ClassifyModel {
backbone,
head,
img_size: c.img_size,
}))
}
TaskCfg::Detect(d) => {
if !matches!(d.head.as_str(), "yolo") {
return Err(AvError::config("rfdetr 头在 M6 落地(PLAN §8)"));
}
let head_levels = validate_head_levels(&d.head_levels)?;
let backbone = build_detect_backbone(&(p / "backbone"), &backbone_cfg, d.img_size)?;
let channels: AvResult<Vec<i64>> = head_levels
.iter()
.map(|&s| backbone.stride_channels(s))
.collect();
let channels = channels?;
let head = DetectHead::with_mode(
&(p / "head"),
&head_levels,
&channels,
d.num_classes as i64,
d.obb_mode,
);
Ok(TaskModel::Detect(DetectModel {
backbone,
head,
img_size: d.img_size,
head_levels,
assigner: d.assigner.clone(),
loss_w_cls: d.loss_cls_weight as f64,
loss_w_ciou: d.loss_ciou_weight as f64,
loss_w_dfl: d.loss_dfl_weight as f64,
obb: d.obb_mode.then_some(ObbParams {
angle_domain: AngleDomain::Le90,
}),
}))
}
TaskCfg::Seg(s) => {
if !matches!(s.head.as_str(), "yolact") {
return Err(AvError::config(
"seg.head = \"direct\"(精度档 RoIAlign 逐实例掩码头)按 M4 落地,\
当前请使用 head = \"yolact\"(实时档)",
));
}
let backbone = build_seg_backbone(&(p / "backbone"), &backbone_cfg, s.img_size)?;
let mask_branch = MaskBranch::new(
&(p / "mask"),
backbone.stride_channels(16)?,
s.num_classes as i64,
s.num_protos as i64,
);
Ok(TaskModel::Seg(SegModel {
backbone,
mask_branch,
num_classes: s.num_classes as i64,
img_size: s.img_size,
loss_weight: s.loss_weight as f64,
loss_w_bce: s.loss_bce_weight as f64,
loss_w_dice: s.loss_dice_weight as f64,
}))
}
TaskCfg::Keypoint(k) => {
if !matches!(k.decode.as_str(), "direct") {
return Err(AvError::config(format!(
"keypoint.decode = \"{}\" 未落地(PLAN §4.4 精度/实时档按里程碑补齐),\
当前请使用 decode = \"direct\"(直接回归档:stride 8 每 cell 回归 \
cls+box+K×(dx,dy)+可见性,OKS 可微损失可直连)",
k.decode
)));
}
let backbone = build_keypoint_backbone(&(p / "backbone"), &backbone_cfg, k.img_size)?;
let head = KeypointHead::new(
&(p / "head"),
backbone.stride_channels(8)?,
k.num_keypoints as i64,
);
Ok(TaskModel::Keypoint(KeypointModel {
backbone,
head,
img_size: k.img_size,
loss_weight: k.loss_weight as f64,
loss_w_oks: k.loss_oks_weight as f64,
}))
}
TaskCfg::Obb(_) => Err(AvError::config(
"独立 kind = \"obb\" 任务暂不可装配:ObbCfg 缺少 num_classes/img_size 字段\
(av-core 本期只读)。请使用 detect 任务 + obb_mode = true(PLAN §4.2:\
OBB 与普通检测共用检测头,仅角度分支与损失分叉),模型侧能力已完备;\
如需独立 obb 任务,请在 av-core 的 ObbCfg 增补 num_classes/img_size \
(及损失权重)字段后再接入",
)),
}
}
fn atanh_clamp(x: f32) -> f32 {
let x = x.clamp(-0.95, 0.95);
0.5 * ((1.0 + x) / (1.0 - x)).ln()
}
thread_local! {
pub static LOSS_DEBUG: RefCell<Option<String>> = const { RefCell::new(None) };
}
fn tensor_to_vec_f32(t: &Tensor) -> Vec<f32> {
let t = t
.to_device(Device::Cpu)
.to_kind(Kind::Float)
.contiguous()
.reshape([-1]);
let n = t.size()[0] as usize;
let mut dst = vec![0f32; n];
t.copy_data(&mut dst, n);
dst
}
fn tensor_to_vec_i64(t: &Tensor) -> Vec<i64> {
let t = t
.to_device(Device::Cpu)
.to_kind(Kind::Int64)
.contiguous()
.reshape([-1]);
let n = t.size()[0] as usize;
let mut dst = vec![0i64; n];
t.copy_data(&mut dst, n);
dst
}
fn tensor_to_vec_u8(t: &Tensor) -> Vec<u8> {
let t = t
.to_device(Device::Cpu)
.to_kind(Kind::Uint8)
.contiguous()
.reshape([-1]);
let n = t.size()[0] as usize;
let mut dst = vec![0u8; n];
t.copy_data(&mut dst, n);
dst
}
fn pyramid_level(
py: &av_core::traits::FeaturePyramid,
stride: i64,
) -> AvResult<&av_core::traits::FeatureMap> {
py.levels
.iter()
.find(|l| l.stride as i64 == stride)
.ok_or_else(|| AvError::shape(format!("骨干缺少 stride {stride} 特征层")))
}
fn validate_head_levels(head_levels: &[u32]) -> AvResult<Vec<u32>> {
if head_levels.is_empty() {
return Err(AvError::config("detect.head_levels 不能为空"));
}
if !head_levels.windows(2).all(|w| w[0] < w[1]) {
return Err(AvError::config(format!(
"detect.head_levels 必须严格升序且无重复,got {head_levels:?}"
)));
}
Ok(head_levels.to_vec())
}
fn level_breaks(img_size: u32, head_levels: &[u32]) -> Vec<f32> {
head_levels
.iter()
.map(|&l| img_size as f32 * l as f32 / 32.0)
.collect()
}
fn select_level(head_levels: &[u32], breaks: &[f32], side: f32) -> u32 {
debug_assert_eq!(head_levels.len(), breaks.len());
head_levels
.iter()
.zip(breaks.iter())
.find(|(_s, &b)| side <= b)
.map(|(&s, _)| s)
.unwrap_or(*head_levels.last().expect("head_levels 非空"))
}
pub const DFL_SHIFT: f32 = (REG_MAX as f32 - 1.0) / 2.0;
pub const LOSS_W_CLS: f64 = 1.0;
pub const LOSS_W_CIOU: f64 = 5.0;
pub const LOSS_W_DFL: f64 = 1.5;
pub const LOSS_W_KFIOU_SCALE: f64 = 10.0;
fn dfl_project(dist: &Tensor) -> Tensor {
let size = dist.size();
let (n, h, w) = (size[0], size[2], size[3]);
let device = dist.device();
let prob = dist.reshape([n, 4, REG_MAX, h, w]).softmax(2, Kind::Float);
let bins = Tensor::arange(REG_MAX, (Kind::Float, device)).reshape([1i64, 1, REG_MAX, 1, 1]);
let expect = (&prob * &bins).sum_dim_intlist(&[2i64][..], false, Kind::Float);
expect - (DFL_SHIFT as f64)
}
fn decode_pred_cxcywh(box_raw: &Tensor, s: f32) -> Tensor {
let size = box_raw.size();
let (_n, h, w) = (size[0], size[2], size[3]);
let device = box_raw.device();
let xs = Tensor::arange(w, (Kind::Float, device)).reshape([1i64, 1, 1, w]);
let ys = Tensor::arange(h, (Kind::Float, device)).reshape([1i64, 1, h, 1]);
let tx = box_raw.select(1, 0).unsqueeze(1);
let ty = box_raw.select(1, 1).unsqueeze(1);
let tw = box_raw.select(1, 2).unsqueeze(1);
let th = box_raw.select(1, 3).unsqueeze(1);
let cx = (&xs + 0.5 + tx.tanh()) * (s as f64);
let cy = (&ys + 0.5 + ty.tanh()) * (s as f64);
let bw = tw.exp() * (s as f64);
let bh = th.exp() * (s as f64);
Tensor::cat(&[&cx, &cy, &bw, &bh], 1)
}
fn decode_pred_xyxy(box_raw: &Tensor, s: f32) -> Tensor {
let b = decode_pred_cxcywh(box_raw, s);
let cx = b.select(1, 0).unsqueeze(1);
let cy = b.select(1, 1).unsqueeze(1);
let bw = b.select(1, 2).unsqueeze(1);
let bh = b.select(1, 3).unsqueeze(1);
Tensor::cat(
&[
&(&cx - &bw / 2.0),
&(&cy - &bh / 2.0),
&(&cx + &bw / 2.0),
&(&cy + &bh / 2.0),
],
1,
)
}
fn normalize_theta_tensor(t: &Tensor, domain: AngleDomain) -> Tensor {
let pi = std::f64::consts::PI;
let (lo, _hi) = domain.range();
let folds = ((t - (lo as f64)) / pi).floor();
t - folds * pi
}
fn decode_theta(t_theta: &Tensor, domain: AngleDomain) -> Tensor {
let half_pi = std::f64::consts::FRAC_PI_2;
let th = t_theta.tanh() * half_pi;
normalize_theta_tensor(&th, domain)
}
fn ciou_element(pred: &Tensor, gt: &Tensor) -> Tensor {
let px1 = pred.select(1, 0);
let py1 = pred.select(1, 1);
let px2 = pred.select(1, 2);
let py2 = pred.select(1, 3);
let gx1 = gt.select(1, 0);
let gy1 = gt.select(1, 1);
let gx2 = gt.select(1, 2);
let gy2 = gt.select(1, 3);
let ix1 = px1.maximum(&gx1);
let iy1 = py1.maximum(&gy1);
let ix2 = px2.minimum(&gx2);
let iy2 = py2.minimum(&gy2);
let inter = (ix2 - ix1).clamp_min(0.0) * (iy2 - iy1).clamp_min(0.0);
let parea = (&px2 - &px1) * (&py2 - &py1);
let garea = (&gx2 - &gx1) * (&gy2 - &gy1);
let union = (&parea + &garea - &inter).clamp_min(1e-7);
let iou = &inter / &union;
let pcx = (&px1 + &px2) / 2.0;
let pcy = (&py1 + &py2) / 2.0;
let gcx = (&gx1 + &gx2) / 2.0;
let gcy = (&gy1 + &gy2) / 2.0;
let dx = &pcx - &gcx;
let dy = &pcy - г
let rho2 = &dx * &dx + &dy * &dy;
let ex1 = px1.minimum(&gx1);
let ey1 = py1.minimum(&gy1);
let ex2 = px2.maximum(&gx2);
let ey2 = py2.maximum(&gy2);
let cdx = ex2 - ex1;
let cdy = ey2 - ey1;
let c2 = &(&cdx * &cdx + &cdy * &cdy) + 1e-7;
let pw = (&px2 - &px1).clamp_min(1e-4);
let ph = (&py2 - &py1).clamp_min(1e-4);
let gw = (&gx2 - &gx1).clamp_min(1e-4);
let gh = (&gy2 - &gy1).clamp_min(1e-4);
let ratio_diff = (&gw / &gh).atan() - (&pw / &ph).atan();
let v = &ratio_diff * &ratio_diff * (4.0 / (std::f64::consts::PI * std::f64::consts::PI));
let alpha = &v / ((1.0 - &iou) + &v + 1e-7);
1.0 - &iou + (&rho2 / &c2) + (&alpha * &v)
}
fn dfl_element(dist: &Tensor, target: &Tensor, pos_w: &Tensor) -> Tensor {
let size = dist.size();
let (n, h, w) = (size[0], size[2], size[3]);
let logp = dist
.reshape([n, 4, REG_MAX, h, w])
.log_softmax(2, Kind::Float);
let tl = target.floor(); let trf = &tl + 1.0;
let wl = &trf - target; let wr = target - &tl; let mut total: Option<Tensor> = None;
for k in 0..4i64 {
let side = logp.select(1, k); let tl_idx = tl.select(1, k).to_kind(Kind::Int64).unsqueeze(1);
let tr_idx = trf.select(1, k).to_kind(Kind::Int64).unsqueeze(1);
let ce_l = side.gather(1, &tl_idx, false).neg(); let ce_r = side.gather(1, &tr_idx, false).neg(); let term =
(wl.select(1, k).unsqueeze(1) * &ce_l + wr.select(1, k).unsqueeze(1) * &ce_r) * pos_w;
total = Some(match total {
None => term,
Some(t) => t + term,
});
}
let num = total.expect("至少一条边").sum(Kind::Float);
let den = pos_w.sum(Kind::Float).clamp_min(1.0);
num / den
}
struct LevelSnap {
stride: f32,
h: usize,
w: usize,
pred_boxes: Vec<Vec<[f32; 4]>>,
scores: Vec<Vec<Vec<f32>>>,
}
type SnapPerImage = (Vec<[f32; 4]>, Vec<Vec<f32>>);
fn snapshot_level(
cls: &Tensor,
box_raw: &Tensor,
stride: f32,
gts: &[Vec<[f32; 4]>],
gt_labels: &[Vec<u32>],
) -> LevelSnap {
let size = cls.size();
let h = size[2] as usize;
let w = size[3] as usize;
let n = gts.len();
let cells = h * w;
let (sig, bx) = tch::no_grad(|| {
(
cls.sigmoid().to_device(Device::Cpu),
box_raw.to_device(Device::Cpu),
)
});
let c_len = sig.size()[1] as usize;
let bx_v = tensor_to_vec_f32(&bx);
let sig_v = tensor_to_vec_f32(&sig);
let bidx = |ni: usize, k: usize, hi: usize, wi: usize| ((ni * 4 + k) * h + hi) * w + wi;
let per_image: Vec<SnapPerImage> = (0..n)
.into_par_iter()
.map(|ni| {
let mut pb = vec![[0f32; 4]; cells];
for hi in 0..h {
for wi in 0..w {
let cell = hi * w + wi;
let tx = bx_v[bidx(ni, 0, hi, wi)];
let ty = bx_v[bidx(ni, 1, hi, wi)];
let tw = bx_v[bidx(ni, 2, hi, wi)];
let th = bx_v[bidx(ni, 3, hi, wi)];
let cx = (wi as f32 + 0.5 + tx.tanh()) * stride;
let cy = (hi as f32 + 0.5 + ty.tanh()) * stride;
let bw = tw.exp() * stride;
let bh = th.exp() * stride;
pb[cell] = [cx - bw / 2.0, cy - bh / 2.0, cx + bw / 2.0, cy + bh / 2.0];
}
}
let mut sc = Vec::with_capacity(gt_labels[ni].len());
for &label in gt_labels[ni].iter() {
let base = (ni * c_len + label as usize) * cells;
sc.push(sig_v[base..base + cells].to_vec());
}
(pb, sc)
})
.collect();
let (pred_boxes, scores): (Vec<_>, Vec<_>) = per_image.into_iter().unzip();
LevelSnap {
stride,
h,
w,
pred_boxes,
scores,
}
}
impl ClassifyModel {
pub fn logits(&self, x: &Tensor) -> AvResult<Tensor> {
Ok(self.head.logits(&self.backbone.forward_pooled(x)?))
}
pub fn loss(&self, x: &Tensor, batch: &TrainBatch) -> AvResult<Tensor> {
let TrainBatch::Classify { labels } = batch else {
return Err(AvError::train("分类模型收到非分类批数据"));
};
Ok(self.logits(x)?.cross_entropy_for_logits(labels))
}
pub fn predict(&self, x: &Tensor) -> AvResult<(Vec<u32>, Vec<f32>)> {
tch::no_grad(|| {
let logits = self.logits(x)?;
let labels = logits.argmax(-1, false);
let probs = logits.softmax(-1, Kind::Float);
let conf = probs
.gather(-1, &labels.reshape([-1, 1]), false)
.reshape([-1]);
let labels = tensor_to_vec_i64(&labels);
Ok((
labels.into_iter().map(|v| v as u32).collect(),
tensor_to_vec_f32(&conf),
))
})
}
}
type LevelRaw = Vec<(i64, Tensor, Tensor, Tensor)>;
type ObbLevelRaw = Vec<(i64, Tensor, Tensor, Tensor, Tensor)>;
impl DetectModel {
fn forward_levels(&self, x: &Tensor) -> AvResult<LevelRaw> {
let py = self.backbone.forward_features(x)?;
let feats: Vec<&Tensor> = self
.head_levels
.iter()
.map(|&s| pyramid_level(&py, s as i64).map(|l| &l.tensor))
.collect::<AvResult<Vec<&Tensor>>>()?;
let outs = self.head.forward(&feats);
Ok(outs
.into_iter()
.zip(self.head_levels.iter())
.map(|((cls, dist, _theta), &s)| {
let box_raw = dfl_project(&dist);
(s as i64, cls, dist, box_raw)
})
.collect())
}
fn forward_levels_obb(&self, x: &Tensor) -> AvResult<ObbLevelRaw> {
let py = self.backbone.forward_features(x)?;
let feats: Vec<&Tensor> = self
.head_levels
.iter()
.map(|&s| pyramid_level(&py, s as i64).map(|l| &l.tensor))
.collect::<AvResult<Vec<&Tensor>>>()?;
let outs = self.head.forward(&feats);
let mut out = Vec::with_capacity(self.head_levels.len());
for ((cls, dist, theta), &s) in outs.into_iter().zip(self.head_levels.iter()) {
let theta = theta.ok_or_else(|| AvError::shape("OBB 模型缺少角度分支输出"))?;
let box_raw = dfl_project(&dist);
out.push((s as i64, cls, box_raw, theta, dist));
}
Ok(out)
}
pub(crate) fn raw_preds(&self, x: &Tensor) -> AvResult<Vec<(i64, Tensor, Tensor)>> {
Ok(self
.forward_levels(x)?
.into_iter()
.map(|(s, cls, _dist, box_raw)| (s, cls, box_raw))
.collect())
}
pub fn loss(&self, x: &Tensor, batch: &TrainBatch) -> AvResult<Tensor> {
match batch {
TrainBatch::Obb { boxes, labels } => self.loss_obb(x, boxes, labels),
TrainBatch::Detect { boxes, labels } => {
if self.obb.is_some() {
return Err(AvError::train(
"OBB 模型(obb_mode=true)需要 TrainBatch::Obb(含角度 gt)",
));
}
if self.assigner.eq_ignore_ascii_case("tal") {
self.loss_tal(x, boxes, labels)
} else {
self.loss_center_l1(x, boxes, labels)
}
}
TrainBatch::Classify { .. } => Err(AvError::train("检测模型收到非检测批数据")),
TrainBatch::Seg { .. } => Err(AvError::train(
"检测模型收到分割批数据(TrainBatch::Seg 需要 TaskCfg::Seg 模型)",
)),
TrainBatch::Keypoint { .. } => Err(AvError::train(
"检测模型收到关键点批数据(TrainBatch::Keypoint 需要 TaskCfg::Keypoint 模型)",
)),
}
}
fn loss_tal(
&self,
x: &Tensor,
boxes: &[Vec<[f32; 4]>],
labels: &[Vec<u32>],
) -> AvResult<Tensor> {
let n = x.size()[0] as usize;
let device = x.device();
let tal_cfg = TalConfig::default();
let c_len_total = self.head.num_classes as usize;
let mut gts: Vec<Vec<[f32; 4]>> = vec![Vec::new(); n];
let mut gt_labels: Vec<Vec<u32>> = vec![Vec::new(); n];
for i in 0..n.min(boxes.len()).min(labels.len()) {
for (b, &l) in boxes[i].iter().zip(&labels[i]) {
if (l as usize) < c_len_total {
gts[i].push(*b);
gt_labels[i].push(l);
}
}
}
let levels = self.forward_levels(x)?;
let snaps: Vec<LevelSnap> = levels
.iter()
.map(|(stride, cls, _dist, box_raw)| {
snapshot_level(cls, box_raw, *stride as f32, >s, >_labels)
})
.collect();
let cell_offsets: Vec<usize> = {
let mut offs = Vec::with_capacity(snaps.len() + 1);
let mut acc = 0usize;
for snap in &snaps {
offs.push(acc);
acc += snap.h * snap.w;
}
offs
};
let n_cells_total = *cell_offsets.last().expect("至少一层");
let mut all_centers: Vec<[f32; 2]> = Vec::with_capacity(n_cells_total * 2);
for snap in &snaps {
for hi in 0..snap.h {
for wi in 0..snap.w {
all_centers.push([
(wi as f32 + 0.5) * snap.stride,
(hi as f32 + 0.5) * snap.stride,
]);
}
}
}
let assignments: Vec<Vec<Option<assigner::PosCell>>> = gts
.par_iter()
.enumerate()
.map(|(ni, gt_i)| {
let mut all_boxes: Vec<[f32; 4]> = Vec::with_capacity(all_centers.len());
let mut rows: Vec<Vec<f32>> = Vec::new();
for snap in &snaps {
all_boxes.extend_from_slice(&snap.pred_boxes[ni]);
}
let g_cnt = snaps[0].scores[ni].len();
for g in 0..g_cnt {
let mut row = Vec::with_capacity(all_centers.len());
for snap in &snaps {
row.extend_from_slice(&snap.scores[ni][g]);
}
rows.push(row);
}
let row_slices: Vec<&[f32]> = rows.iter().map(|r| r.as_slice()).collect();
assigner::assign_single_image(&all_boxes, &all_centers, &row_slices, gt_i, &tal_cfg)
})
.collect();
let mut total: Option<Tensor> = None;
for (li, (_stride, cls, dist, box_raw)) in levels.iter().enumerate() {
let size = cls.size();
let (c_len, h, w) = (size[1] as usize, size[2] as usize, size[3] as usize);
let s = snaps[li].stride;
let cells = h * w;
let cell_base = cell_offsets[li];
let mut cls_t = vec![0f32; n * c_len * cells];
let mut cls_w = vec![1f32; n * c_len * cells];
let mut gt_box = vec![0f32; n * 4 * cells];
let mut pos_w = vec![0f32; n * cells];
let mut dfl_t = vec![DFL_SHIFT; n * 4 * cells];
let mut pos_cnt = 0usize;
let mut pos_cls_idx: Vec<usize> = Vec::new();
let clamp_dfl = |v: f32| v.clamp(1e-4, REG_MAX as f32 - 1.0 - 1e-4);
for (gi, asg) in assignments.iter().enumerate() {
for (cell, pos) in asg.iter().enumerate() {
let Some(p) = pos else { continue };
let local = cell as isize - cell_base as isize;
if local < 0 || local as usize >= cells {
continue; }
let local = local as usize;
let (hi, wi) = (local / w, local % w);
let flat2 = gi * cells + hi * w + wi;
let gt = gts[gi][p.gt];
let (gcx, gcy) = ((gt[0] + gt[2]) / 2.0, (gt[1] + gt[3]) / 2.0);
let gw = (gt[2] - gt[0]).max(1e-3);
let gh = (gt[3] - gt[1]).max(1e-3);
let cls_flat =
(gi * c_len + gt_labels[gi][p.gt] as usize) * cells + hi * w + wi;
cls_t[cls_flat] = p.weight;
pos_cls_idx.push(cls_flat);
pos_w[flat2] = p.weight;
pos_cnt += 1;
for (k, v) in gt.iter().enumerate() {
gt_box[(gi * 4 + k) * cells + hi * w + wi] = *v;
}
let offx = gcx / s - (wi as f32 + 0.5);
let offy = gcy / s - (hi as f32 + 0.5);
let t = [
clamp_dfl(atanh_clamp(offx) + DFL_SHIFT),
clamp_dfl(atanh_clamp(offy) + DFL_SHIFT),
clamp_dfl((gw / s).ln() + DFL_SHIFT),
clamp_dfl((gh / s).ln() + DFL_SHIFT),
];
for (k, tv) in t.iter().enumerate() {
dfl_t[(gi * 4 + k) * cells + hi * w + wi] = *tv;
}
}
}
if pos_cnt > 0 {
let total_elems = n * c_len * cells;
let boost = (((total_elems - pos_cnt) as f32) / (pos_cnt as f32)).clamp(1.0, 50.0);
for &idx in &pos_cls_idx {
cls_w[idx] = boost;
}
}
let shape_c = [n as i64, c_len as i64, h as i64, w as i64];
let cls_t_t = Tensor::from_slice(&cls_t)
.to_device(device)
.reshape(shape_c);
let cls_w_t = Tensor::from_slice(&cls_w)
.to_device(device)
.reshape(shape_c);
let cls_loss = cls.binary_cross_entropy_with_logits(
&cls_t_t,
Some(&cls_w_t),
None::<&Tensor>,
tch::Reduction::Mean,
);
let shape_b = [n as i64, 4i64, h as i64, w as i64];
let gt_box_t = Tensor::from_slice(>_box)
.to_device(device)
.reshape(shape_b);
let pred_xyxy = decode_pred_xyxy(box_raw, s);
let ciou_elem = ciou_element(&pred_xyxy, >_box_t); let pos_w_t = Tensor::from_slice(&pos_w)
.to_device(device)
.reshape([n as i64, 1i64, h as i64, w as i64]);
let ciou_loss =
(&ciou_elem * &pos_w_t).sum(Kind::Float) / pos_w_t.sum(Kind::Float).clamp_min(1.0);
let dfl_t_t = Tensor::from_slice(&dfl_t)
.to_device(device)
.reshape(shape_b);
let dfl_loss = dfl_element(dist, &dfl_t_t, &pos_w_t);
let level_loss = &cls_loss * self.loss_w_cls
+ &(&ciou_loss * self.loss_w_ciou)
+ &(&dfl_loss * self.loss_w_dfl);
if let Some(dbg) = LOSS_DEBUG.with(|d| d.take()) {
eprintln!(
"[loss-debug] s={s} cls={:.4} ciou={:.4} dfl={:.4} pos={pos_cnt} {dbg}",
cls_loss.double_value(&[]),
ciou_loss.double_value(&[]),
dfl_loss.double_value(&[]),
);
}
total = Some(match total {
None => level_loss,
Some(t) => t + level_loss,
});
}
total.ok_or_else(|| AvError::train("检测损失为空:无特征层"))
}
fn loss_obb(
&self,
x: &Tensor,
boxes: &[Vec<[f32; 5]>],
labels: &[Vec<u32>],
) -> AvResult<Tensor> {
let obb = self
.obb
.ok_or_else(|| AvError::train("非 OBB 模型收到 TrainBatch::Obb"))?;
let n = x.size()[0] as usize;
let device = x.device();
let tal_cfg = TalConfig::default();
let c_len_total = self.head.num_classes as usize;
let mut gts5: Vec<Vec<[f32; 5]>> = vec![Vec::new(); n];
let mut gts_env: Vec<Vec<[f32; 4]>> = vec![Vec::new(); n];
let mut gt_labels: Vec<Vec<u32>> = vec![Vec::new(); n];
for i in 0..n.min(boxes.len()).min(labels.len()) {
for (b, &l) in boxes[i].iter().zip(&labels[i]) {
if (l as usize) >= c_len_total {
continue; }
let (cx, cy, w, h) = (b[0], b[1], b[2].max(1e-3), b[3].max(1e-3));
let th = obb.angle_domain.normalize(b[4]);
gts5[i].push([cx, cy, w, h, th]);
let (hw, hh) = envelope_half_extents(w, h, th);
gts_env[i].push([cx - hw, cy - hh, cx + hw, cy + hh]);
gt_labels[i].push(l);
}
}
let levels = self.forward_levels_obb(x)?;
let snaps: Vec<LevelSnap> = levels
.iter()
.map(|(stride, cls, box_raw, _theta, _dist)| {
snapshot_level(cls, box_raw, *stride as f32, >s_env, >_labels)
})
.collect();
let cell_offsets: Vec<usize> = {
let mut offs = Vec::with_capacity(snaps.len() + 1);
let mut acc = 0usize;
for snap in &snaps {
offs.push(acc);
acc += snap.h * snap.w;
}
offs
};
let n_cells_total = *cell_offsets.last().expect("至少一层");
let mut all_centers: Vec<[f32; 2]> = Vec::with_capacity(n_cells_total * 2);
for snap in &snaps {
for hi in 0..snap.h {
for wi in 0..snap.w {
all_centers.push([
(wi as f32 + 0.5) * snap.stride,
(hi as f32 + 0.5) * snap.stride,
]);
}
}
}
let assignments: Vec<Vec<Option<assigner::PosCell>>> = gts_env
.par_iter()
.enumerate()
.map(|(ni, gt_i)| {
let mut all_boxes: Vec<[f32; 4]> = Vec::with_capacity(all_centers.len());
let mut rows: Vec<Vec<f32>> = Vec::new();
for snap in &snaps {
all_boxes.extend_from_slice(&snap.pred_boxes[ni]);
}
let g_cnt = snaps[0].scores[ni].len();
for g in 0..g_cnt {
let mut row = Vec::with_capacity(all_centers.len());
for snap in &snaps {
row.extend_from_slice(&snap.scores[ni][g]);
}
rows.push(row);
}
let row_slices: Vec<&[f32]> = rows.iter().map(|r| r.as_slice()).collect();
assigner::assign_single_image(&all_boxes, &all_centers, &row_slices, gt_i, &tal_cfg)
})
.collect();
let mut total: Option<Tensor> = None;
for (li, (_stride, cls, box_raw, t_theta, dist)) in levels.iter().enumerate() {
let size = cls.size();
let (c_len, h, w) = (size[1] as usize, size[2] as usize, size[3] as usize);
let s = snaps[li].stride;
let cells = h * w;
let cell_base = cell_offsets[li];
let mut cls_t = vec![0f32; n * c_len * cells];
let mut cls_w = vec![1f32; n * c_len * cells];
let mut gt5 = vec![0f32; n * 5 * cells];
for gi in 0..n {
for hi in 0..h {
for wi in 0..w {
let base = gi * 5 * cells;
gt5[base + hi * w + wi] = (wi as f32 + 0.5) * s;
gt5[base + cells + hi * w + wi] = (hi as f32 + 0.5) * s;
gt5[base + 2 * cells + hi * w + wi] = 1.0;
gt5[base + 3 * cells + hi * w + wi] = 1.0;
gt5[base + 4 * cells + hi * w + wi] = 0.0;
}
}
}
let mut pos_w = vec![0f32; n * cells];
let mut dfl_t = vec![DFL_SHIFT; n * 4 * cells];
let mut pos_cnt = 0usize;
let mut pos_cls_idx: Vec<usize> = Vec::new();
let clamp_dfl = |v: f32| v.clamp(1e-4, REG_MAX as f32 - 1.0 - 1e-4);
for (gi, asg) in assignments.iter().enumerate() {
for (cell, pos) in asg.iter().enumerate() {
let Some(p) = pos else { continue };
let local = cell as isize - cell_base as isize;
if local < 0 || local as usize >= cells {
continue; }
let local = local as usize;
let (hi, wi) = (local / w, local % w);
let flat = hi * w + wi;
let gt = gts5[gi][p.gt];
let (gcx, gcy, gw, gh) = (gt[0], gt[1], gt[2], gt[3]);
let cls_flat = (gi * c_len + gt_labels[gi][p.gt] as usize) * cells + flat;
cls_t[cls_flat] = p.weight;
pos_cls_idx.push(cls_flat);
pos_w[gi * cells + flat] = p.weight;
pos_cnt += 1;
for (k, v) in gt.iter().enumerate() {
gt5[(gi * 5 + k) * cells + flat] = *v;
}
let offx = gcx / s - (wi as f32 + 0.5);
let offy = gcy / s - (hi as f32 + 0.5);
let t = [
clamp_dfl(atanh_clamp(offx) + DFL_SHIFT),
clamp_dfl(atanh_clamp(offy) + DFL_SHIFT),
clamp_dfl((gw / s).ln() + DFL_SHIFT),
clamp_dfl((gh / s).ln() + DFL_SHIFT),
];
for (k, tv) in t.iter().enumerate() {
dfl_t[(gi * 4 + k) * cells + flat] = *tv;
}
}
}
if pos_cnt > 0 {
let total_elems = n * c_len * cells;
let boost = (((total_elems - pos_cnt) as f32) / (pos_cnt as f32)).clamp(1.0, 50.0);
for &idx in &pos_cls_idx {
cls_w[idx] = boost;
}
}
let shape_c = [n as i64, c_len as i64, h as i64, w as i64];
let cls_t_t = Tensor::from_slice(&cls_t)
.to_device(device)
.reshape(shape_c);
let cls_w_t = Tensor::from_slice(&cls_w)
.to_device(device)
.reshape(shape_c);
let cls_loss = cls.binary_cross_entropy_with_logits(
&cls_t_t,
Some(&cls_w_t),
None::<&Tensor>,
tch::Reduction::Mean,
);
let gt5_t = Tensor::from_slice(>5)
.to_device(device)
.reshape([n as i64, 5i64, h as i64, w as i64]);
let pred_cxcywh = decode_pred_cxcywh(box_raw, s);
let pred_theta = decode_theta(t_theta, obb.angle_domain);
let pred5 = Tensor::cat(&[&pred_cxcywh, &pred_theta], 1);
let kf_elem = kfiou_element(&pred5, >5_t, KfTransform::Log1p); let pos_w_hw = Tensor::from_slice(&pos_w)
.to_device(device)
.reshape([n as i64, h as i64, w as i64]);
let kf_loss =
(&kf_elem * &pos_w_hw).sum(Kind::Float) / pos_w_hw.sum(Kind::Float).clamp_min(1.0);
let shape_b = [n as i64, 4i64, h as i64, w as i64];
let dfl_t_t = Tensor::from_slice(&dfl_t)
.to_device(device)
.reshape(shape_b);
let dfl_loss = dfl_element(dist, &dfl_t_t, &pos_w_hw.unsqueeze(1));
let level_loss = &cls_loss * self.loss_w_cls
+ &(&kf_loss * (self.loss_w_ciou * LOSS_W_KFIOU_SCALE))
+ &(&dfl_loss * self.loss_w_dfl);
if let Some(dbg) = LOSS_DEBUG.with(|d| d.take()) {
eprintln!(
"[loss-debug][obb] s={s} cls={:.4} kfiou={:.4} dfl={:.4} pos={pos_cnt} {dbg}",
cls_loss.double_value(&[]),
kf_loss.double_value(&[]),
dfl_loss.double_value(&[]),
);
}
total = Some(match total {
None => level_loss,
Some(t) => t + level_loss,
});
}
total.ok_or_else(|| AvError::train("OBB 损失为空:无特征层"))
}
fn loss_center_l1(
&self,
x: &Tensor,
boxes: &[Vec<[f32; 4]>],
labels: &[Vec<u32>],
) -> AvResult<Tensor> {
let n = x.size()[0] as usize;
let device = x.device();
let head_levels = self.head_levels.clone();
let breaks = level_breaks(self.img_size, &head_levels);
let mut total: Option<Tensor> = None;
for (s, cls, box_raw) in self
.forward_levels(x)?
.iter()
.map(|(s, c, _d, b)| (s, c, b))
{
let size = cls.size();
let (c_len, h, w) = (size[1] as usize, size[2] as usize, size[3] as usize);
let s = *s as f32;
let mut cls_target = vec![0f32; n * c_len * h * w];
let mut cls_weight = vec![1f32; n * c_len * h * w];
let mut box_target = vec![0f32; n * 4 * h * w];
let mut pos_cnt = 0usize;
for (gi, (img_boxes, img_labels)) in boxes.iter().zip(labels.iter()).enumerate().take(n)
{
for (gt, &label) in img_boxes.iter().zip(img_labels.iter()) {
let (x1, y1, x2, y2) = (gt[0], gt[1], gt[2], gt[3]);
let (cx, cy) = ((x1 + x2) / 2.0, (y1 + y2) / 2.0);
let side = (x2 - x1).max(y2 - y1);
let s_sel = select_level(&head_levels, &breaks, side) as f32;
if s_sel != s {
continue;
}
if label as usize >= c_len {
continue; }
pos_cnt += 1;
let wi = ((cx / s) as usize).min(w - 1);
let hi = ((cy / s) as usize).min(h - 1);
let cw = ((gi * c_len + label as usize) * h + hi) * w + wi;
cls_target[cw] = 1.0;
cls_weight[cw] = 50.0; let offx = cx / s - (wi as f32 + 0.5);
let offy = cy / s - (hi as f32 + 0.5);
let tbox = [
atanh_clamp(offx),
atanh_clamp(offy),
((x2 - x1) / s).ln(),
((y2 - y1) / s).ln(),
];
for (k, tv) in tbox.into_iter().enumerate() {
box_target[((gi * 4 + k) * h + hi) * w + wi] = tv;
}
}
}
let shape = [n as i64, c_len as i64, h as i64, w as i64];
let cls_t = Tensor::from_slice(&cls_target)
.to_device(device)
.reshape(shape);
let cls_w = Tensor::from_slice(&cls_weight)
.to_device(device)
.reshape(shape);
let cls_loss = cls.binary_cross_entropy_with_logits(
&cls_t,
Some(&cls_w),
None::<&Tensor>,
tch::Reduction::Mean,
);
let shape4 = [n as i64, 4i64, h as i64, w as i64];
let box_t = Tensor::from_slice(&box_target)
.to_device(device)
.reshape(shape4);
let box_mask = Tensor::from_slice(&{
let mut m = vec![0f32; n * 4 * h * w];
for (gi, (img_boxes, _)) in boxes.iter().zip(labels.iter()).enumerate().take(n) {
for gt in img_boxes {
let (cx, cy) = ((gt[0] + gt[2]) / 2.0, (gt[1] + gt[3]) / 2.0);
let side = (gt[2] - gt[0]).max(gt[3] - gt[1]);
let s_sel = select_level(&head_levels, &breaks, side) as f32;
if s_sel != s {
continue;
}
let wi = ((cx / s) as usize).min(w - 1);
let hi = ((cy / s) as usize).min(h - 1);
for k in 0..4 {
m[((gi * 4 + k) * h + hi) * w + wi] = 1.0;
}
}
}
m
})
.to_device(device)
.reshape(shape4);
let scale = if pos_cnt == 0 {
0.0
} else {
(n * 4 * h * w) as f32 / (pos_cnt * 4) as f32
};
let box_loss = (box_raw - &box_t).abs() * &box_mask;
let box_loss =
box_loss.mean_dim(&[0i64, 1, 2, 3][..], false, Kind::Float) * &Tensor::from(scale);
let level_loss = &cls_loss + &(&box_loss * &Tensor::from(5f32));
if let Some(dbg) = LOSS_DEBUG.with(|d| d.take()) {
eprintln!(
"[loss-debug] s={s} cls={:.4} box={:.4} (scale={scale:.0}) pos={pos_cnt} {}",
cls_loss.double_value(&[]),
box_loss.double_value(&[]),
dbg
);
}
total = Some(match total {
None => level_loss,
Some(t) => &t + &level_loss,
});
}
total.ok_or_else(|| AvError::train("检测损失为空:无特征层"))
}
pub fn predict(&self, x: &Tensor, conf: f32, iou: f32) -> AvResult<Vec<Vec<Detection>>> {
tch::no_grad(|| {
if let Some(obb) = self.obb {
let n = x.size()[0] as usize;
let mut all: Vec<Vec<Detection>> = vec![Vec::new(); n];
for (s, cls, box_raw, t_theta, _dist) in self.forward_levels_obb(x)? {
for (i, dets) in
decode_level_obb(s, &cls, &box_raw, &t_theta, conf, obb.angle_domain)?
.into_iter()
.enumerate()
{
all[i].extend(dets);
}
}
for dets in all.iter_mut() {
*dets = rotate_nms(std::mem::take(dets), iou, RotNmsMetric::Polygon);
}
return Ok(all);
}
let n = x.size()[0] as usize;
let mut all: Vec<Vec<Detection>> = vec![Vec::new(); n];
for (s, cls, box_raw) in self.raw_preds(x)? {
for (i, dets) in decode_level(s, &cls, &box_raw, conf)?
.into_iter()
.enumerate()
{
all[i].extend(dets);
}
}
for dets in all.iter_mut() {
*dets = nms(std::mem::take(dets), iou);
}
Ok(all)
})
}
}
fn decode_level_obb(
s: i64,
cls: &Tensor,
box_raw: &Tensor,
t_theta: &Tensor,
conf: f32,
domain: AngleDomain,
) -> AvResult<Vec<Vec<Detection>>> {
let size = cls.size();
let (n, c, h, w) = (
size[0] as usize,
size[1] as usize,
size[2] as usize,
size[3] as usize,
);
let probs_v = tensor_to_vec_f32(&cls.sigmoid());
let boxes_v = tensor_to_vec_f32(box_raw);
let theta_v = tensor_to_vec_f32(t_theta);
let pidx = |ni: usize, ci: usize, hi: usize, wi: usize| ((ni * c + ci) * h + hi) * w + wi;
let bidx = |ni: usize, k: usize, hi: usize, wi: usize| ((ni * 4 + k) * h + hi) * w + wi;
let tidx = |ni: usize, hi: usize, wi: usize| (ni * h + hi) * w + wi;
let sf = s as f32;
let half_pi = std::f32::consts::FRAC_PI_2;
let mut out = vec![Vec::new(); n];
for ni in 0..n {
for hi in 0..h {
for wi in 0..w {
let (best_ci, best_p) = (0..c)
.map(|ci| (ci, probs_v[pidx(ni, ci, hi, wi)]))
.max_by(|a, b| a.1.total_cmp(&b.1))
.unwrap_or((0, 0.0));
if best_p.is_nan() || best_p < conf {
continue;
}
let tx = boxes_v[bidx(ni, 0, hi, wi)];
let ty = boxes_v[bidx(ni, 1, hi, wi)];
let tw = boxes_v[bidx(ni, 2, hi, wi)];
let th = boxes_v[bidx(ni, 3, hi, wi)];
let cx = (wi as f32 + 0.5 + tx.tanh()) * sf;
let cy = (hi as f32 + 0.5 + ty.tanh()) * sf;
let bw = tw.exp() * sf;
let bh = th.exp() * sf;
let theta = domain.normalize(theta_v[tidx(ni, hi, wi)].tanh() * half_pi);
out[ni].push(Detection {
bbox: Aabb::new(cx - bw / 2.0, cy - bh / 2.0, cx + bw / 2.0, cy + bh / 2.0),
score: best_p,
class_id: best_ci as u32,
angle: Some(theta),
keypoints: None,
});
}
}
}
Ok(out)
}
fn decode_level(
s: i64,
cls: &Tensor,
box_raw: &Tensor,
conf: f32,
) -> AvResult<Vec<Vec<Detection>>> {
let size = cls.size();
let (n, c, h, w) = (
size[0] as usize,
size[1] as usize,
size[2] as usize,
size[3] as usize,
);
let probs_v = tensor_to_vec_f32(&cls.sigmoid());
let boxes_v = tensor_to_vec_f32(box_raw);
let pidx = |ni: usize, ci: usize, hi: usize, wi: usize| ((ni * c + ci) * h + hi) * w + wi;
let bidx = |ni: usize, k: usize, hi: usize, wi: usize| ((ni * 4 + k) * h + hi) * w + wi;
let sf = s as f32;
let mut out = vec![Vec::new(); n];
for ni in 0..n {
for hi in 0..h {
for wi in 0..w {
let (best_ci, best_p) = (0..c)
.map(|ci| (ci, probs_v[pidx(ni, ci, hi, wi)]))
.max_by(|a, b| a.1.total_cmp(&b.1))
.unwrap_or((0, 0.0));
if best_p.is_nan() || best_p < conf {
continue;
}
let tx = boxes_v[bidx(ni, 0, hi, wi)];
let ty = boxes_v[bidx(ni, 1, hi, wi)];
let tw = boxes_v[bidx(ni, 2, hi, wi)];
let th = boxes_v[bidx(ni, 3, hi, wi)];
let cx = (wi as f32 + 0.5 + tx.tanh()) * sf;
let cy = (hi as f32 + 0.5 + ty.tanh()) * sf;
let bw = tw.exp() * sf;
let bh = th.exp() * sf;
out[ni].push(Detection {
bbox: Aabb::from_xywh(cx - bw / 2.0, cy - bh / 2.0, bw, bh),
score: best_p,
class_id: best_ci as u32,
angle: None,
keypoints: None,
});
}
}
}
Ok(out)
}
impl SegModel {
pub fn mask_size(&self) -> u32 {
self.img_size / 4
}
pub fn img_size(&self) -> u32 {
self.img_size
}
fn forward_branch(&self, x: &Tensor) -> AvResult<(Tensor, Tensor)> {
let py = self.backbone.forward_features(x)?;
let f16 = pyramid_level(&py, 16)?;
let m = self.mask_size() as i64;
Ok(self.mask_branch.forward(&f16.tensor, (m, m)))
}
pub fn loss(&self, x: &Tensor, batch: &TrainBatch) -> AvResult<Tensor> {
let TrainBatch::Seg { masks, labels } = batch else {
return Err(AvError::train("分割模型收到非分割批数据"));
};
let n = x.size()[0] as usize;
let device = x.device();
let (proto, coefcls) = self.forward_branch(x)?;
let cs = coefcls.size();
let gh = cs[2];
let gw = cs[3];
let c = self.num_classes;
let k = self.mask_branch.num_protos;
let ps = proto.size();
let (mh, mw) = (ps[2], ps[3]);
let cells = (gh * gw) as usize;
let total_elems = n * (c * gh * gw) as usize;
let mut cls_t = vec![0f32; total_elems];
let mut cls_w = vec![1f32; total_elems];
let mut pos_idx: Vec<usize> = Vec::new();
struct MaskInst {
img: i64,
hi: i64,
wi: i64,
gt: Vec<u8>,
}
let mut insts: Vec<MaskInst> = Vec::new();
for i in 0..n.min(masks.len()).min(labels.len()) {
for (g, &label) in labels[i].iter().enumerate() {
let Some(gt_mask) = masks[i].get(g) else {
continue;
};
if (label as i64) >= c {
continue; }
if gt_mask.len() != (mh * mw) as usize {
continue; }
let (mut sx, mut sy, mut area) = (0f32, 0f32, 0usize);
for (pi, &v) in gt_mask.iter().enumerate() {
if v != 0 {
sx += (pi % mw as usize) as f32;
sy += (pi / mw as usize) as f32;
area += 1;
}
}
if area == 0 {
continue;
}
let (cx, cy) = (sx / area as f32, sy / area as f32);
let cell_w = mw as f32 / gw as f32;
let cell_h = mh as f32 / gh as f32;
let wi = ((cx / cell_w) as usize).min(gw as usize - 1);
let hi = ((cy / cell_h) as usize).min(gh as usize - 1);
let flat = hi * gw as usize + wi;
let cls_flat = (i * c as usize + label as usize) * cells + flat;
cls_t[cls_flat] = 1.0;
pos_idx.push(cls_flat);
insts.push(MaskInst {
img: i as i64,
hi: hi as i64,
wi: wi as i64,
gt: gt_mask.clone(),
});
}
}
let mut mask_losses: Vec<Tensor> = Vec::new();
if !insts.is_empty() {
let proto_sig = proto.sigmoid(); let plane = (mh * mw) as usize;
let mut gt_flat = vec![0f32; insts.len() * plane];
for (gi, inst) in insts.iter().enumerate() {
for (pj, &v) in inst.gt.iter().enumerate() {
gt_flat[gi * plane + pj] = v as f32;
}
}
let gt_all = Tensor::from_slice(>_flat).to_device(device).reshape([
insts.len() as i64,
mh,
mw,
]);
for (gi, inst) in insts.iter().enumerate() {
let MaskInst { img, hi, wi, .. } = inst;
let coef_map = coefcls.select(0, *img).narrow(0, c, k); let coef = coef_map.select(1, *hi).select(1, *wi); let proto_i = proto_sig.select(0, *img); let logit = (&proto_i * &coef.reshape([k, 1i64, 1i64])).sum_dim_intlist(
&[0i64][..],
false,
Kind::Float,
);
let gt_t = gt_all.select(0, gi as i64); let bce = logit.binary_cross_entropy_with_logits(
>_t,
None::<&Tensor>,
None::<&Tensor>,
tch::Reduction::Mean,
);
let dl = dice_loss(&logit.sigmoid(), >_t);
mask_losses.push(&bce * self.loss_w_bce + &dl * self.loss_w_dice);
}
}
if !pos_idx.is_empty() {
let boost =
(((total_elems - pos_idx.len()) as f32) / (pos_idx.len() as f32)).clamp(1.0, 50.0);
for &idx in &pos_idx {
cls_w[idx] = boost;
}
}
let shape_c = [n as i64, c, gh, gw];
let cls_t_t = Tensor::from_slice(&cls_t)
.to_device(device)
.reshape(shape_c);
let cls_w_t = Tensor::from_slice(&cls_w)
.to_device(device)
.reshape(shape_c);
let cls_loss = coefcls.slice(1, 0, c, 1).binary_cross_entropy_with_logits(
&cls_t_t,
Some(&cls_w_t),
None::<&Tensor>,
tch::Reduction::Mean,
);
let mut total = cls_loss;
if !mask_losses.is_empty() {
let cnt = mask_losses.len() as f64;
let mut acc: Option<Tensor> = None;
for l in mask_losses {
acc = Some(match acc {
None => l,
Some(t) => t + l,
});
}
let mask_mean = acc.expect("mask_losses 非空") / cnt;
total += mask_mean;
}
Ok(total * self.loss_weight)
}
pub fn predict(&self, x: &Tensor, conf: f32, nms_iou: f32) -> AvResult<Vec<Vec<SegInstance>>> {
tch::no_grad(|| {
let (proto, coefcls) = self.forward_branch(x)?;
let n = x.size()[0] as usize;
let cs = coefcls.size();
let (c, k, gh, gw) = (self.num_classes, self.mask_branch.num_protos, cs[2], cs[3]);
let ps = proto.size();
let (mh, mw) = (ps[2] as usize, ps[3] as usize);
let plane = (mh * mw) as i64;
let cells = (gh * gw) as usize;
let device = proto.device();
let proto_sig = proto.sigmoid(); let cls_sig = coefcls.slice(1, 0, c, 1).sigmoid(); let coef_map = coefcls
.slice(1, c, c + k, 1)
.reshape([n as i64, k, cells as i64]);
let mut out: Vec<Vec<SegInstance>> = vec![Vec::new(); n];
for (ni, out_ni) in out.iter_mut().enumerate() {
let (scores, cls_idx) = cls_sig.select(0, ni as i64).max_dim(0, false); let scores = scores.reshape([cells as i64]);
let cls_idx = cls_idx.reshape([cells as i64]);
let cand_all = scores
.ge(tch::Scalar::from(conf as f64))
.nonzero()
.select(1, 0); let m_all = cand_all.size()[0] as usize;
if m_all == 0 {
continue;
}
let cand_scores = tensor_to_vec_f32(&scores.index_select(0, &cand_all));
let mut order: Vec<usize> = (0..m_all).collect();
order.sort_by(|&a, &b| cand_scores[b].total_cmp(&cand_scores[a]));
order.truncate(MAX_SEGS_PER_IMAGE);
let idx_top =
Tensor::from_slice(&(order.iter().map(|&p| p as i64).collect::<Vec<_>>()[..]))
.to_device(device);
let top_scores = tensor_to_vec_f32(&scores.index_select(0, &idx_top));
let top_cls = tensor_to_vec_i64(&cls_idx.index_select(0, &idx_top));
let cand_coef = coef_map
.select(0, ni as i64)
.index_select(1, &idx_top)
.transpose(0, 1); let proto_n = proto_sig.select(0, ni as i64).reshape([k, plane]);
let logits = cand_coef.matmul(&proto_n); let masks_v = tensor_to_vec_u8(
&logits
.greater(tch::Scalar::from(0.0f64))
.to_kind(Kind::Uint8),
);
let plane_u = plane as usize;
let mut insts: Vec<SegInstance> = Vec::with_capacity(order.len());
for (j, mask) in masks_v.chunks_exact(plane_u).enumerate() {
if !mask.contains(&1) {
continue; }
insts.push(SegInstance {
label: top_cls[j] as u32,
score: top_scores[j],
mask: mask.to_vec(),
});
}
*out_ni = mask_nms(insts, nms_iou);
}
Ok(out)
})
}
}
impl KeypointModel {
pub fn img_size(&self) -> u32 {
self.img_size
}
pub fn kp_stride(&self) -> u32 {
crate::keypoint::KP_HEAD_STRIDE
}
fn forward_head(&self, x: &Tensor) -> AvResult<Tensor> {
let py = self.backbone.forward_features(x)?;
let f8 = pyramid_level(&py, self.kp_stride() as i64)?;
Ok(self.head.forward(&f8.tensor))
}
pub fn loss(&self, x: &Tensor, batch: &TrainBatch) -> AvResult<Tensor> {
let TrainBatch::Keypoint {
boxes,
kpts,
labels: _,
} = batch
else {
return Err(AvError::train("关键点模型收到非关键点批数据"));
};
let n = x.size()[0] as usize;
let device = x.device();
let k = self.head.num_keypoints as usize;
let out = self.forward_head(x)?;
let size = out.size();
let (h, w) = (size[2] as usize, size[3] as usize);
let s = self.kp_stride() as f32;
let cells = h * w;
let ku = 2 * k;
let mut cls_t = vec![0f32; n * cells];
let mut cls_w = vec![1f32; n * cells];
let mut box_t = vec![0f32; n * 4 * cells];
let mut pos_w = vec![0f32; n * cells]; let mut off_t = vec![0f32; n * ku * cells];
let mut off_w = vec![0f32; n * ku * cells]; let mut vis_t = vec![0f32; n * k * cells];
let mut vis_w = vec![0f32; n * k * cells]; let mut gtx = vec![0f32; n * k * cells]; let mut gty = vec![0f32; n * k * cells];
let mut vis_map = vec![0f32; n * k * cells];
let mut scale_map = vec![1f32; n * cells]; let mut oks_pos = vec![0f32; n * cells]; let mut pos_cnt = 0usize;
let mut pos_cls_idx: Vec<usize> = Vec::new();
for (gi, (img_boxes, img_kpts)) in boxes.iter().zip(kpts.iter()).enumerate().take(n) {
for (gt, gk) in img_boxes.iter().zip(img_kpts.iter()) {
if gk.len() != k {
return Err(AvError::data(format!(
"关键点实例点数 {} 与模型 num_keypoints = {k} 不一致(检查数据与配置)",
gk.len()
)));
}
let (gcx, gcy) = (gt[0], gt[1]);
let (gw, gh) = (gt[2].max(1e-3), gt[3].max(1e-3));
let wi = ((gcx / s) as usize).min(w - 1);
let hi = ((gcy / s) as usize).min(h - 1);
let flat = hi * w + wi;
let cell = gi * cells + flat;
cls_t[cell] = 1.0;
pos_cls_idx.push(cell);
pos_cnt += 1;
pos_w[cell] = 1.0;
scale_map[cell] = (gw * gh).sqrt().max(1e-3);
let tb = [
atanh_clamp(gcx / s - (wi as f32 + 0.5)),
atanh_clamp(gcy / s - (hi as f32 + 0.5)),
(gw / s).ln(),
(gh / s).ln(),
];
for (c, tv) in tb.iter().enumerate() {
box_t[(gi * 4 + c) * cells + flat] = *tv;
}
let mut any_vis = false;
for (j, kp) in gk.iter().enumerate() {
let vis = kp[2] > 0.0;
any_vis |= vis;
let vw = if vis { 1.0 } else { 0.0 };
let off_base = (gi * ku + 2 * j) * cells + flat;
off_t[off_base] = kp[0] / s - (wi as f32 + 0.5);
off_w[off_base] = vw;
let off_base = off_base + cells;
off_t[off_base] = kp[1] / s - (hi as f32 + 0.5);
off_w[off_base] = vw;
let vi = (gi * k + j) * cells + flat;
vis_t[vi] = vw;
vis_w[vi] = 1.0; gtx[vi] = kp[0];
gty[vi] = kp[1];
vis_map[vi] = vw;
}
if any_vis {
oks_pos[cell] = 1.0;
}
}
}
if pos_cnt > 0 {
let total_elems = n * cells;
let boost = (((total_elems - pos_cnt) as f32) / (pos_cnt as f32)).clamp(1.0, 50.0);
for &idx in &pos_cls_idx {
cls_w[idx] = boost;
}
}
let shape_c = [n as i64, 1i64, h as i64, w as i64];
let cls_t_t = Tensor::from_slice(&cls_t)
.to_device(device)
.reshape(shape_c);
let cls_w_t = Tensor::from_slice(&cls_w)
.to_device(device)
.reshape(shape_c);
let loss_cls = out
.narrow(1, crate::keypoint::KP_CLS_CH, 1)
.binary_cross_entropy_with_logits(
&cls_t_t,
Some(&cls_w_t),
None::<&Tensor>,
tch::Reduction::Mean,
);
if pos_cnt == 0 {
return Ok(loss_cls * self.loss_weight);
}
let shape_b = [n as i64, 4i64, h as i64, w as i64];
let box_t_t = Tensor::from_slice(&box_t)
.to_device(device)
.reshape(shape_b);
let pos_w_t = Tensor::from_slice(&pos_w)
.to_device(device)
.reshape([n as i64, 1i64, h as i64, w as i64]);
let box_l1 = (out.narrow(1, crate::keypoint::KP_BOX_CH, 4) - &box_t_t).abs() * &pos_w_t;
let box_l1 = box_l1.sum(Kind::Float) / (&pos_w_t * 4.0).sum(Kind::Float).clamp_min(1.0);
let shape_o = [n as i64, ku as i64, h as i64, w as i64];
let off_t_t = Tensor::from_slice(&off_t)
.to_device(device)
.reshape(shape_o);
let off_w_t = Tensor::from_slice(&off_w)
.to_device(device)
.reshape(shape_o);
let off_l1 =
(out.narrow(1, crate::keypoint::KP_OFF_CH, ku as i64) - &off_t_t).abs() * &off_w_t;
let off_l1 = off_l1.sum(Kind::Float) / off_w_t.sum(Kind::Float).clamp_min(1.0);
let shape_v = [n as i64, k as i64, h as i64, w as i64];
let vis_t_t = Tensor::from_slice(&vis_t)
.to_device(device)
.reshape(shape_v);
let vis_w_t = Tensor::from_slice(&vis_w)
.to_device(device)
.reshape(shape_v);
let loss_vis = out
.narrow(1, crate::keypoint::KP_OFF_CH + ku as i64, k as i64)
.binary_cross_entropy_with_logits(
&vis_t_t,
Some(&vis_w_t),
None::<&Tensor>,
tch::Reduction::Mean,
);
let gtx_t = Tensor::from_slice(>x).to_device(device).reshape(shape_v);
let gty_t = Tensor::from_slice(>y).to_device(device).reshape(shape_v);
let vis_map_t = Tensor::from_slice(&vis_map)
.to_device(device)
.reshape(shape_v);
let scale_map_t = Tensor::from_slice(&scale_map)
.to_device(device)
.reshape([n as i64, 1i64, h as i64, w as i64]);
let oks_pos_t = Tensor::from_slice(&oks_pos)
.to_device(device)
.reshape([n as i64, 1i64, h as i64, w as i64]);
let sigma = Tensor::from_slice(&sigma_table(k))
.to_device(device)
.reshape([1i64, k as i64, 1i64, 1i64]);
let xs =
Tensor::arange(w as i64, (Kind::Float, device)).reshape([1i64, 1i64, 1i64, w as i64]);
let ys =
Tensor::arange(h as i64, (Kind::Float, device)).reshape([1i64, 1i64, h as i64, 1i64]);
let off_raw = out
.narrow(1, crate::keypoint::KP_OFF_CH, ku as i64)
.reshape([n as i64, k as i64, 2i64, h as i64, w as i64]);
let dx = off_raw.select(2, 0); let dy = off_raw.select(2, 1);
let px = (&xs + 0.5) * (s as f64) + &dx * (s as f64);
let py = (&ys + 0.5) * (s as f64) + &dy * (s as f64);
let ex = &px - >x_t;
let ey = &py - >y_t;
let d2 = &ex * &ex + &ey * &ey;
let denom = &(&(&scale_map_t * &scale_map_t) * 2.0) * &sigma * σ let e = (&d2 / &denom).neg().exp();
let num = (&e * &vis_map_t).sum_dim_intlist(&[1i64][..], false, Kind::Float); let den = vis_map_t
.sum_dim_intlist(&[1i64][..], false, Kind::Float)
.clamp_min(1.0);
let ok_map = &num / &den;
let ok_mean =
(&ok_map * &oks_pos_t).sum(Kind::Float) / oks_pos_t.sum(Kind::Float).clamp_min(1.0);
let loss_oks = ok_mean * -1.0 + 1.0;
if let Some(dbg) = LOSS_DEBUG.with(|d| d.take()) {
eprintln!(
"[loss-debug][kp] cls={:.4} box={:.4} off={:.4} vis={:.4} oks={:.4} pos={pos_cnt} {dbg}",
loss_cls.double_value(&[]),
box_l1.double_value(&[]),
off_l1.double_value(&[]),
loss_vis.double_value(&[]),
loss_oks.double_value(&[]),
);
}
Ok((&loss_cls * LOSS_W_KP_CLS
+ &(&box_l1 * LOSS_W_KP_BOX)
+ &(&off_l1 * LOSS_W_KP_OFF)
+ &(&loss_vis * LOSS_W_KP_VIS)
+ &(&loss_oks * self.loss_w_oks))
* self.loss_weight)
}
pub fn predict(&self, x: &Tensor, conf: f32, iou: f32) -> AvResult<Vec<Vec<Detection>>> {
tch::no_grad(|| {
let out = self.forward_head(x)?;
let k = self.head.num_keypoints as usize;
let size = out.size();
let (n, h, w) = (size[0] as usize, size[2] as usize, size[3] as usize);
let sf = self.kp_stride() as f32;
let lim = self.img_size as f32;
let cls_v = tensor_to_vec_f32(&out.narrow(1, crate::keypoint::KP_CLS_CH, 1).sigmoid());
let box_v = tensor_to_vec_f32(&out.narrow(1, crate::keypoint::KP_BOX_CH, 4));
let off_v = tensor_to_vec_f32(&out.narrow(1, crate::keypoint::KP_OFF_CH, ku_of(k)));
let vis_v = tensor_to_vec_f32(
&out.narrow(1, crate::keypoint::KP_OFF_CH + 2 * k as i64, k as i64)
.sigmoid(),
);
let cidx = |ni: usize, hi: usize, wi: usize| (ni * h + hi) * w + wi;
let bidx = |ni: usize, c: usize, hi: usize, wi: usize| ((ni * 4 + c) * h + hi) * w + wi;
let oidx =
|ni: usize, c: usize, hi: usize, wi: usize| ((ni * 2 * k + c) * h + hi) * w + wi;
let vidx = |ni: usize, j: usize, hi: usize, wi: usize| ((ni * k + j) * h + hi) * w + wi;
let mut out_imgs: Vec<Vec<Detection>> = vec![Vec::new(); n];
for ni in 0..n {
for hi in 0..h {
for wi in 0..w {
let p = cls_v[cidx(ni, hi, wi)];
if p < conf {
continue;
}
let tx = box_v[bidx(ni, 0, hi, wi)];
let ty = box_v[bidx(ni, 1, hi, wi)];
let tw = box_v[bidx(ni, 2, hi, wi)];
let th = box_v[bidx(ni, 3, hi, wi)];
let cx = (wi as f32 + 0.5 + tx.tanh()) * sf;
let cy = (hi as f32 + 0.5 + ty.tanh()) * sf;
let bw = tw.exp() * sf;
let bh = th.exp() * sf;
let mut kps = Vec::with_capacity(k);
for j in 0..k {
let kx = (wi as f32 + 0.5 + off_v[oidx(ni, 2 * j, hi, wi)]) * sf;
let ky = (hi as f32 + 0.5 + off_v[oidx(ni, 2 * j + 1, hi, wi)]) * sf;
let v = if vis_v[vidx(ni, j, hi, wi)] > 0.5 {
2.0
} else {
0.0
};
kps.push([kx.clamp(0.0, lim), ky.clamp(0.0, lim), v]);
}
out_imgs[ni].push(Detection {
bbox: Aabb::new(
cx - bw / 2.0,
cy - bh / 2.0,
cx + bw / 2.0,
cy + bh / 2.0,
),
score: p,
class_id: 0,
angle: None,
keypoints: Some(kps),
});
}
}
let mut dets = std::mem::take(&mut out_imgs[ni]);
dets.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
dets.truncate(MAX_KPS_PER_IMAGE);
out_imgs[ni] = nms(dets, iou);
}
Ok(out_imgs)
})
}
}
fn ku_of(k: usize) -> i64 {
(2 * k) as i64
}
impl TaskModel {
pub fn num_classes(&self) -> i64 {
match self {
TaskModel::Classify(m) => m.head.num_classes,
TaskModel::Detect(m) => m.head.num_classes,
TaskModel::Seg(m) => m.num_classes,
TaskModel::Keypoint(_) => 1,
}
}
pub fn img_size(&self) -> u32 {
match self {
TaskModel::Classify(m) => m.img_size,
TaskModel::Detect(m) => m.img_size,
TaskModel::Seg(m) => m.img_size,
TaskModel::Keypoint(m) => m.img_size,
}
}
pub fn loss(&self, x: &Tensor, batch: &TrainBatch) -> AvResult<Tensor> {
match self {
TaskModel::Classify(m) => m.loss(x, batch),
TaskModel::Detect(m) => m.loss(x, batch),
TaskModel::Seg(m) => m.loss(x, batch),
TaskModel::Keypoint(m) => m.loss(x, batch),
}
}
pub fn predict(&self, x: &Tensor, conf: f32, iou: f32) -> AvResult<PredictOutput> {
match self {
TaskModel::Classify(m) => {
let (labels, confs) = m.predict(x)?;
Ok(PredictOutput::Classify { labels, confs })
}
TaskModel::Detect(m) => Ok(PredictOutput::Detect {
per_image: m.predict(x, conf, iou)?,
}),
TaskModel::Seg(m) => Ok(PredictOutput::Seg {
per_image: m.predict(x, conf, iou)?,
}),
TaskModel::Keypoint(m) => Ok(PredictOutput::Keypoint {
per_image: m.predict(x, conf, iou)?,
}),
}
}
pub fn set_train(&self, train: bool) {
match self {
TaskModel::Classify(m) => m.backbone.set_train(train),
TaskModel::Detect(m) => m.backbone.set_train(train),
TaskModel::Seg(m) => m.backbone.set_train(train),
TaskModel::Keypoint(m) => m.backbone.set_train(train),
}
}
}
#[cfg(all(test, feature = "torch"))]
mod tests {
use super::*;
fn detect_cfg_toml(assigner: &str) -> av_core::config::RunConfig {
let toml = format!(
concat!(
"[model]
backbone = {{ family = \"simple-cnn\", depth = 1.0, width = 1.0, pretrained = \"none\" }}
",
"[data.sources.train]\ndir = \"d\"\n",
"[[model.tasks]]\nkind = \"detect\"\nnum_classes = 2\nimg_size = 64\n",
"assigner = \"{a}\"\n"
),
a = assigner
);
av_core::config::RunConfig::from_toml_str(&toml).expect("smoke 配置必须合法")
}
fn sample_batch() -> TrainBatch {
TrainBatch::Detect {
boxes: vec![
vec![[10.0, 12.0, 30.0, 34.0]],
vec![[20.0, 20.0, 44.0, 41.0]],
],
labels: vec![vec![0], vec![1]],
}
}
fn sample_seg_batch() -> TrainBatch {
let mh = 24usize;
let mut m = vec![0u8; mh * mh];
for y in 4..12 {
for x in 4..12 {
m[y * mh + x] = 1;
}
}
TrainBatch::Seg {
masks: vec![vec![m]],
labels: vec![vec![0]],
}
}
#[test]
fn dfl_project_expectation_matches_onehot() {
let mut v = vec![0f32; (4 * REG_MAX) as usize];
for k in 0..4usize {
v[k * REG_MAX as usize + 8] = 20.0;
}
let dist = Tensor::from_slice(&v).reshape([1i64, 4 * REG_MAX, 1, 1]);
let proj = dfl_project(&dist);
assert!(
(proj.double_value(&[0, 0, 0, 0]) - (8.0 - DFL_SHIFT as f64)).abs() < 1e-3,
"got {}",
proj.double_value(&[0, 0, 0, 0])
);
}
#[test]
fn decode_pred_xyxy_zero_offset_gives_cell_center_box() {
let box_raw = Tensor::zeros([1i64, 4, 2, 3], (Kind::Float, Device::Cpu));
let xy = decode_pred_xyxy(&box_raw, 8.0);
assert!((xy.double_value(&[0, 0, 0, 0]) - 0.0).abs() < 1e-4);
assert!((xy.double_value(&[0, 2, 0, 0]) - 8.0).abs() < 1e-4);
assert!((xy.double_value(&[0, 0, 1, 2]) - 16.0).abs() < 1e-4);
assert!((xy.double_value(&[0, 3, 1, 2]) - 16.0).abs() < 1e-4);
}
#[test]
fn ciou_matches_hand_computed() {
let pred = Tensor::from_slice(&[0.0f32, 0.0, 10.0, 10.0]).reshape([1i64, 4, 1, 1]);
let gt = Tensor::from_slice(&[0.0f32, 0.0, 10.0, 10.0]).reshape([1i64, 4, 1, 1]);
let l = ciou_element(&pred, >);
assert!(
(l.double_value(&[0, 0])).abs() < 1e-4,
"got {}",
l.double_value(&[0, 0])
);
let pred = Tensor::from_slice(&[0.0f32, 0.0, 10.0, 10.0]).reshape([1i64, 4, 1, 1]);
let gt = Tensor::from_slice(&[5.0f32, 5.0, 15.0, 15.0]).reshape([1i64, 4, 1, 1]);
let l = ciou_element(&pred, >);
assert!(
(l.double_value(&[0, 0]) - 0.968254).abs() < 1e-3,
"got {}",
l.double_value(&[0, 0])
);
let pred = Tensor::from_slice(&[0.0f32, 0.0, 20.0, 10.0]).reshape([1i64, 4, 1, 1]);
let gt = Tensor::from_slice(&[0.0f32, 0.0, 10.0, 10.0]).reshape([1i64, 4, 1, 1]);
let l = ciou_element(&pred, >);
assert!(
(l.double_value(&[0, 0]) - 0.553250).abs() < 1e-3,
"got {}",
l.double_value(&[0, 0])
);
}
#[test]
fn dfl_loss_left_right_cross_entropy_hand_computed() {
let mut v = vec![0f32; (4 * REG_MAX) as usize];
for k in 0..4usize {
v[k * REG_MAX as usize + 8] = 10.0;
}
let dist = Tensor::from_slice(&v).reshape([1i64, 4 * REG_MAX, 1, 1]);
let target = Tensor::from_slice(&[7.75f32; 4]).reshape([1i64, 4, 1, 1]);
let pos_w = Tensor::ones([1i64, 1, 1, 1], (Kind::Float, Device::Cpu));
let loss = dfl_element(&dist, &target, &pos_w);
let z = 10f64.exp() + 15.0;
let expected = 4.0 * (z.ln() - 7.5);
assert!(
(loss.double_value(&[]) - expected).abs() < 1e-3,
"got {} expected {expected}",
loss.double_value(&[])
);
}
#[test]
fn tal_loss_smoke_forward_backward() {
let cfg = detect_cfg_toml("tal");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("装配应成功");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let loss = model.loss(&x, &sample_batch()).expect("TAL 损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward(); }
#[test]
fn legacy_l1_loss_smoke_forward_backward() {
let cfg = detect_cfg_toml("center");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("装配应成功");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let loss = model
.loss(&x, &sample_batch())
.expect("L1 旧路径损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward();
}
fn detect_cfg_toml_levels(assigner: &str, head_levels: &[u32]) -> av_core::config::RunConfig {
let levels = head_levels
.iter()
.map(|l| l.to_string())
.collect::<Vec<_>>()
.join(", ");
let toml = format!(
concat!(
"[model]
backbone = {{ family = \"simple-cnn\", depth = 1.0, width = 1.0, pretrained = \"none\" }}
",
"[data.sources.train]\ndir = \"d\"\n",
"[[model.tasks]]\nkind = \"detect\"\nnum_classes = 2\nimg_size = 64\n",
"assigner = \"{a}\"\nhead_levels = [{levels}]\n"
),
a = assigner,
levels = levels
);
av_core::config::RunConfig::from_toml_str(&toml).expect("P2 配置必须合法")
}
#[test]
fn detect_head_levels_default_is_two_tier() {
let cfg = detect_cfg_toml("tal");
let Some(av_core::config::TaskCfg::Detect(d)) = cfg.model.tasks.first() else {
panic!("应为 detect 任务");
};
assert_eq!(d.head_levels, vec![8, 16]);
}
#[test]
fn p2_feature_pyramid_and_head_shape_contract() {
let cfg = detect_cfg_toml_levels("tal", &[4, 8, 16]);
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("装配应成功");
let TaskModel::Detect(m) = &model else {
panic!("应为检测模型");
};
assert_eq!(m.head.strides, vec![4, 8, 16], "头层级应为 P2/P3/P4");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let py = m.backbone.forward_features(&x).expect("骨干前向应成功");
assert_eq!(py.levels.len(), 3);
let strides: Vec<u32> = py.levels.iter().map(|l| l.stride).collect();
assert_eq!(strides, vec![4, 8, 16]);
let vars = vs.variables();
let s4_key = vars
.keys()
.find(|k| k.contains("head.s4.cls1.weight"))
.expect("P2 头变量 head.s4.cls1.weight 应存在")
.clone();
assert_eq!(
vars[&s4_key].size()[1],
m.backbone.stride_channels(4).expect("stride 4 应存在"),
"P2 层输入通道须与骨干 stride 4 特征一致"
);
let raw = m.raw_preds(&x).expect("检测前向应成功");
assert_eq!(raw.len(), 3);
for (li, (s, cls, box_raw)) in raw.iter().enumerate() {
let stride = *s;
assert_eq!(stride as u32, [4u32, 8, 16][li]);
let grid = 64 / stride;
assert_eq!(
cls.size(),
vec![2, 2, grid, grid],
"s{stride} cls 形状应为 [N,C,img/s,img/s]"
);
assert_eq!(
box_raw.size(),
vec![2, 4, grid, grid],
"s{stride} box 形状应为 [N,4,img/s,img/s]"
);
}
let per_image = m.predict(&x, 0.0, 0.5).expect("P2 推理应成功");
assert_eq!(per_image.len(), 2);
}
#[test]
fn default_head_variable_names_unchanged() {
let cfg = detect_cfg_toml("tal");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
build_model(&vs.root(), &cfg).expect("装配应成功");
let names: Vec<String> = vs.variables().keys().cloned().collect();
assert!(names.iter().any(|n| n.contains("head.s8.cls1.weight")));
assert!(names.iter().any(|n| n.contains("head.s16.cls1.weight")));
assert!(!names.iter().any(|n| n.contains("head.s4")), "默认无 P2 层");
}
#[test]
fn level_assignment_boundaries_two_tier_matches_legacy() {
let levels = vec![8u32, 16];
let breaks = level_breaks(64, &levels);
assert_eq!(breaks, vec![16.0, 32.0]);
assert_eq!(select_level(&levels, &breaks, 0.5), 8);
assert_eq!(
select_level(&levels, &breaks, 16.0),
8,
"边界值(含)归低层"
);
assert_eq!(select_level(&levels, &breaks, 16.1), 16);
assert_eq!(select_level(&levels, &breaks, 64.0), 16);
}
#[test]
fn level_assignment_boundaries_three_tier_p2() {
let levels = vec![4u32, 8, 16];
let breaks = level_breaks(640, &levels);
assert_eq!(breaks, vec![80.0, 160.0, 320.0]);
assert_eq!(select_level(&levels, &breaks, 1.0), 4);
assert_eq!(select_level(&levels, &breaks, 80.0), 4, "边界值(含)归 P2");
assert_eq!(select_level(&levels, &breaks, 80.5), 8);
assert_eq!(
select_level(&levels, &breaks, 160.0),
8,
"边界值(含)归中层"
);
assert_eq!(select_level(&levels, &breaks, 160.5), 16);
assert_eq!(select_level(&levels, &breaks, 320.0), 16);
assert_eq!(
select_level(&levels, &breaks, 640.0),
16,
"超出回落最大 stride 层"
);
}
#[test]
fn head_levels_rejects_empty_unsorted_or_duplicated() {
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
for (levels, expect_msg) in [
(&[0u32; 0][..], "不能为空"),
(&[8u32, 8][..], "严格升序"),
(&[16u32, 8][..], "严格升序"),
] {
let cfg = detect_cfg_toml_levels("tal", levels);
let err = match build_model(&vs.root(), &cfg) {
Err(e) => e,
Ok(_) => panic!("head_levels = {levels:?} 应被拒绝"),
};
assert!(
err.to_string().contains(expect_msg),
"head_levels = {levels:?} 应报「{expect_msg}」,got: {err}"
);
}
let cfg = detect_cfg_toml_levels("tal", &[4, 32]);
let err = match build_model(&vs.root(), &cfg) {
Err(e) => e,
Ok(_) => panic!("不支持的 stride 应被拒绝"),
};
assert!(
err.to_string().contains("stride 32"),
"不支持的 stride 应报骨干形状错,got: {err}"
);
}
#[test]
fn p2_tal_loss_smoke_forward_backward() {
let cfg = detect_cfg_toml_levels("tal", &[4, 8, 16]);
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("装配应成功");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let loss = model.loss(&x, &sample_batch()).expect("P2 TAL 损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward(); }
#[test]
fn p2_center_l1_loss_smoke_forward_backward() {
let cfg = detect_cfg_toml_levels("center", &[4, 8, 16]);
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("装配应成功");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let loss = model.loss(&x, &sample_batch()).expect("P2 L1 损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward();
}
#[test]
fn resnet18_detect_backbone_wiring() {
let toml = concat!(
"[model]\n",
"backbone = { family = \"resnet18\", depth = 1.0, width = 1.0, pretrained = \"none\" }\n",
"neck = { type = \"identity\", channels = [] }\n",
"[data.sources.train]\ndir = \"d\"\n",
"[[model.tasks]]\nkind = \"detect\"\nhead = \"yolo\"\n",
"num_classes = 2\nimg_size = 96\nassigner = \"tal\"\n"
);
let cfg = av_core::config::RunConfig::from_toml_str(toml).expect("resnet18 检测配置应合法");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("resnet18 检测装配应成功");
let TaskModel::Detect(m) = &model else {
panic!("应为检测模型");
};
assert_eq!(m.backbone.stride_channels(4).unwrap(), 64, "layer1");
assert_eq!(m.backbone.stride_channels(8).unwrap(), 128, "layer2");
assert_eq!(m.backbone.stride_channels(16).unwrap(), 256, "layer3");
let vars = vs.variables();
assert!(vars.contains_key("backbone.layer1.0.conv1.weight"));
assert!(vars.contains_key("backbone.layer4.1.bn2.bias"));
assert!(
!vars.keys().any(|n| n.starts_with("backbone.c1")),
"检测 resnet18 不应出现 simple-cnn 变量名"
);
let s8 = vars
.iter()
.find(|(n, _)| n.contains("head.s8.cls1.weight"))
.expect("s8 头应存在");
assert_eq!(s8.1.size()[1], 128, "s8 头输入通道应等于骨干 stride 8 通道");
let x = Tensor::randn([2, 3, 96, 96], (Kind::Float, Device::Cpu));
let loss = model
.loss(&x, &sample_batch())
.expect("resnet18 检测损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward();
}
#[test]
fn resnet18_seg_backbone_wiring() {
let toml = concat!(
"[model]
",
"backbone = { family = \"resnet18\", depth = 1.0, width = 1.0, pretrained = \"none\" }
",
"[[model.tasks]]
kind = \"seg\"
head = \"yolact\"
",
"num_classes = 2
img_size = 96
",
"[data.sources.train]
dir = \"d\"
"
);
let cfg = av_core::config::RunConfig::from_toml_str(toml).expect("resnet18 seg 配置应合法");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("resnet18 seg 装配应成功");
let TaskModel::Seg(m) = &model else {
panic!("应为分割模型");
};
assert_eq!(m.backbone.stride_channels(16).unwrap(), 256, "layer3");
let vars = vs.variables();
assert!(vars.contains_key("backbone.layer1.0.conv1.weight"));
assert!(
!vars.keys().any(|n| n.starts_with("backbone.c1")),
"seg resnet18 不应出现 simple-cnn 变量名"
);
let x = Tensor::randn([2, 3, 96, 96], (Kind::Float, Device::Cpu));
let batch = sample_seg_batch();
let loss = model.loss(&x, &batch).expect("resnet18 seg 损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward();
}
fn obb_cfg_toml() -> av_core::config::RunConfig {
let toml = concat!(
"[model]
backbone = { family = \"simple-cnn\", depth = 1.0, width = 1.0, pretrained = \"none\" }
",
"[data.sources.train]\ndir = \"d\"\n",
"[[model.tasks]]\nkind = \"detect\"\nnum_classes = 2\nimg_size = 64\n",
"assigner = \"tal\"\nobb_mode = true\n"
);
av_core::config::RunConfig::from_toml_str(toml).expect("obb smoke 配置必须合法")
}
fn obb_sample_batch() -> TrainBatch {
TrainBatch::Obb {
boxes: vec![
vec![[20.0, 20.0, 16.0, 8.0, 0.5]],
vec![[32.0, 30.0, 12.0, 12.0, -0.6]],
],
labels: vec![vec![0], vec![1]],
}
}
#[test]
fn obb_loss_smoke_forward_backward() {
let cfg = obb_cfg_toml();
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("OBB 模型装配应成功");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let loss = model.loss(&x, &obb_sample_batch()).expect("OBB 损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward(); }
#[test]
fn obb_p2_loss_smoke_forward_backward() {
let toml = concat!(
"[model]
backbone = { family = \"simple-cnn\", depth = 1.0, width = 1.0, pretrained = \"none\" }
",
"[data.sources.train]\ndir = \"d\"\n",
"[[model.tasks]]\nkind = \"detect\"\nnum_classes = 2\nimg_size = 64\n",
"assigner = \"tal\"\nobb_mode = true\nhead_levels = [4, 8, 16]\n"
);
let cfg = av_core::config::RunConfig::from_toml_str(toml).expect("配置必须合法");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("OBB P2 模型装配应成功");
let TaskModel::Detect(m) = &model else {
panic!("应为检测模型");
};
assert_eq!(m.head.strides, vec![4, 8, 16]);
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let loss = model
.loss(&x, &obb_sample_batch())
.expect("OBB P2 损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward();
}
#[test]
fn obb_batch_type_mismatch_is_rejected() {
let cfg = obb_cfg_toml();
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let obb_model = build_model(&vs.root(), &cfg).expect("装配应成功");
assert!(obb_model
.loss(
&Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu)),
&sample_batch()
)
.is_err());
let vs2 = tch::nn::VarStore::new(tch::Device::Cpu);
let det_model = build_model(&vs2.root(), &detect_cfg_toml("tal")).expect("装配应成功");
assert!(det_model
.loss(
&Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu)),
&obb_sample_batch()
)
.is_err());
}
#[test]
fn obb_predict_outputs_angle_and_rot_nms_runs() {
let cfg = obb_cfg_toml();
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("OBB 模型装配应成功");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let TaskModel::Detect(m) = &model else {
panic!("OBB 配置应装配出检测模型");
};
let per_image = m.predict(&x, 0.0, 0.5).expect("OBB 推理应成功");
assert_eq!(per_image.len(), 2);
assert!(per_image.iter().any(|d| !d.is_empty()), "conf=0 应给出候选");
let (lo, hi) = av_core::conventions::AngleDomain::Le90.range();
for dets in &per_image {
for d in dets {
let th = d.angle.expect("OBB 候选必须带角度");
assert!(th >= lo - 1e-4 && th < hi, "θ={th} 越出 le90 域");
assert!(d.bbox.w() > 0.0 && d.bbox.h() > 0.0, "外接框须合法");
}
}
}
#[test]
fn obb_bs2_overfit_both_images_regression() {
let cfg = obb_cfg_toml();
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("OBB 模型装配应成功");
let TaskModel::Detect(m) = &model else {
panic!("OBB 配置应装配出检测模型");
};
let mut opt = {
use tch::nn::OptimizerConfig;
tch::nn::Adam::default().build(&vs, 3e-3).expect("优化器")
};
let boxes: Vec<Vec<[f32; 5]>> = vec![
vec![[20.0, 20.0, 16.0, 8.0, 0.5]],
vec![[44.0, 40.0, 10.0, 18.0, -0.9]],
];
let labels = vec![vec![0u32], vec![1u32]];
let x0 = Tensor::randn([3, 64, 64], (Kind::Float, Device::Cpu)) * 0.1;
let x1 = Tensor::randn([3, 64, 64], (Kind::Float, Device::Cpu)) * 0.1 + 0.5;
let x = Tensor::stack(&[x0, x1], 0);
let batch = TrainBatch::Obb {
boxes: boxes.clone(),
labels: labels.clone(),
};
for _ in 0..250 {
let loss = m.loss(&x, &batch).expect("OBB 损失应成功");
opt.zero_grad();
loss.backward();
opt.clip_grad_norm(10.0);
opt.step();
}
use av_core::geometry::RotBox;
tch::no_grad(|| {
for (gi, gts) in boxes.iter().enumerate() {
let xi = x.copy().slice(0, gi as i64, gi as i64 + 1, 1);
let dets = m.predict(&xi, 0.25, 0.5).expect("OBB 推理应成功");
let gr = RotBox {
cx: gts[0][0],
cy: gts[0][1],
w: gts[0][2],
h: gts[0][3],
theta: gts[0][4],
};
let best = dets[0]
.iter()
.map(|d| {
gr.iou(&RotBox {
cx: (d.bbox.x1 + d.bbox.x2) / 2.0,
cy: (d.bbox.y1 + d.bbox.y2) / 2.0,
w: d.bbox.x2 - d.bbox.x1,
h: d.bbox.y2 - d.bbox.y1,
theta: d.angle.unwrap_or(0.0),
})
})
.fold(0.0f32, f32::max);
assert!(
best > 0.5,
"图 {gi} 过拟合后最优旋转 IoU = {best}(bs>1 批索引回归?)"
);
}
});
}
#[test]
fn standalone_obb_task_reports_guidance() {
let toml = concat!(
"[model]
backbone = { family = \"simple-cnn\", depth = 1.0, width = 1.0, pretrained = \"none\" }
",
"[data.sources.train]\ndir = \"d\"\n",
"[[model.tasks]]\nkind = \"obb\"\n"
);
let cfg = av_core::config::RunConfig::from_toml_str(toml).expect("配置应合法");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let err = match build_model(&vs.root(), &cfg) {
Err(e) => e,
Ok(_) => panic!("独立 obb 任务应报接入指引错误"),
};
assert!(
err.to_string().contains("obb_mode"),
"错误信息应指向 detect + obb_mode,got: {err}"
);
}
fn seg_cfg_toml(head: &str) -> av_core::config::RunConfig {
let toml = format!(
concat!(
"[model]
backbone = {{ family = \"simple-cnn\", depth = 1.0, width = 1.0, pretrained = \"none\" }}
",
"[data.sources.train]\ndir = \"d\"\n",
"[[model.tasks]]\nkind = \"seg\"\nhead = \"{h}\"\n",
"num_classes = 3\nimg_size = 64\nnum_protos = 8\n"
),
h = head
);
av_core::config::RunConfig::from_toml_str(&toml).expect("seg smoke 配置必须合法")
}
fn seg_sample_batch() -> TrainBatch {
let fill_rect = |m: &mut Vec<u8>, (x0, y0, x1, y1): (usize, usize, usize, usize)| {
for y in y0..y1 {
for x in x0..x1 {
m[y * 16 + x] = 1;
}
}
};
let mut m0 = vec![0u8; 256];
fill_rect(&mut m0, (3, 3, 9, 9));
let mut m1 = vec![0u8; 256];
fill_rect(&mut m1, (1, 1, 7, 7));
let mut m2 = vec![0u8; 256];
fill_rect(&mut m2, (8, 8, 14, 14));
TrainBatch::Seg {
masks: vec![vec![m0], vec![m1, m2]],
labels: vec![vec![0], vec![1, 2]],
}
}
#[test]
fn seg_loss_smoke_forward_backward() {
let cfg = seg_cfg_toml("yolact");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("Seg 模型装配应成功");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let loss = model.loss(&x, &seg_sample_batch()).expect("Seg 损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward(); }
#[test]
fn seg_loss_empty_batch_stays_connected() {
let cfg = seg_cfg_toml("yolact");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("Seg 模型装配应成功");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let batch = TrainBatch::Seg {
masks: vec![Vec::new(), Vec::new()],
labels: vec![Vec::new(), Vec::new()],
};
let loss = model.loss(&x, &batch).expect("空批损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward();
}
#[test]
fn seg_predict_outputs_binary_masks_with_nms() {
let cfg = seg_cfg_toml("yolact");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("Seg 模型装配应成功");
let TaskModel::Seg(m) = &model else {
panic!("seg 配置应装配出 SegModel");
};
assert_eq!(m.mask_size(), 16);
let x = Tensor::randn([1, 3, 64, 64], (Kind::Float, Device::Cpu));
let out = m.predict(&x, 0.0, 0.5).expect("Seg 推理应成功");
assert_eq!(out.len(), 1);
for inst in &out[0] {
assert_eq!(inst.mask.len(), 256);
assert!(inst.mask.iter().all(|&v| v <= 1), "掩码应二值");
assert!(inst.mask.contains(&1), "空掩码应被丢弃");
assert!((0.0..=1.0).contains(&inst.score), "score 应在 [0,1]");
assert!((inst.label as i64) < 3, "label 应在类数内");
}
let kept = &out[0];
for (a, b) in kept
.iter()
.enumerate()
.flat_map(|(i, a)| kept[i + 1..].iter().map(move |b| (a, b)))
{
if a.label == b.label {
assert!(
mask_iou(&a.mask, &b.mask) < 0.5,
"NMS 后同类掩码 IoU 应 < 0.5"
);
}
}
}
#[test]
fn mask_nms_suppresses_same_class_near_duplicates() {
let inst = |label: u32, score: f32, m: Vec<u8>| SegInstance {
label,
score,
mask: m,
};
let a = inst(5, 0.9, vec![1, 1, 1, 1]); let b = inst(5, 0.8, vec![1, 1, 0, 0]);
let c = inst(5, 0.7, vec![0, 0, 0, 1]); let d = inst(6, 0.6, vec![1, 1, 1, 1]); let kept = mask_nms(vec![b, c, d, a], 0.5);
assert_eq!(kept.len(), 3, "b 应被 a 抑制,c/d 保留");
assert_eq!(kept[0].score, 0.9);
assert!(kept.iter().any(|k| k.label == 6), "异类不应被抑制");
assert!(!kept.iter().any(|k| k.score == 0.8), "近重复同类应被抑制");
assert!(mask_nms(Vec::new(), 0.5).is_empty());
let solo = inst(5, 0.9, vec![1, 1, 1, 1]);
assert_eq!(mask_nms(vec![solo], 0.5).len(), 1);
}
#[test]
fn seg_batch_type_mismatch_is_rejected() {
let cfg = seg_cfg_toml("yolact");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let seg_model = build_model(&vs.root(), &cfg).expect("装配应成功");
assert!(seg_model
.loss(
&Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu)),
&sample_batch()
)
.is_err());
let vs2 = tch::nn::VarStore::new(tch::Device::Cpu);
let det_model = build_model(&vs2.root(), &detect_cfg_toml("tal")).expect("装配应成功");
assert!(det_model
.loss(
&Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu)),
&seg_sample_batch()
)
.is_err());
}
#[test]
fn seg_direct_head_reports_not_supported() {
let cfg = seg_cfg_toml("direct");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let err = match build_model(&vs.root(), &cfg) {
Err(e) => e,
Ok(_) => panic!("direct 精度档应报未支持"),
};
assert!(err.to_string().contains("direct"), "got: {err}");
}
#[test]
fn seg_overfit_tiny_batch_reduces_loss() {
use tch::nn::OptimizerConfig;
let cfg = seg_cfg_toml("yolact");
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("装配应成功");
let s = 64usize;
let mut buf = vec![0f32; 2 * 3 * s * s];
let mut paint = |img: usize, ch: usize, (x0, y0, x1, y1): (usize, usize, usize, usize)| {
for y in y0..y1 {
for x in x0..x1 {
buf[((img * 3 + ch) * s + y) * s + x] += 1.5;
}
}
};
paint(0, 0, (12, 12, 36, 36)); paint(1, 1, (4, 4, 28, 28)); paint(1, 2, (32, 32, 56, 56)); let x = Tensor::from_slice(&buf)
.to_kind(Kind::Float)
.reshape([2i64, 3, s as i64, s as i64]);
let batch = seg_sample_batch();
let mut opt = tch::nn::Adam::default()
.build(&vs, 1e-2)
.expect("优化器应构建");
let first = model
.loss(&x, &batch)
.expect("损失应成功")
.double_value(&[]);
let TaskModel::Seg(sm) = &model else {
panic!("seg 配置应装配出 SegModel");
};
let mut last = first;
let mut candidates_ok = false;
for step in 0..400 {
opt.zero_grad();
let loss = model.loss(&x, &batch).expect("损失应成功");
loss.backward();
opt.step();
last = loss.double_value(&[]);
if (step + 1) % 40 == 0 {
let preds = sm.predict(&x, 0.0, 0.5).expect("过拟合推理应成功");
candidates_ok = preds.iter().all(|dets| !dets.is_empty());
if candidates_ok {
break;
}
}
}
assert!(last.is_finite(), "损失必须有限");
assert!(
candidates_ok && last < first,
"过拟合后 predict 应每图产出掩码候选且损失下降:first={first} last={last}"
);
}
fn kp_cfg_toml() -> av_core::config::RunConfig {
let toml = concat!(
"[model]
backbone = { family = \"simple-cnn\", depth = 1.0, width = 1.0, pretrained = \"none\" }
",
"[data.sources.train]\ndir = \"d\"\n",
"[[model.tasks]]\nkind = \"keypoint\"\ndecode = \"direct\"\n",
"num_keypoints = 3\nimg_size = 64\n"
);
av_core::config::RunConfig::from_toml_str(toml).expect("kp smoke 配置必须合法")
}
fn kp_sample_batch() -> TrainBatch {
TrainBatch::Keypoint {
boxes: vec![
vec![[32.0, 32.0, 16.0, 16.0]],
vec![[20.0, 40.0, 12.0, 20.0], [44.0, 16.0, 12.0, 12.0]],
],
kpts: vec![
vec![vec![
[30.0, 28.0, 2.0],
[34.0, 28.0, 2.0],
[32.0, 36.0, 1.0],
]],
vec![
vec![[18.0, 32.0, 2.0], [22.0, 32.0, 0.0], [20.0, 44.0, 2.0]],
vec![[42.0, 12.0, 2.0], [46.0, 12.0, 2.0], [44.0, 20.0, 2.0]],
],
],
labels: vec![vec![0], vec![0, 0]],
}
}
#[test]
fn kp_loss_smoke_forward_backward() {
let cfg = kp_cfg_toml();
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("Keypoint 模型装配应成功");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let loss = model
.loss(&x, &kp_sample_batch())
.expect("关键点损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward(); }
#[test]
fn kp_empty_batch_stays_connected() {
let cfg = kp_cfg_toml();
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("装配应成功");
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let batch = TrainBatch::Keypoint {
boxes: vec![Vec::new(), Vec::new()],
kpts: vec![Vec::new(), Vec::new()],
labels: vec![Vec::new(), Vec::new()],
};
let loss = model.loss(&x, &batch).expect("空批损失应成功");
assert!(loss.double_value(&[]).is_finite());
loss.backward();
}
#[test]
fn kp_point_count_mismatch_is_rejected() {
let cfg = kp_cfg_toml();
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("装配应成功");
let batch = TrainBatch::Keypoint {
boxes: vec![vec![[32.0, 32.0, 16.0, 16.0]]],
kpts: vec![vec![vec![[30.0, 28.0, 2.0]]]], labels: vec![vec![0]],
};
let err = model
.loss(
&Tensor::randn([1, 3, 64, 64], (Kind::Float, Device::Cpu)),
&batch,
)
.unwrap_err();
assert!(err.to_string().contains("num_keypoints"), "got: {err}");
}
#[test]
fn kp_batch_type_mismatch_is_rejected() {
let cfg = kp_cfg_toml();
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let kp_model = build_model(&vs.root(), &cfg).expect("装配应成功");
assert!(kp_model
.loss(
&Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu)),
&sample_batch()
)
.is_err());
let vs2 = tch::nn::VarStore::new(tch::Device::Cpu);
let det_model = build_model(&vs2.root(), &detect_cfg_toml("tal")).expect("装配应成功");
assert!(det_model
.loss(
&Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu)),
&kp_sample_batch()
)
.is_err());
}
#[test]
fn kp_predict_outputs_detections_with_keypoints() {
let cfg = kp_cfg_toml();
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("装配应成功");
let TaskModel::Keypoint(m) = &model else {
panic!("keypoint 配置应装配出 KeypointModel");
};
assert_eq!(m.kp_stride(), 8);
let x = Tensor::randn([1, 3, 64, 64], (Kind::Float, Device::Cpu));
let per_image = m.predict(&x, 0.0, 0.5).expect("关键点推理应成功");
assert_eq!(per_image.len(), 1);
for d in &per_image[0] {
let kps = d.keypoints.as_ref().expect("实例必须带关键点");
assert_eq!(kps.len(), 3);
for kp in kps {
assert!(
(0.0..=64.0).contains(&kp[0]) && (0.0..=64.0).contains(&kp[1]),
"关键点应 clamp 在画布内: {kp:?}"
);
assert!(
kp[2] == 0.0 || kp[2] == 2.0,
"v 应二值化为 0/2,got {}",
kp[2]
);
}
assert!(d.bbox.x2 >= d.bbox.x1 && d.bbox.y2 >= d.bbox.y1, "框应合法");
assert!((0.0..=1.0).contains(&d.score), "score 应在 [0,1]");
}
let kept = &per_image[0];
for (a, b) in kept
.iter()
.enumerate()
.flat_map(|(i, a)| kept[i + 1..].iter().map(move |b| (a, b)))
{
assert!(
a.bbox.iou(&b.bbox) <= 0.5 + 1e-6,
"NMS 后框 IoU 应 ≤ 0.5,got {}",
a.bbox.iou(&b.bbox)
);
}
}
#[test]
fn kp_overfit_tiny_batch_reduces_loss() {
use tch::nn::OptimizerConfig;
let cfg = kp_cfg_toml();
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = build_model(&vs.root(), &cfg).expect("装配应成功");
let s = 64usize;
let batch = kp_sample_batch();
let TrainBatch::Keypoint { boxes, kpts, .. } = &batch else {
unreachable!()
};
let mut buf = vec![0f32; 2 * 3 * s * s];
for (img, img_boxes) in boxes.iter().enumerate() {
for (b, gk) in img_boxes.iter().zip(gk_of(&kpts[img])) {
let (x0, y0) = ((b[0] - b[2] / 2.0) as usize, (b[1] - b[3] / 2.0) as usize);
for y in y0..y0 + b[3] as usize {
for x in x0..x0 + b[2] as usize {
buf[((img * 3) * s + y) * s + x] += 1.5;
}
}
for kp in gk {
if kp[2] > 0.0 {
let (ky, kx) = (kp[1] as usize, kp[0] as usize);
buf[((img * 3 + 1) * s + ky) * s + kx] += 2.0;
}
}
}
}
let x = Tensor::from_slice(&buf)
.to_kind(Kind::Float)
.reshape([2i64, 3, s as i64, s as i64]);
let mut opt = tch::nn::Adam::default()
.build(&vs, 1e-2)
.expect("优化器应构建");
let first = model
.loss(&x, &batch)
.expect("损失应成功")
.double_value(&[]);
let mut last = first;
for _ in 0..150 {
opt.zero_grad();
let loss = model.loss(&x, &batch).expect("损失应成功");
loss.backward();
opt.step();
last = loss.double_value(&[]);
}
assert!(last.is_finite(), "损失必须有限");
assert!(
last < first * 0.8,
"过拟合后损失应显著下降:first={first} last={last}"
);
}
fn gk_of(kpts: &[Vec<[f32; 3]>]) -> &[Vec<[f32; 3]>] {
kpts
}
}
#[cfg(test)]
mod backbone_dispatch_tests {
use super::*;
fn cfg_with_backbone(family: &str, task: TaskCfg) -> RunConfig {
RunConfig {
model: av_core::config::ModelConfig {
backbone: av_core::config::BackboneCfg {
family: family.into(),
..av_core::config::BackboneCfg::default()
},
tasks: vec![task],
..av_core::config::ModelConfig::default()
},
..RunConfig::default()
}
}
fn kp_task(img_size: u32) -> TaskCfg {
TaskCfg::Keypoint(av_core::config::KeypointCfg {
img_size,
decode: "direct".into(),
..av_core::config::KeypointCfg::default()
})
}
fn det_task(img_size: u32) -> TaskCfg {
TaskCfg::Detect(av_core::config::DetectCfg {
img_size,
num_classes: 3,
..av_core::config::DetectCfg::default()
})
}
#[test]
fn keypoint_builds_on_resnet18_and_dino() {
let vs = tch::nn::VarStore::new(Device::Cpu);
let cfg = cfg_with_backbone("resnet18", kp_task(224));
let m = build_model(&vs.root(), &cfg).unwrap();
assert!(
matches!(m, TaskModel::Keypoint(_)),
"resnet18 应装配 keypoint"
);
let vs2 = tch::nn::VarStore::new(Device::Cpu);
let cfg = cfg_with_backbone("dinov2", kp_task(448));
let m = build_model(&vs2.root(), &cfg).unwrap();
assert!(
matches!(m, TaskModel::Keypoint(_)),
"dino-v2 应装配 keypoint"
);
}
#[test]
fn detect_builds_on_dino() {
let vs = tch::nn::VarStore::new(Device::Cpu);
let cfg = cfg_with_backbone("dinov2", det_task(448));
let m = build_model(&vs.root(), &cfg).unwrap();
assert!(matches!(m, TaskModel::Detect(_)), "dino-v2 应装配 detect");
}
#[test]
fn dino_detect_rejects_non_448_multiple_imgsz() {
let vs = tch::nn::VarStore::new(Device::Cpu);
let cfg = cfg_with_backbone("dinov2", det_task(640)); assert!(build_model(&vs.root(), &cfg).is_err(), "640 必须被校验拒绝");
}
}