首页
/ Transformers 中的 Conditional DETR:从条件交叉注意力原理到对象检测实战

Transformers 中的 Conditional DETR:从条件交叉注意力原理到对象检测实战

2026-09-07 12:24:58作者:乔或婵

导读

Conditional DETR(Conditional DETR for Fast Training Convergence)是 DETR(Detection Transformer)家族中首个针对"训练收敛慢"这一核心痛点提出的改进方案,其核心是一套**条件交叉注意力(conditional cross-attention)**机制,可将训练收敛速度相对原始 DETR 提升 6.7× 至 10×。本文以仓库 Conditional DETR 模型文档 为主体,结合 src/transformers/models/conditional_detr 下的源码与配置实现展开,带你理解其条件查询(conditional spatial query)的底层原理、完整数据流、ConditionalDetrConfig 的全部关键参数,以及如何用它完成对象检测、语义/实例/全景分割的推理与微调。

一、背景:DETR 的慢收敛问题与 Conditional DETR 的解法

DETR 将 Transformer 的编码器-解码器结构引入目标检测,无需锚框或 NMS 即可得到优秀效果,但其关键痛点是训练收敛非常慢。原论文指出,这一现象的根本原因在于:DETR 的交叉注意力高度依赖"内容嵌入(content embeddings)"来定位物体的四个极值点并回归边界框,这要求高质量的内容嵌入,从而显著提高了训练难度。

Conditional DETR(论文发表于 2021-08-13,由 DepuMeng 等作者于 2022-09-22 贡献到 Transformers)给出了对症下药的方案:从解码器嵌入中学习一个条件空间查询(conditional spatial query),用于解码器的多头交叉注意力。其好处在于——借助条件空间查询,每个交叉注意力头只需关注一个"带状区域"(band),该区域对应物体的某个极值点或物体框内部的一块区域。这大幅缩小了用于目标分类与框回归的空间范围,降低了对内容嵌入质量的依赖,从而显著降低训练难度。

依据论文与官方模型文档,Conditional DETR 在 R50/R101 骨干下收敛速度比 DETR 快 6.7×,在更强骨干 DC5-R50/DC5-R101 下可达 10×。当前仓库默认权重检查点为 microsoft/conditional-detr-resnet-50

二、条件交叉注意力的源码级拆解

在 Transformers 中,Conditional DETR 的实现文件为 modeling_conditional_detr.py。它与原始 DETR 的最大差异集中在解码器侧的交叉注意力模块 ConditionalDetrDecoderCrossAttention第 578 行起)。

2.1 内容与位置分离投影

常规注意力对 Query/Key 各做一次整体投影,而条件交叉注意力为 Query 与 Key 分别维护内容投影位置投影

# Content and position projections
self.q_content_proj = nn.Linear(hidden_size, hidden_size)
self.q_pos_proj = nn.Linear(hidden_size, hidden_size)
self.k_content_proj = nn.Linear(hidden_size, hidden_size)
self.k_pos_proj = nn.Linear(hidden_size, hidden_size)
self.v_proj = nn.Linear(hidden_size, hidden_size)
self.q_pos_sine_proj = nn.Linear(hidden_size, hidden_size)

2.2 拼接而非相加:维度翻倍的条件位置编码

与 DETR 将位置编码"加"到内容嵌入上不同,Conditional DETR 将条件空间嵌入与内容向量在特征维度上拼接(concatenate),使 Query/Key 维度临时翻倍:

  • Query 侧:query_input(内容投影结果)与经 q_pos_sine_proj 投影的 query_sine_embed(条件正弦位置嵌入)拼接;
  • Key 侧:key_input(内容)与 key_pos(编码器侧位置嵌入投影)拼接;
  • Value 侧仍是普通内容投影。

正因为拼接后 head_dim 翻倍,缩放因子在实现里按 expanded_head_dim = (hidden_size * 2) // num_attention_heads 计算,即 self.scaling = expanded_head_dim**-0.5,而最终输出投影 o_proj 仍把结果压回 hidden_size 维度。

需要特别注意的是 forwardquery_position_embeddings 的可选逻辑:

  • 第一层解码层(is_first=True:额外的位置嵌入会被到 Query/Key 内容上(query_input = query_input + self.q_pos_proj(query_position_embeddings)key_input = key_input + key_pos),然后内容再与正弦嵌入拼接;
  • 后续解码层:不再注入对象查询位置嵌入,仅使用被 query_scale 变换过的条件正弦嵌入。

该开关对应解码器在初始化时执行的特殊处理——第 1218-1220 行 会把第 2 层起的各层 encoder_attn.q_pos_proj 置为 None,仅保留第一层使用。

2.3 解码器中条件查询的生成

条件空间查询并非凭空而来,而是在 ConditionalDetrDecoder.forward 中按如下流程动态构造:

  1. 参考点预测:将对象查询位置嵌入(object_queries_position_embeddings,一个形状为 (batch, num_queries, hidden) 的可学习嵌入)送入 ref_point_head(一个 2 层 MLP 预测头),得到 reference_points_before_sigmoid
  2. Sigmoid 归一化reference_points = sigmoid(...) 得到位于 (0,1) 区间的参考点,其前两维作为对象中心 obj_center
  3. 条件正弦嵌入:用 encode_sinusoidal_position_embedding(obj_center, num_pos_feats=d_model // 2) 生成初始条件正弦嵌入;
  4. 逐层变换:第 0 层保持 pos_transformation = 1,后续每一层先用 query_scale(一个 2 层 MLP 预测头)从当前 hidden_states 预测变换量,再执行 query_sine_embed = query_sine_embed_before_transformation * pos_transformation,从而为每一层解码器生成"以参考点为中心"的条件空间查询。

配合每个解码器层内的 LayerNorm 堆叠(auxiliary_loss=True 时逐层归一化输出用于辅助损失),这套机制在工程实现上完整复刻了论文中"通过条件空间查询收窄每个注意力头的关注条带"的设计。

三、端到端数据流:从像素到检测结果

以不带任务头的裸模型 ConditionalDetrModel 为例,其 forward 的完整链路(第 1368 行起)分为五个阶段:

  1. 骨干提取特征pixel_values + pixel_mask 送入 ConditionalDetrConvEncoder(默认 ResNet-50 风格),取出最后一张特征图与其对应的降采样 mask;
  2. 1×1 通道压缩input_projection = nn.Conv2d(backbone最后通道数, d_model, kernel_size=1) 将特征图通道压缩到 d_model(默认 256);
  3. 空间位置编码:按 position_embedding_type 选择 ConditionalDetrSinePositionEmbedding(正弦,normalize=True)或 ConditionalDetrLearnedPositionEmbedding(可学习),生成 2D 空间位置嵌入;
  4. Flatten 送编码器:特征图 (N,C,H,W) 展平并转置为 (batch, H*W, hidden),连同展平后的 mask 与空间位置嵌入一起送入 6 层 Transformer 编码器;
  5. 零初始化查询送解码器:对象查询的位置部分来自 self.query_position_embeddings = nn.Embedding(num_queries, d_model) 的权重(广播到 batch 维度),而内容部分初始化为全零张量(queries = torch.zeros_like(...)),二者一并送入解码器进行自注意力与条件交叉注意力更新。

模型文档中的示例代码印证了最后一个阶段输出的语义——解码器最后一层的隐藏状态即为每个查询的最终嵌入,形状为 (batch_size, num_queries, hidden_size)

from transformers import AutoImageProcessor, AutoModel
from PIL import Image
import httpx
from io import BytesIO

url = "http://images.cocodataset.org/val2017/000000039769.jpg"
with httpx.stream("GET", url) as response:
    image = Image.open(BytesIO(response.read()))

image_processor = AutoImageProcessor.from_pretrained("microsoft/conditional-detr-resnet-50")
model = AutoModel.from_pretrained("microsoft/conditional-detr-resnet-50")

inputs = image_processor(images=image, return_tensors="pt")
outputs = model(**inputs)

# 最终查询嵌入:shape (batch_size, num_queries, hidden_size)
last_hidden_states = outputs.last_hidden_state
list(last_hidden_states.shape)  # [1, 300, 256]

裸模型的输出 ConditionalDetrModelOutput 在标准 Seq2SeqModelOutput 基础上新增了两个字段:intermediate_hidden_states(仅在 config.auxiliary_loss=True 时返回的逐解码层归一化中间激活)与 reference_points(每层解码器的参考点,形状 (decoder_layers, batch, num_queries, 2))。

四、配置详解:ConditionalDetrConfig

configuration_conditional_detr.py 定义了 ConditionalDetrConfigmodel_type = "conditional_detr")。除模型文档 autodoc 中显式列出的三个参数外,源码给出了完整字段及其默认值,下面按用途分组:

4.1 检测语义相关

参数 默认值 说明
num_queries 300 对象查询(检测槽位)数量,即单图可检测的最大物体数;COCO 推荐 100
auxiliary_loss False 是否使用逐解码器层的辅助解码损失
position_embedding_type "sine" 图像特征之上的位置编码类型,可选 "sine""learned"
dilation False 是否在最后一个卷积块用空洞卷积替代 stride(即 DC5 变体);仅在使用 timm 骨干时受支持

4.2 网络结构相关

参数 默认值 说明
d_model 256 编码器/解码器隐藏维度(也是 box 头输入维度)
encoder_layers / decoder_layers 6 / 6 编码器 / 解码器层数
encoder_ffn_dim / decoder_ffn_dim 2048 / 2048 FFN 中间层维度
encoder_attention_heads / decoder_attention_heads 8 / 8 注意力头数
activation_function "relu" MLP 激活函数
dropout / attention_dropout / activation_dropout 0.1 / 0.0 / 0.0 各类 dropout 概率
encoder_layerdrop / decoder_layerdrop 0.0 / 0.0 逐层随机丢层概率(LayerDrop)
num_channels 3 输入图像通道数
is_encoder_decoder True 编码器-解码器结构标志

4.3 初始化与损失权重

参数 默认值 说明
init_std 0.02 正态初始化标准差
init_xavier_std 1.0 Xavier 初始化标准差
class_cost / bbox_cost / giou_cost 2 / 5 / 2 训练集匹配阶段(二分匹配)中分类、L1 框损失、GIoU 损失的代价系数
cls_loss_coefficient 2 分类损失系数
bbox_loss_coefficient 5 L1 框回归损失系数
giou_loss_coefficient 2 GIoU 损失系数
mask_loss_coefficient / dice_loss_coefficient 1 / 1 分割任务中 mask 与 DICE 损失系数
focal_alpha 0.25 Focal Loss 的 alpha 参数(分类损失)

4.4 兼容性与骨干配置

为与 Transformer 通用属性命名习惯对齐,attribute_maphidden_size 映射到 d_modelnum_attention_heads 映射到 encoder_attention_headsnum_hidden_layers 映射到 encoder_layers,方便 AutoConfig 与通用工具代码读取。__post_init__ 中通过 consolidate_backbone_kwargs_to_config 处理骨干配置:默认骨干为 resnet50、类型为 resnetout_features=["stage4"];若启用 dilation,则额外要求 output_stride=16(即 DC5)。同时会把 timm 骨干的默认参数(features_only=Trueuse_pretrained_backbone=Falseout_indices=[1,2,3,4])固化,保证与原始实现行为一致。

初始化一个 microsoft/conditional-detr-resnet-50 风格配置并实例化模型的方式如下:

from transformers import ConditionalDetrConfig, ConditionalDetrModel

configuration = ConditionalDetrConfig()
model = ConditionalDetrModel(configuration)  # 随机权重
configuration = model.config

五、对象检测实战:ConditionalDetrForObjectDetection

5.1 检测头与输出格式

ConditionalDetrForObjectDetection 在裸模型之上堆叠两个检测头(第 1510-1522 行):

self.model = ConditionalDetrModel(config)
self.class_labels_classifier = nn.Linear(config.d_model, config.num_labels)
self.bbox_predictor = ConditionalDetrMLPPredictionHead(
    input_dim=config.d_model, hidden_dim=config.d_model, output_dim=4, num_layers=3
)

其中 class_labels_classifier 输出形状为 (batch, num_queries, num_classes + 1) 的分类 logits(最后一维含"无物体"类别),bbox_predictor 是一个 3 层 MLP,输出 4 维归一化框坐标。

其输出类型 ConditionalDetrObjectDetectionOutput 的字段(详见 第 92-126 行):

  • loss:分类负对数似然 + 边界框损失(L1 与尺度不变 GIoU 的线性组合)的总损失,仅在传入 labels 时返回;
  • loss_dict:各分项损失字典,便于日志记录;
  • logits:形状 (batch_size, num_queries, num_classes + 1) 的分类 logits;
  • pred_boxes:形状 (batch_size, num_queries, 4)归一化框坐标,编码为 (center_x, center_y, width, height),数值在 [0, 1] 区间、相对各自图像尺寸(不考虑 padding);
  • auxiliary_outputs:当 config.auxiliary_loss=True 且提供 labels 时,给出每个解码器层的 logitspred_boxes 列表;
  • last_hidden_state 及编码器/解码器的隐藏状态与注意力堆栈。

5.2 从原始输出到可视化坐标

因为 pred_boxes 是相对图像归一化的 (cx, cy, w, h),要得到可直接绘制的像素框,必须使用 ConditionalDetrImageProcessor.post_process_object_detectionimage_processing_conditional_detr.py 第 806 行)完成反归一化并转换为 (x1, y1, x2, y2) 格式。其签名与语义为:

  • outputs:模型的 ConditionalDetrObjectDetectionOutput(仅支持 PyTorch);
  • threshold:保留预测的分数阈值(方法默认 0.5);
  • target_sizes:形状 (batch_size, 2) 的 Tensor 或 (height, width) 元组列表,表示每个样本的原始尺寸;传 None 则不做尺寸还原;
  • top_k:每图最多保留的框数量(默认 100)。

典型推理代码如下:

from transformers import AutoImageProcessor, ConditionalDetrForObjectDetection
import torch

image_processor = AutoImageProcessor.from_pretrained("microsoft/conditional-detr-resnet-50")
model = ConditionalDetrForObjectDetection.from_pretrained("microsoft/conditional-detr-resnet-50")

inputs = image_processor(images=image, return_tensors="pt")
outputs = model(**inputs)

# 还原为 (x1, y1, x2, y2) 的像素坐标框
target_sizes = torch.tensor([image.size[::-1]])
results = image_processor.post_process_object_detection(
    outputs, threshold=0.7, target_sizes=target_sizes
)[0]
for score, label, box in zip(results["scores"], results["labels"], results["boxes"]):
    box = [round(i, 2) for i in box.tolist()]
    print(f"Detected {model.config.id2label[label.item()]} with confidence "
          f"{round(score.item(), 3)} at location {box}")

六、扩展到分割任务

ConditionalDetrForSegmentation 面向语义 / 实例 / 全景分割。从其输出类型 ConditionalDetrSegmentationOutput 看,除了与检测一致的总损失、loss_dictlogitspred_boxesauxiliary_outputs 外,还包含 pred_masks(形状 (batch_size, num_queries, height/4, width/4),即 1/4 分辨率的 mask logits)。从源码结构可以推断,其分割头沿用了 DETR 家族经典的 mask 头组合:文件内包含 ConditionalDetrMHAttentionMap(按查询生成注意力图)、ConditionalDetrConvBlockConditionalDetrFPNFusionStage(多尺度 FPN 融合)与 ConditionalDetrMaskHeadSmallConv(小型反卷积上采样头)等子模块(见 第 889-1047 行)。

ConditionalDetrImageProcessor 为分割任务提供了三种后处理入口:

  • post_process_semantic_segmentation第 865 行):把各查询的 mask 折叠成语义分割图;
  • post_process_instance_segmentation第 936 行):按检测结果切分为实例 mask;
  • post_process_panoptic_segmentation第 1019 行):生成包含 segmentationsegments_info 的全景分割结果。

对应后处理参数可参考同一文件内的 docstring 与配套测试 test_image_processing_conditional_detr.py

七、图像处理器:预处理与批处理细节

ConditionalDetrImageProcessor 基于 torchvision 后端的预处理管线(preprocess第 672 行起)按以下阶段处理单张或批量图像:

  1. 注解准备:当传入 COCO 格式注解时,通过 prepare_annotation 将 COCO 目标转换为 Conditional DETR 期望的单图 target;可配合 return_segmentation_masksmasks_path(指向 mask PNG 目录)加载分割掩码;
  2. Resize:按 sizeSizeDict)与 resample 缩放图像,并同步缩放注解框;
  3. Rescale + Normalize:按 rescale_factorimage_meanimage_std 完成像素缩放与标准化(实现上做了融合以提效);
  4. 注解归一化do_convert_annotationsTrue 时将框转为 (cx, cy, w, h) 的归一化形式(normalize_annotation);
  5. Pad:若 do_pad=True,将各图像补齐到统一尺寸(未显式给出 pad_size 时取 batch 内最大高宽),并同步生成 pixel_mask1 表示真实像素、0 表示 padding)。

最终 BatchFeature 包含 pixel_values 与(padding 时)pixel_mask;若提供了注解,则还会在 labels 键下返回逐样本的目标字典列表。当图像尺寸不一时,pixel_mask 会一路传递到编码器,用于屏蔽 padding 位置的注意力。

八、微调与资源索引

官方为该模型维护了完整的微调工具链:

需要说明的是,模型实现文件由 modular_conditional_detr.py 通过 modular 机制自动生成(文件头有明确的 CI 约束注释),如需改动实现逻辑应修改 modular 源文件而非直接编辑生成文件。

九、小结:一张关键参数的速查表

关注点 默认配置 备注
检测槽位 num_queries=300 COCO 建议 100
结构 6 层编码器 + 6 层解码器,d_model=256,8 头 FFN 维度 2048
位置编码 position_embedding_type="sine" 可选 "learned"
条件交叉注意力 内容/位置分离投影,维度翻倍拼接 仅首层叠加对象查询位置嵌入
骨干 默认 ResNet-50(stage4 可用 dilation 开启 DC5
损失组合 L1 + GIoU + 分类(可加辅助损失) 系数见 *_loss_coefficient
输出坐标 归一化 (cx, cy, w, h) post_process_object_detection 还原像素框

Conditional DETR 是理解"为什么 DETR 收敛慢、如何用条件空间查询缓解"的最佳范本之一,而本仓库中从配置到模型、图像处理器、转换脚本与测试的完整闭环,也为后续在此基础上改进(如 DAB-DETR、DN-DETR、DINO 中的去噪与动态锚点)提供了清晰的代码级参照。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.75 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
857
1.35 K
docsdocs
暂无描述
Markdown
898
5.82 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
921
1.84 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.8 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
531
596
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.02 K
519
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.36 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
391