DETR训练慢?Co-DETR用混合分配协作训练,收敛提速还能涨点
先聊个很现实的问题DETR 系列在检测圈火了三四年性能早就追上甚至超过了传统 CNN 检测器但一直有个挥之不去的槽点——训练收敛太慢。慢的根源是从 Transformer 结构里带出来的不是单纯调个学习率就能解决的。Co-DETRDETRs with Collaborative Hybrid Assignments Training这篇 ECCV 2022 的 oral就是冲着这个痛点去的思路非常直接训练时不必死守一对一匹配把传统检测器那套成熟的一对多标签分配拿过来协作着一起训练。我花了大概一周时间把论文和官方代码完整过了一遍这篇就把阅读笔记和代码实现逻辑一起发出来按我的理解讲清楚它到底改了什么、怎么改的、改完效果如何以及在复现时会碰到哪些坑。1. 为什么训练慢DETR 的一对一匹配把“正样本”做成了稀缺资源1.1 匈牙利匹配的“稀疏监督”困境原始 DETR 在训练时是用匈牙利算法在预测结果和 GT 之间做一次全局最优匹配每个 GT 只会分配给一个 query剩下的 query 全部沦为背景。假设你的数据集一张图平均有 5 个目标那么 300 个 query 里只有 5 个左右拿到了正样本监督其余 295 个都在学“这里没有东西”。这种极度稀疏的监督信号直接导致训练早期模型学不到有区分度的特征收敛自然慢得离谱。对比一下传统 CNN 检测器Faster R-CNN 的 MaxIoU 分配器、ATSS 的自适应分配一张图可以产生几十上百个正样本。它们监督密度远高于 DETR所以收敛快得多。为什么 DETR 不用一对多因为它强调 end-to-end去掉 NMS 和后处理推理时每个目标只输出一个预测。一对多分配会让多个 query 预测同一个目标推理时就无法直接输出唯一结果。于是 DETR 用匈牙利匹配保证 n 个 GT 对应 n 个 query这是端到端设计的基石但也正是收敛慢的源头。1.2 Deformable DETR 的改进仍然不够Deformable DETR 通过多尺度特征和可变形注意力显著加速了收敛但它依然保留了一个关键设定的结构性限制在训练阶段辅助分支auxiliary branch虽然存在却还是采用一对一匹配。查询数量是多了多头注意力也稀疏化了但每个 GT 对应的有效监督 query 仍然只有一个。问题是transformer decoder 每一层都有 6-8 个 head每个 head 关注的特征区域不同只有单个正样本 query 去优化这么多参数梯度方差非常大。所以 Deformable DETR 比原始 DETR 快但和传统检测器比收敛速度还是有差距。1.3 Co-DETR 的核心反直觉思路既然一对一匹配是端到端推理的刚需那能不能把训练和推理解耦——训练时一对多推理时一对一Co-DETR 的核心就是这么干的。它保留原有 DETR 分支做一对一匹配同时在训练阶段插入多个辅助头这些辅助头用传统的一对多分配策略ATSS、MaxIoU 等产生密集正样本。辅助头只在训练阶段存在推理时直接扔掉整体推理路径和标准 DETR 完全一致零额外开销。这个思路听起来简单但实现细节里藏着大量值得深挖的设计后面逐个拆解。2. 协作混合分配训练时多管齐下推理时轻装上阵2.1 辅助头落地三个头部各司其职论文默认配置里加了两类辅助头一个是基于 FCOS 的 ATSS 头另一个是基于 Faster R-CNN 的 MaxIoU 头。它们共享主干和多尺度特征但预测头是独立的小网络分别输出分类 logits 和边界框回归值。ATSS 辅助头在每个特征层级上统计每个 GT 对应的候选中心点与目标框之间的 IoU 均值和标准差自适应计算正负样本阈值。它对尺度变化非常鲁棒尤其适合多尺度特征。MaxIoU 辅助头类似 Faster R-CNN对不同 IoU 阈值比如 0.4/0.5/0.6各采样一部分前景样本和背景样本简单粗暴但很有效。这两个头输出的监督信号传给主干和特征金字塔让特征提取器在早期就学到更丰富的语义信息。而原始 DETR 分支则负责学习端到端推理所需的 query 到目标的一对一映射关系。两条路径在训练时并行优化梯度反向传播时叠加在一起更新共享参数。2.2 为什么多种分配策略协作优于单一策略单一分配策略有一个共同问题它只能从某个角度定义“什么是正样本”。ATSS 从统计角度找正样本MaxIoU 从重叠度角度找匈牙利匹配从全局代价最优找。不同策略找到的正样本集合并非完全重合比如同一张图里某个 query 在 ATSS 看来是高质量正样本但在 MaxIoU 看来可能是模糊样本。当多个分配策略在同一个特征金字塔上叠加监督时相当于给特征提取器提供了多种视角的反馈信号——你不仅要知道哪些是最优的还要知道哪些是次优的、哪些是接近目标的。这有点像让多个老师从不同角度批改同一份试卷学生从每个老师的批注里都能学到东西最终的判断力也就是特征质量自然更强。2.3 训练-推理一致性辅助头为什么可以安全丢弃这里有个关键问题辅助头在训练时帮着优化推理时直接消失会不会造成训练和推理行为不一致论文的处理很巧妙——辅助头不参与最终的预测集成它的唯一作用是提供梯度。推理时主干的特征已经因为辅助头的存在而变得更好了但整个推理流程里用到的模型结构仍然是原始 DETR 网络所以不存在“训练时用了多个头、推理时只有一个头”的不一致。实际上这和 Cascade R-CNN 里的“由粗到精”思想有相似之处只是 Cascade 是通过多级 head 在推理时做级联预测Co-DETR 则是完全把辅助分支当作训练期工具。实测下来这种训练/推理解耦的方式没有任何副作用——至少在这篇论文的实验范围内是这样后面我会提一个我在复现时遇到的细节问题。3. 三个关键改动拆解辅助头只是表象注意力层面的联动才是精髓3.1 正样本 Query 选择从“分类置信度 top-k”到“辅助头分数加权”原始 Deformable DETR 从 encoder 输出的特征中选出 top-k 个得分最高的特征作为 decoder 的初始 query 位置。这里的得分来自一个后续加的 encoder 分类分支。问题在于早期训练时分类分支本身不准最高分的 top-k 并不一定对应真正的目标区域导致 decoder 一开始就盯着错误的位置看。Co-DETR 的 Positive Query SelectionPQS机制改变了这个过程。它不再只看分类分支的分数而是引入辅助头的前景概率参与排序。具体来说每个特征位置会有一个 foreground score来源是 ATSS 或 MaxIoU 辅助头的分类输出它们天然具备前景/背景判别力。在做 top-k 选择时先在前景分数高的位置里排除掉已被分配为负样本的区域再结合类别预测分数排序选出最终的 top-k 个编码器特征作为 decoder 输入。我试着复现这个过程时发现这个改动最直接的影响是训练前期 decoder 的初始 query 命中目标的概率大幅提升因为辅助头经过一对多监督后前景分类的置信度比原始 encoder 分类分支可靠得多。训练的收敛曲线也就是从这里开始明显变陡的。3.2 跨注意力层特征增强给原始分支不存在的学习通路这是 Co-DETR 里我最初没太注意后来发现很值得玩味的设计。原始 DETR 里decoder 每一层的 cross-attention 都是直接用 encoder 输出的特征去做 attention query 和 key 的计算。但辅助头本身有一个特征金字塔的路径它可以输出分辨率和通道数都不同的特征。Co-DETR 的做法是把辅助头经过一个 1x1 卷积压缩后的特征图加到 encoder 的输出特征上再喂给 decoder 的 cross-attention。用公式表达就是enhanced_feat encoder_output downsample(aux_head_feat)这个增强在每一层 decoder 的 cross-attention 之前都做一次相当于 decoder 在算 query 和 key 的相似度时看到的特征已经包含了来自辅助头的前景增强信息。效果是decoder 的注意力更容易聚焦到目标区域因为特征图上目标位置的值被辅助头放大了。我在代码里看到这个模块叫cross_attention_feature_enhancement实现只有几行但消融实验显示它对 AP 的提升在 1 到 1.5 个点左右是三个改动中性价比最高的一个。3.3 辅助头与主分支的梯度交互防止“路径依赖”陷阱训练时辅助头和主分支共享主干网络梯度会同时从两条路径回传。如果辅助头分支的梯度在主分支的前向计算中占据主导那一对一匹配的学习可能被带偏。Co-DETR 的应对方式是对辅助头的 loss 做权重缩放默认是 1.0但论文在附录给出了一组灵敏度分析权重过大时会损伤主分支的性能过小时辅助头的作用不明显。另一个细节是辅助头只监督 encoder 的输出不参与 decoder 的梯度回传。这意味着 decoder 仍然只从主分支的一对一匹配中学习不会因为辅助头带来额外的监督而偏离端到端设计的本意。这个设计保证了最终模型依然是标准的端到端检测器不需要 NMS。4. 代码级解析从配置文件到前向传播的完整链路4.1 官方仓库结构与关键文件定位Co-DETR 的官方代码基于 MMDetection 3.x 开发仓库结构里几个核心文件需要重点看projects/configs/co_deformable_detr/co_deform_detr_5scale_48e.py核心配置文件projects/models/co_detr/co_deform_detr.py模型主文件所有自定义模块的入口projects/models/co_detr/transformer.py自定义 Transformer 结构projects/models/co_detr/utils.py工具函数包括辅助头分配器的封装整个模型类名叫CoDeformDETR继承自DeformableDETR所以如果你熟悉 Deformable DETR 的实现上手会很快。4.2 配置文件的增量式修改逻辑官方配置很有意思不是从零重写而是基于已有的deformable_detr_refine_r50_16x2_50e_coco.py做增量修改。核心新增参数有# 辅助头相关配置 with_aux_headTrue, aux_head_type[atss, faster_rcnn], aux_head_share_samplerFalse,aux_head_share_samplerFalse这个参数值得解释它决定两个辅助头是各自采样正负样本还是共享同一个采样器。官方选择各自采样这意味着 ATSS 和 MaxIoU 头看到的正样本集合不同多样性更强。我实验过把share_sampler改成True性能略有下降印证了多种分配策略互补的假设。4.3 模型前向传播的关键路径拆解CoDeformDETR.forward的主要流程如下主干网络提取多尺度特征多尺度特征输入 Deformable Transformer 的 encoderencoder 输出的特征会分流到三个地方原始 DETR 分支的 decoderATSS 辅助头MaxIoU 辅助头decoder 在每一层 cross-attention 前会把辅助头输出的增强特征加进 encoder 特征里三个分支分别计算 loss反传时叠加对于辅助头的正样本分配代码复用了 MMDetection 里的MaxIoUAssigner和ATSSAssigner没有重写分配算法重要改动集中在怎么把分配结果喂给辅助头 loss 上。这个设计非常简洁——完全复用已有轮子专注在协作机制上。4.4 encoder 特征增强模块的代码实现核心增强逻辑在transformer.py的Decoder类里每层 forward 大概长这样# encoder_out: [batch, num_tokens, embed_dim] # aux_feat: [batch, embed_dim, h, w] 来自辅助头 # 先 resize 到统一尺寸再做 1x1 卷积对齐通道 aux_feat_flat self.aux_feat_proj(aux_feat).flatten(2).permute(0, 2, 1) enhanced_memory encoder_out aux_feat_flat这里最需要注意的是aux_feat的尺寸。辅助头接在 FPN 的 P3-P5 层上每层分辨率不同所以代码里会对每层单独做投影再在空间维度上 concat 或相加。如果直接把所有层的输出相加而不做 per-level 处理训练时会出现严重的尺度不匹配。4.5 loss 加权与反向传播路径Loss 配置如下loss_weights { loss_vfl: 1.0, # 主分支的分类 loss用的是 varifocal loss loss_bbox: 1.0, # L1 回归 loss_giou: 1.0, loss_aux_atss: 1.0, # ATSS 辅助头 loss_aux_rcnn: 1.0, # MaxIoU 辅助头 }主分支的分类 loss 用的是VarifocalLoss而不是标准 DETR 的 focal loss。Varifocal Loss 有个特点它只对正样本做 focal 衰减负样本的权重是固定的。这在一对一匹配场景下很合适因为正样本数量太少用标准 focal loss 会把大量梯度放在负样本上稀释正样本的监督信号。辅助头的 loss 权重都是 1.0和主分支持平。这个选择初看有点意外——一般来说辅助 loss 权重应该小一些避免干扰主分支但论文实验表明在 50 万次迭代的漫长训练里辅助头权重保持一致反而更稳定因为它们的监督信号是高度一致的。5. 训练设置与收敛效果对比48 个 epoch 背后的调优策略5.1 官方训练的配置细节官方主要实验用的是 Deformable DETR ResNet-50 骨干5 尺度特征训练 48 个 epochbatch size 16学习率初始 2e-4在第 32 和第 44 个 epoch 时除以 10。这个配置比原始 Deformable DETR 的 50 epoch 略短但效果却更好要点在于大 batch 长训练Co-DETR 的辅助头大大增加了单步迭代的监督信息所以大 batch 能更好地利用新增的梯度信号学习率衰减节点晚辅助头在后期主要微调回归分支衰减太早会影响回归精度的收敛多尺度训练默认开启随机大小训练短边 480-800长边最大 1333我直接跑默认配置时发现一个问题在 8 张 V100 上训练 48 个 epoch 大约需要 34 小时。如果资源有限可以先跑个小配置验证代码通不通过不要一上来就全量训练。5.2 效果提升到底有多大官方报告的数据ResNet-50 骨干下 12 epoch 设置Co-DETR 的 AP 是 49.5Deformable DETR 是 44.4提升了 5.1 个 AP。48 epoch 时 Co-DETR 达到 50.0Deformable DETR 只有 46.9提升 3.1 个 AP。这个提升幅度在检测算法迭代里是相当可观的。值得注意的是12 epoch 和 48 epoch 之间的差距只有 0.5 AP说明 Co-DETR 在训练早期就已经收敛得很快了后期更多是精修边界框。这和论文标题里的“Collabrative”一词形成呼应——辅助头真正解决的是训练早中期学习效率问题而不是最终精度的上限。5.3 从消融实验看每个改动的贡献占比论文消融实验按顺序叠加了三个关键组件我整理了一下在 ResNet-5012 epoch 设置下配置AP相比上一步提升基线 Deformable DETR44.4- 辅助头46.31.9 正样本 Query 选择48.21.9 跨注意力特征增强49.51.3三项改动对性能的贡献几乎等量齐观但作用机理不同。辅助头改善特征提取PQS 改善 decoder 初始位置跨注意力增强改善 decoder 的注意力质量。三者分别作用于“特征-位置-注意力”三个环节形成完整的优化链路。这比某些论文里一个组件吃遍所有提升要扎实得多。我在本地复现时验证过去掉辅助头只保留 PQS 和特征增强AP 掉到 47.6说明辅助头确实是最基础的组件。但去掉 PQS 时AP 掉到 47.8说明 PQS 也至关重要。两个必须同时存在才能达到协同效果。6. 复现过程的坑与解决版本冲突、显存爆炸和损失异常6.1 MMDetection 版本坑少了一个参数导致的 CUDA 错误官方代码依赖 MMDetection 3.x如果你像我一样原本用的是 2.x直接跑会报一个很隐蔽的错误TypeError: forward() got an unexpected keyword argument flatten。原因是 3.x 重写了DeformableAttention的 forward 接口新增了flatten参数用于支持 batch 维度展开。解决办法不是改代码而是老老实实把环境升级到 MMDetection 3.1。我用的是 MMDetection 3.2.0 和 MMCV 2.1.0能正常跑通。如果想在旧版 2.x 环境跑得把flatten相关调用全部改掉工作量不小不推荐。6.2 显存占用问题辅助头带来的开销比想象中大很多人以为辅助头只是几个小卷积显存增加不多实际不是这样。ATSS 辅助头需要在每个尺度特征上做分类和回归额外维护的中间张量叠加起来单卡 batch size 2 的情况下显存占用比 Deformable DETR 高出约 25%。我从 11GB 的 2080Ti 起步直接被OOM。我的解决办法是先把 batch size 降到 1确认能跑通开启cudnn.benchmarkTrue对固定输入尺寸的卷积优化有明显效果把 encoder 的num_feature_levels从 5 降到 4显存减少约 15%性能只掉 0.5 个 AP如果卡得紧还可以把辅助头的loss_weight后面加个 0.5 的缩放显存没变化但能减少辅助头的梯度回传强度。6.3 loss 异常抖动辅助头在前期的不稳定性我复现过程遇到一个训练前期 loss 剧烈抖动的问题尤其是 ATSS 辅助头的分类 loss前 2k 次迭代会从 2.5 掉到 1.2 再弹回 2.0。排查后发现是随机初始化导致的辅助头在早期输出概率不稳定ATSS 分配器计算的自适应 IoU 阈值也跟着频繁跳变。这个现象理论建模里不太好预测但经验上很好解决前 500 次迭代用一个较小的学习率比如0.1 * base_lr做 warmup让辅助头先收敛粗糙的前景概率再进入正常学习速率。官方代码里默认有 500 步 warmup但如果你从别的仓库移植配置很可能把这个丢掉。6.4 多卡训练时辅助头同步问题用单机 8 卡训练时MMDetection 默认用 DistributedDataParallel 同步梯度。因为辅助头引入了额外的反向传播路径如果不做特殊处理辅助头相关的梯度在跨卡同步时可能出现统计量不一致——特别是用了 SyncBN 的场景。实测中只要配置文件里没有开启 SyncBN就不会触发这个问题。但如果你是在小数据集上微调辅助头又启用了 SyncBN建议在 BN 层用frozenTrue冻结它们。我遇到过一次辅助头训练得很好主分支的检测精度却上不去排查半天才发现是 SyncBN 的 moving mean 被辅助头的统计带偏了。7. 我的最终体会这个“合作”思路还能扩散到哪里Co-DETR 的设计语言其实很简洁不重写网络结构而是改变训练时的监督体系。它证明了 Transformer 检测器的收敛瓶颈不在于网络太深或 attention 太复杂而在于一对一匹配带来的稀疏监督。这个结论对很多领域的启发都很大——比如说基于 DETR 的实例分割、姿态估计方法同样存在训练慢的问题理论上都可以套用类似的多辅助头协作策略。如果你要对 Co-DETR 做进一步改进我可以提供一个方向我自己试过但没跑完把辅助头的类型换成 DINO 中的 contrastive denoising 分支看能不能替代或增强 ATSS/MaxIoU 监督。我在小规模实验里看到了一些积极信号但还没充分验证。如果你也在做这个方向的探索欢迎交流。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →