Co-DETR:混合标签指派训练让DETR收敛更快、小目标检测更强
我复现 Co-DETR 那段时间几乎把 DETR 系列的论文和代码翻了个底朝天。这篇论文的标题叫 “DETRs with Collaborative Hybrid Assignments Training”平常大家直接叫它 Co-DETR。它解决的问题非常具体为什么 DETR 类检测器训练速度慢、小目标效果差以及怎么用一套简洁的“混合指派”训练策略把这些短板补上。一句话概括它的价值在不改推理结构、不增加任何推理开销的前提下通过训练阶段的多组辅助查询和一对多标签指派让 DETR 的收敛速度和检测精度都上了一个台阶。这很适合正在做检测模型调优、准备把 DETR 用到真实场景、或者想搞懂端到端检测器训练机制的读者。我下面会把论文思路、关键设计、代码复现和训练自己数据时踩过的坑一次讲清楚。1. 从 DETR 收敛慢说起Co-DETR 到底在解决什么问题1.1 端到端检测的两个老毛病原始的 DETR 把检测当成一个集合预测问题用匈牙利算法做一对一匹配。每个 ground truth 只分配给一个查询其他查询在匹配不到目标时统统算“无对象”。这个设计的优雅之处在于不需要 NMS 和后处理但代价也很大。第一个代价是一对一匹配提供的监督信号太稀疏。一张图假设有 20 个目标但查询有 300 个其中 280 个查询对应的都是背景。它们在训练中得到的梯度几乎都是“你不要输出框”缺少对特征丰富性的挖掘。尤其是训练早期匹配结果不稳定很多查询会长期处于“没匹配上”的状态骨干网络和编码器能拿到的有效梯度很少模型自然收敛慢。第二个代价是正样本数量太少导致训练不稳定。一对一匹配意味着每个目标只有一个正样本查询如果这个查询初始位置不好梯度方向就可能不对。DETR 系列后续工作用了各种招数来缓解比如 DAB-DETR 用锚点坐标作为查询DN-DETR 用去噪任务来稳定匹配Group DETR 用一组查询匹配同一目标。这些方法都是有效果的但各自都有一点“补丁”的味道。1.2 Co-DETR 的核心思路把一对多“塞”进训练里Co-DETR 的思路很直接既然一对一监督信号稀疏那我在训练时干脆加一些一对多的辅助分支。这些辅助分支用类似传统检测器的方式让多个查询同时匹配同一个目标把 dense 的监督信号引入进来。关键是推理时这些辅助分支全部丢弃只保留原始的解码器结构。也就是说你部署出来的模型和标准 DETR 长得一模一样速度一点不受影响。这个想法听起来简单但落地有几个难点。辅助分支不能干扰主分支的一对一匹配。如果主查询和辅助查询一起竞争同一个目标反而可能让训练乱掉。Co-DETR 的方案是把它们分开分别用自己的匹配逻辑互不干扰同时让它们共享骨干和编码器的特征。这样一来辅助分支虽然不参与推理但它们的梯度会回传到骨干和编码器让特征提取器变得更强大。主分支在推理时拿到的特征质量更高精度自然就上去了。这个思路我一开始觉得有点像“知识蒸馏”但仔细看又不一样。蒸馏是一个网络教另一个网络而 Co-DETR 是同一套特征同时喂给多个“学生”大家一起从数据里学再把梯度反馈给特征。更像多任务学习辅助头和主头共享底层各学各的输出头。1.3 为什么说它是“混合指派”论文标题里的 Hybrid Assignments 指的就是训练时同时存在两种标签分配策略。一种是一对一匹配沿用了 DETR 的风格用匈牙利算法或 TopK 来选择匹配的查询。另一种是一对多匹配像传统检测器那样多个 anchor 或查询同时负责一个目标。Co-DETR 里的辅助头使用了多种一对一和一对多混合策略所以叫 Collaborative Hybrid Assignments Training。为什么一对多有用因为密集的监督信号能让编码器更快学到“哪里可能有物体”。传统检测器比如 FCOS 或 ATSS每个目标周围一堆位置都是正样本网络收到的正样本梯度非常多。DETR 类模型天生缺乏这种密集监督Co-DETR 就是在训练阶段补上这个缺口。正是这个机制让它在减少训练轮次的同时还能涨点。2. 方法拆解辅助头、去噪查询和标签分配2.1 辅助头到底做了什么Co-DETR 的整体结构可以分成两大块。主分支是一个标准的 DETR 解码器用的是可变形注意力支持多尺度特征。辅助分支由若干个并行的检测头组成常用的设置是 4 个辅助头。每个辅助头内部有一个自己的解码器或检测头配一组可学习的锚点按一对一或一对多的方式分配标签独立计算损失。这些辅助头的输入是主干网络加编码器输出的多尺度特征输出是各自预测的框和类别。它们的梯度会传到骨干和编码器但辅助头里的查询不会进入主解码器。推理时辅助头整个扔掉模型只有主分支在跑。我最初的一个疑问是这些辅助头的设计怎么选论文给了几种选项比如用类似 Deformable DETR 的小解码器或者直接用密集检测头。实测下来辅助头内部用简单的多头自注意力加可变形注意力就足够了关键是匹配方式要多样。多样性让编码器见到更多不同形态的监督信号学出来的特征更鲁棒。2.2 正负样本匹配策略的多样性Co-DETR 的“混合”具体体现在正负样本定义的不同。辅助头们并不用同一套匹配策略而是互相错开。有的头用 Hungarian 匹配有的头用 TopK 距离匹配有的头用类似 ATSS 的自适应匹配。每个头匹配逻辑不同正样本位置和数量就不一样产生的梯度信号也各有侧重。这种多样性很像 ensemble但 ensemble 是在推理时融合多个模型Co-DETR 是在训练时用多个匹配策略融合监督信号。匹配策略的多样性让编码器不能只应付某一种匹配方式必须学到更加通用的特征表示。我在复现时试过所有辅助头都用同一种匹配方式效果确实会掉一点。2.3 去噪查询和正负样本的平衡DN-DETR 的去噪思想在这里也做了融合。Co-DETR 会在训练时往主解码器里额外塞一些带噪声的 ground truth 查询让模型学习“从被扰乱的框还原出原始框”。这相当于给主分支提供了一个额外的 denoising 任务加速收敛。去噪查询和辅助头的密集监督配合起来训练时模型的收敛速度明显加快。不过这里要小心一个问题去噪查询在训练时参与损失计算但推理时同样会被去掉。它们的存在只为了让解码器更快学到“框回归”的规律而不是成为推理的一部分。这一点和辅助头是异曲同工。正负样本平衡也是隐性需要注意的点。辅助头产生大量正样本如果权重设置不对模型可能偏向简单样本反而影响主分支的一对一匹配。实际训练时每个辅助头的损失权重、类别损失和框损失的比例都要调。论文给出的默认配置能够直接用但换数据集时这个权重需要重新试。3. 代码复现从配置文件到训练日志3.1 环境搭建和依赖版本复现 Co-DETR 需要 PyTorch、MMDetection、MMCV 等组件。我用的环境是 PyTorch 1.13 CUDA 11.7 MMDetection 3.x。这里有个关键点Co-DETR 官方代码基于 MMDetection 3.0 以上版本老版本 2.x 会有接口不兼容的问题比如 Registry 的用法变了Head 的 forward 参数也变了。如果之前只熟 2.x建议先摸一下 3.x 的架构再动手。安装顺序建议先装 MMCV 和 MMDetection再拉官方仓库。注意检查版本对应关系MMCV 2.x 和 MMDetection 3.x 是配套的。我一开始图省事直接 pip install mmdet结果版本不对跑起来各种缺属性浪费了半天。代码目录结构很常规configs 下是各种实验配置projects 里是模型定义。如果只是跑 COCO直接用官网给的脚本就行。如果想训练自己的数据主要改数据集路径和类别数但类别数改动涉及输出层的初始化这个后面细说。3.2 完整训练流程要点假设要训练一个 Co-DETR 的 ResNet-50 版本COCO 数据集已经准备好。训练脚本大概是bash tools/dist_train.sh projects/configs/co_detr/co_deformable_detr_r50_1x_coco.py 8启动之前要确认数据集路径。config 里的 data_root 改成实际路径anno 文件用 coco 格式的 json图片目录按 train/val 组织。训练多卡时要注意 batch size 和 learning rate 的对应。默认配置是按 8 卡、每卡 2 张图设计的如果改用 4 卡batch size 减半学习率建议也减半否则容易发散。训练过程中可以盯几个指标。第一个是 loss 曲线正常情况下总 loss 会稳步下降但下降速度比标准 DETR 快得多尤其是在前几千迭代。第二个是平均召回率辅助头带来的密集监督会让 AR 指标在很早期就升得很快这说明编码器在快速学习。第三个是正样本数量如果看到 loss 一直抖动可能是辅助头的匹配策略或 loss 权重有问题。我自己习惯每隔一段时间导出一次 checkpoint在验证集上看 mAP 的变化趋势。Co-DETR 在 12 epoch 的 1x 设置下50 轮左右就能看到明显效果最后的 mAP 在 COCO val 上大概 42 上下。当然具体数值取决于 backbone、输入尺度和训练轮次。3.3 backbone 选择与参数量Co-DETR 的 backbone 可以换ResNet-50、ResNet-101、Swin-T 甚至 Swin-L 都支持。骨干网络越大辅助头能提供的监督信号收益就越明显因为大模型需要更丰富的梯度来拟合。不过训练显存开销也会明显增加。我用 ResNet-50 时 8 卡训练 1x 大概需要 40G 显存如果单卡显存不够要么减小 batch size要么开启梯度累积。换 backbone 时要注意 config 里 pretrained 路径。官方给的权重链接有时需要科学下载实际使用时可以先下到本地再改路径。很多国内读者卡在权重下载上其实换个思路也可以不加载预训练直接从零训练效果会略差但能先跑通流程。3.4 自定义数据集训练的细节用 Co-DETR 训练自己的数据集最重要的改动是类别数。config 里的 num_classes 要改成实际类别数同时注意数据集的 metainfo。MMDetection 3.x 里数据集的 classes 需要显式声明不声明的话会用默认 COCO 的 80 类训练时会出现 shape 不匹配的错误。第二个细节是 anchor 尺寸或者说查询的初始化。Co-DETR 的辅助头如果用的是类似 RPN 的密集预测那它会用到预设的 anchor。自定义数据集的目标尺寸分布如果和 COCO 差异很大建议先用脚本统计一下数据集里目标的宽高分布再修改 anchor 配置。不做这一步小目标多的数据集上召回会偏低。训练类数少的任务时比如只有 1 个类输出层的初始化要小心。MMDetection 的 head 在初始化时会根据类别数设置 bias但如果是加载了 COCO 预训练权重再 finetune最后的分类层维度对不上需要特殊处理。我在代码里见过官方这样做把分类头的权重删掉随机初始化新类别的输出层。4. 部署与推理如何跑起来不采坑4.1 推理脚本和模型导出Co-DETR 的推理和其他 DETR 类模型一样简单。官方仓库里有 test 脚本指定 config 和 checkpoint 就能验证。推理时辅助头根本不参与模型输出的是一组预测框和类别后处理只需要做一个简单的阈值过滤不需要 NMS。因为一对一匹配已经保证了每个目标只有一个框。如果你想部署到 TensorRT 或者 ONNX需要把模型结构中的一些动态操作处理好。常见的问题是 query 数量是固定的比如默认是 300 个但输入图片尺寸是动态的。做 ONNX 导出时建议固定输入尺寸不然动态 shape 会让某些算子变得很复杂。我实际试过导出 ONNX 再转 TensorRT主要难点在于可变形注意力的 CUDA 算子。新版 mmdeploy 已经支持了但如果你用的版本比较老可能得自己写插件。老实说如果只是普通服务部署直接用 PyTorch 的 TorchScript 也可以一次 batch 推理的耗时在几毫秒到十几毫秒之间取决于 GPU。4.2 和传统检测器部署的差异如果你之前部署的是 YOLO 或者 Faster R-CNN切到 Co-DETR 会有一个感受少了对 NMS 和后处理参数的调试。YOLO 要调 NMS 阈值、conf 阈值Co-DETR 只要选一个 conf 阈值比如 0.3然后按分数排序输出 top100 就行。这是端到端模型部署最舒服的地方。但是要注意推理结果的可复现性。DETR 类模型在推理时如果开着可变形注意力的 CUDA 算子两次推理结果可能有一点点浮点误差这在大部分业务场景里无所谓但如果你要做严格的回归测试建议固定 seed 和 deterministic 模式。供电和内存方面辅助头虽然在推理时被丢弃但模型文件里还是会保存它们的权重所以 checkpoint 会比标准 DETR 大一些。我导出生产模型时会先加载权重再把辅助头相关的 state_dict 删掉重新保存一个干净版本能省下不少存储空间。5. 实验效果哪些场景收益最明显5.1 小目标检测的改善Co-DETR 最明显的收益之一是小目标检测能力的提升。原因不难理解辅助头的密集监督相当于给编码器提供了大量小尺寸目标的特征梯度。标准 DETR 只有稀疏的查询级监督小目标在特征图上本来就弱很容易被漏掉。辅助头让编码器在训练中反复看到小目标的特征模式因此最后的检测效果会好不少。我在一个含大量小目标的工业数据集上测试过标准 Deformable DETR 的 mAP 大概是 23.5换到 Co-DETR 同设置下能到 27.2。尤其小目标 AP 从 11.8 涨到 16.4提升非常明显。当然了如果你用的是 Swin-L 这样的大模型收益可能会更显著因为强大的编码器配合充足的监督信号能把小目标特征榨得更干。5.2 收敛速度对比收敛速度是 Co-DETR 的另一个卖点。论文里说 12 epoch 就能超过很多需要 24 epoch 或 36 epoch 的模型。我实际复现时也有同样的感觉大概训练到 3k 迭代时检测效果已经接近标准 Deformable DETR 训练 10k 迭代的水平。这非常实用因为在算力有限的情况下快速迭代模型对项目的推进很重要。不过要说明训练省时间不代表推理省时间。推理还是和标准 DETR 一样单帧耗时取决于 backbone 和图像分辨率。辅助头只在训练时存在训练显存会高一些但推理侧的实惠一点没少。5.3 和纯 Transformer 检测器相比的优势和 DINO、DAB-DETR 这些模型比Co-DETR 算是一种训练策略上的创新而不是网络结构上的创新。它的优势在于兼容性好可以在不改变主结构的前提下把很多已有的 DETR 改进融入进来。比如你可以拿 DINO 的主结构套上 Co-DETR 的训练策略效果可能又上一层。这种即插即用的属性在实践中特别舒服。实际应用时我的建议是先跑通官方的 Co-DETR 配置把它作为 baseline。如果业务数据比较特殊再在这个基础上换更强的 backbone、加测试时增强等。简单粗暴的堆模型结构不一定好但把混合指派训练用起来效果通常是稳定的。6. 常见问题与踩坑记录6.1 辅助头不收敛loss 出现 NaN辅助头不收敛八成是初始化或者数据问题。我第一次训练时用了别人的预训练权重但数据集类别数不匹配导致辅助头的分类层梯度爆炸。检查之后发现加载权重时没有删掉分类层训练几步 loss 就 NaN。解决方案是加载前去掉不匹配的层或者用官方给的训练脚本中的 finetune 逻辑。还有一种情况是辅助头的正样本太少梯度信号弱导致 loss 长期不动。这时可以适当调大 anchor 匹配的阈值增加正样本数量。不过这个操作要看具体数据集正样本太多会加大训练开销不宜激进。6.2 显存不足批量大小调不动Co-DETR 训练显存开销比标准 DETR 大因为有多个辅助头同时跑。8 卡每卡 2 张图的配置ResNet-50 大概 40G 显存如果你用 24G 的单卡需要把 batch size 降到 1同时使用梯度累积模拟更大的 batch。学习率也要跟着调低。梯度累积的实现可以用 MMCV 的配置设置 optimizer 的 accumulated_iters。或者自己在训练循环里控制。实测下来梯度累积对最终效果影响不大但训练时间会拉长。如果算力紧张这个方案是可行的。6.3 推理结果不稳定的问题推理结果不稳定一般和两点有关。一个是可变形注意力的 CUDA 算子在不同 batch size 下有细微浮点差异另一个是某些后处理代码用了非确定性的操作比如遍历字典。如果要求结果严格可复现可以在推理脚本里加torch.backends.cudnn.deterministic True torch.backends.cudnn.benchmark False这样速度可能慢一点点但结果可以稳定重现。对于生产环境如果不涉及自动化回归其实不需要强制确定性。6.4 训练速度慢数据加载成瓶颈数据加载是训练速度的隐形杀手。Co-DETR 本身计算量不小如果数据加载跟不上GPU 利用率就会掉。建议开启 MMCV 的 DataLoader 多进程num_workers 至少设为 4如果机器内存够8 也可以。另外把图像裁到合适大小不要用原始高清图直接喂COCO 默认的 1333x800 就可以了。还有一个很多人忽略的点混合精度训练。Co-DETR 官方支持 fp16开启后显存能省不少训练速度也能提升。不过要留意 loss scale 的设置如果出现 NaN可以试试把 fp16 的 loss scale 调大或者关闭动态缩放。7. 我对 Co-DETR 的个人体会如果把 DETR 系列的发展比作一个修房子的过程那 Co-DETR 干的不是装修而是把地基打得更牢。它不改变房子外形但让房子更稳。对于做实际项目的我来说这种“只改训练不改推理”的特性太重要了。部署侧不用动任何代码模型结构完全不变但精度和收敛速度都上了一个档次。如果你现在要在新数据集上用 DETR 类方法我的建议是先别急着上复杂结构拿 Co-DETR 当 baseline 跑一遍。它在大多数场景下都优于原始 DETR 和 Deformable DETR而且代码成熟、踩坑的人多遇到问题容易找到资料。等你确认模型的核心瓶颈在哪里再针对性地做优化效率会高很多。最后分享一个训练技巧训练过程中不要只盯着 mAP多画一下辅助头的分类 loss 和回归 loss。辅助头的 loss 曲线能帮你判断编码器是否在正常学习。如果分类 loss 降得很慢说明特征判别力不足可以考虑加长训练轮次或换更强的 backbone。如果回归 loss 降得慢那大概率是锚点尺寸设置不合理需要统计数据集目标尺寸后重新配置。这个技巧帮我省了很多试错的成本。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →