DETR目标检测实战:手把手教你用Transformer实现端到端检测(附代码)
DETR目标检测实战:从零构建基于Transformer的端到端检测系统
在计算机视觉领域,目标检测一直是核心技术难题之一。传统方法如Faster R-CNN、YOLO等虽然效果显著,但都依赖于复杂的后处理流程(如非极大值抑制NMS)和人工设计的锚框(anchor)。2020年,Facebook AI团队提出的DETR(DEtection TRansformer)彻底改变了这一局面,首次将Transformer架构成功应用于目标检测任务,实现了真正的端到端检测。本文将带您从零开始,完整实现一个基于PyTorch的DETR模型,涵盖数据准备、模型构建、训练技巧和推理部署全流程。
1. 环境准备与数据加载
1.1 安装依赖库
首先确保您的Python环境为3.7+,并安装以下核心依赖:
pip install torch==1.10.0+cu113 torchvision==0.11.1+cu113 -f https://download.pytorch.org/whl/torch_stable.html
pip install pycocotools matplotlib scikit-image
1.2 COCO数据集准备
DETR默认使用COCO数据集进行训练和评估。下载并解压数据集到data/coco目录:
data/coco/
├── annotations/ # 存放instances_train2017.json等标注文件
├── train2017/ # 训练集图片
└── val2017/ # 验证集图片
创建自定义数据集加载器,继承PyTorch的Dataset类:
from torch.utils.data import Dataset
from pycocotools.coco import COCO
class CocoDetection(Dataset):
def __init__(self, img_folder, ann_file, transforms):
self.img_folder = img_folder
self.coco = COCO(ann_file)
self.ids = list(sorted(self.coco.imgs.keys()))
self.transforms = transforms
def __getitem__(self, idx):
coco = self.coco
img_id = self.ids[idx]
ann_ids = coco.getAnnIds(imgIds=img_id)
target = coco.loadAnns(ann_ids)
path = coco.loadImgs(img_id)[0]['file_name']
img = Image.open(os.path.join(self.img_folder, path)).convert('RGB')
if self.transforms is not None:
img, target = self.transforms(img, target)
return img, target
1.3 数据增强策略
DETR对数据增强相对敏感,建议采用以下组合:
from torchvision.transforms import Compose, RandomResizedCrop, ToTensor, Normalize
def make_transforms(image_set):
normalize = Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
if image_set == 'train':
return Compose([
RandomHorizontalFlip(),
RandomResizedCrop(800, scale=(0.8, 1.0)),
ToTensor(),
normalize
])
else:
return Compose([
Resize(800),
ToTensor(),
normalize
])
2. DETR模型架构实现
2.1 主干网络(Backbone)
DETR使用标准的CNN(如ResNet50)提取图像特征:
import torch
from torch import nn
from torchvision.models import resnet50
class Backbone(nn.Module):
def __init__(self, hidden_dim=256):
super().__init__()
self.backbone = resnet50(pretrained=True)
self.conv = nn.Conv2d(2048, hidden_dim, 1)
def forward(self, x):
x = self.backbone.conv1(x)
x = self.backbone.bn1(x)
x = self.backbone.relu(x)
x = self.backbone.maxpool(x)
x = self.backbone.layer1(x)
x = self.backbone.layer2(x)
x = self.backbone.layer3(x)
x = self.backbone.layer4(x) # [batch, 2048, h/32, w/32]
x = self.conv(x) # [batch, hidden_dim, h/32, w/32]
return x
2.2 Transformer编码器-解码器
实现Transformer的核心组件:
class Transformer(nn.Module):
def __init__(self, d_model=512, nhead=8, num_encoder_layers=6,
num_decoder_layers=6, dim_feedforward=2048, dropout=0.1):
super().__init__()
encoder_layer = nn.TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout)
self.encoder = nn.TransformerEncoder(encoder_layer, num_encoder_layers)
decoder_layer = nn.TransformerDecoderLayer(d_model, nhead, dim_feedforward, dropout)
self.decoder = nn.TransformerDecoder(decoder_layer, num_decoder_layers)
def forward(self, src, mask, query_embed, pos_embed):
# src: [h*w, batch, d_model]
# mask: [batch, h, w]
# query_embed: [num_queries, d_model]
# pos_embed: [h*w, batch, d_model]
memory = self.encoder(src, src_key_padding_mask=mask, pos=pos_embed)
hs = self.decoder(query_embed.unsqueeze(1).repeat(1, memory.shape[1], 1),
memory, memory_key_padding_mask=mask,
pos=pos_embed, query_pos=query_embed)
return hs.transpose(1, 2) # [num_queries, batch, d_model]
2.3 位置编码与预测头
DETR采用可学习的位置编码和简单的MLP预测头:
class PositionEmbeddingSine(nn.Module):
def __init__(self, num_pos_feats=64, temperature=10000):
super().__init__()
self.num_pos_feats = num_pos_feats
self.temperature = temperature
def forward(self, x):
not_mask = torch.ones_like(x[:,0,:,:])
y_embed = not_mask.cumsum(1, dtype=torch.float32)
x_embed = not_mask.cumsum(2, dtype=torch.float32)
dim_t = torch.arange(self.num_pos_feats, dtype=torch.float32, device=x.device)
dim_t = self.temperature ** (2 * (dim_t // 2) / self.num_pos_feats)
pos_x = x_embed[:, :, :, None] / dim_t
pos_y = y_embed[:, :, :, None] / dim_t
pos_x = torch.stack((pos_x[:, :, :, 0::2].sin(),
pos_x[:, :, :, 1::2].cos()), dim=4).flatten(3)
pos_y = torch.stack((pos_y[:, :, :, 0::2].sin(),
pos_y[:, :, :, 1::2].cos()), dim=4).flatten(3)
pos = torch.cat((pos_y, pos_x), dim=3).permute(0, 3, 1, 2)
return pos
class DETR(nn.Module):
def __init__(self, num_classes=91, num_queries=100):
super().__init__()
self.backbone = Backbone()
self.transformer = Transformer()
self.query_embed = nn.Embedding(num_queries, 256)
self.pos_encoder = PositionEmbeddingSine()
# 预测头
self.class_embed = nn.Linear(256, num_classes + 1)
self.bbox_embed = MLP(256, 256, 4, 3)
def forward(self, x):
features = self.backbone(x) # [batch, 256, h, w]
mask = torch.zeros_like(x[:,0,:,:], dtype=torch.bool)
pos = self.pos_encoder(features)
# 展平特征图
features = features.flatten(2).permute(2, 0, 1) # [h*w, batch, 256]
pos = pos.flatten(2).permute(2, 0, 1)
hs = self.transformer(features, mask, self.query_embed.weight, pos)
outputs_class = self.class_embed(hs)
outputs_coord = self.bbox_embed(hs).sigmoid()
return {'pred_logits': outputs_class[-1],
'pred_boxes': outputs_coord[-1]}
3. 匈牙利匹配与损失函数
3.1 二分图匹配实现
DETR使用匈牙利算法进行预测框与真实框的最优匹配:
from scipy.optimize import linear_sum_assignment
def hungarian_matcher(outputs, targets):
bs, num_queries = outputs["pred_logits"].shape[:2]
# 展平batch维度
out_prob = outputs["pred_logits"].flatten(0, 1).softmax(-1) # [batch*num_queries, num_classes]
out_bbox = outputs["pred_boxes"].flatten(0, 1) # [batch*num_queries, 4]
indices = []
for i in range(bs):
tgt_ids = targets[i]["labels"]
tgt_bbox = targets[i]["boxes"]
# 计算类别损失
cost_class = -out_prob[i*num_queries:(i+1)*num_queries, tgt_ids]
# 计算框损失 (L1 + GIoU)
cost_bbox = torch.cdist(out_bbox[i*num_queries:(i+1)*num_queries], tgt_bbox, p=1)
cost_giou = -generalized_box_iou(box_cxcywh_to_xyxy(out_bbox[i*num_queries:(i+1)*num_queries]),
box_cxcywh_to_xyxy(tgt_bbox))
# 总成本矩阵
C = 1 * cost_class + 5 * cost_bbox + 2 * cost_giou
C = C.reshape(num_queries, -1).cpu()
# 匈牙利算法求解
indices.append(linear_sum_assignment(C))
return [(torch.as_tensor(i, dtype=torch.int64),
torch.as_tensor(j, dtype=torch.int64)) for i, j in indices]
3.2 损失函数设计
DETR的损失函数包含分类损失和框回归损失:
def detr_loss(outputs, targets, matcher):
# 获取匹配结果
indices = matcher(outputs, targets)
# 分类损失 (交叉熵)
src_logits = outputs['pred_logits']
idx = self._get_src_permutation_idx(indices)
target_classes = torch.full(src_logits.shape[:2], 91,
dtype=torch.int64, device=src_logits.device)
for i, (_, J) in enumerate(indices):
target_classes[i, J] = targets[i]['labels']
loss_ce = F.cross_entropy(src_logits.transpose(1, 2), target_classes)
# 框回归损失 (L1 + GIoU)
src_boxes = outputs['pred_boxes']
target_boxes = torch.cat([t['boxes'][i] for t, (_, i) in zip(targets, indices)], dim=0)
loss_bbox = F.l1_loss(src_boxes[idx], target_boxes, reduction='none')
loss_giou = 1 - torch.diag(generalized_box_iou(
box_cxcywh_to_xyxy(src_boxes[idx]),
box_cxcywh_to_xyxy(target_boxes)))
# 总损失
losses = {
'loss_ce': loss_ce,
'loss_bbox': loss_bbox.sum() / num_boxes,
'loss_giou': loss_giou.sum() / num_boxes
}
return losses
4. 训练技巧与优化策略
4.1 学习率调度
DETR对学习率设置敏感,推荐使用以下调度策略:
def build_optimizer(model, lr=1e-4, weight_decay=1e-4):
param_dicts = [
{"params": [p for n, p in model.named_parameters()
if "backbone" not in n and p.requires_grad]},
{"params": [p for n, p in model.named_parameters()
if "backbone" in n and p.requires_grad],
"lr": lr * 0.1},
]
optimizer = torch.optim.AdamW(param_dicts, lr=lr, weight_decay=weight_decay)
# 学习率预热
lr_scheduler = torch.optim.lr_scheduler.SequentialLR(
optimizer,
[
torch.optim.lr_scheduler.LinearLR(optimizer, 0.0001, 1.0, 100),
torch.optim.lr_scheduler.MultiStepLR(optimizer, [40], gamma=0.1)
],
[100]
)
return optimizer, lr_scheduler
4.2 梯度裁剪与混合精度训练
为稳定训练过程,建议启用梯度裁剪和混合精度:
scaler = torch.cuda.amp.GradScaler()
for epoch in range(epochs):
for images, targets in dataloader:
optimizer.zero_grad()
with torch.cuda.amp.autocast():
outputs = model(images)
loss_dict = criterion(outputs, targets)
losses = sum(loss_dict.values())
scaler.scale(losses).backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 0.1)
scaler.step(optimizer)
scaler.update()
lr_scheduler.step()
4.3 关键训练参数配置
| 参数 | 推荐值 | 说明 |
|---|---|---|
| batch_size | 16-32 | 根据GPU显存调整 |
| base_lr | 1e-4 | 主干网络使用0.1倍 |
| weight_decay | 1e-4 | 防止过拟合 |
| epochs | 50-300 | 小数据集需更多epoch |
| warmup_epochs | 100 | 学习率线性预热 |
| dropout | 0.1 | Transformer层dropout率 |
| num_queries | 100 | 默认预测框数量 |
5. 模型推理与部署
5.1 预测后处理
DETR推理过程无需NMS,直接过滤低置信度预测:
def postprocess(outputs, threshold=0.7):
probas = outputs['pred_logits'].softmax(-1)[:, :, :-1] # 去除背景类
keep = probas.max(-1).values > threshold
results = []
for p, b, k in zip(probas, outputs['pred_boxes'], keep):
if k.sum() == 0:
results.append({'scores': [], 'labels': [], 'boxes': []})
continue
scores = p[k]
labels = scores.argmax(-1)
boxes = b[k]
results.append({
'scores': scores.max(-1).values,
'labels': labels,
'boxes': boxes
})
return results
5.2 ONNX导出与优化
将模型导出为ONNX格式便于部署:
dummy_input = torch.randn(1, 3, 800, 800)
torch.onnx.export(
model,
dummy_input,
"detr.onnx",
input_names=["input"],
output_names=["logits", "boxes"],
dynamic_axes={
"input": {0: "batch"},
"logits": {0: "batch"},
"boxes": {0: "batch"}
},
opset_version=12
)
5.3 TensorRT加速
使用TensorRT进一步优化推理速度:
# 转换ONNX到TensorRT引擎
trtexec --onnx=detr.onnx --saveEngine=detr.engine \
--fp16 --workspace=4096 --minShapes=input:1x3x800x800 \
--optShapes=input:4x3x800x800 --maxShapes=input:8x3x800x800
# Python加载引擎
with open("detr.engine", "rb") as f:
runtime = trt.Runtime(trt.Logger(trt.Logger.WARNING))
engine = runtime.deserialize_cuda_engine(f.read())
6. 常见问题与解决方案
6.1 训练收敛慢问题
现象:模型在初期训练阶段损失下降缓慢,AP提升不明显。
解决方案:
- 确保正确实现了学习率预热(前100个step从1e-6线性增加到1e-4)
- 检查主干网络是否冻结了BatchNorm层的统计量
- 增加梯度裁剪(norm=0.1)防止初期梯度爆炸
- 验证匈牙利匹配实现是否正确,特别是GIoU计算部分
6.2 小目标检测效果差
现象:模型对大目标检测效果良好,但小目标漏检率高。
优化策略:
- 在主干网络后添加FPN结构,融合多尺度特征
- 减小输入图像的下采样倍数(如从32倍降到16倍)
- 增加小目标在训练数据中的采样权重
- 使用Deformable DETR等改进版本,专门优化小目标检测
6.3 显存不足问题
现象:训练时出现CUDA out of memory错误。
应对措施:
- 减小batch size(最低可到4,配合梯度累积)
- 使用混合精度训练(AMP)
- 启用checkpointing技术减少中间缓存
- 简化Transformer结构(如减少层数或隐藏维度)
# 梯度累积示例
accumulation_steps = 4
optimizer.zero_grad()
for i, (images, targets) in enumerate(dataloader):
outputs = model(images)
loss = criterion(outputs, targets)
loss = loss / accumulation_steps
loss.backward()
if (i+1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
7. 进阶优化方向
7.1 模型轻量化
对于移动端部署,可考虑以下优化:
- 替换主干网络:使用MobileNetV3或EfficientNet-Lite替代ResNet
- 知识蒸馏:用大模型指导小模型训练
- 量化感知训练:直接训练8整型模型
- 剪枝:移除冗余的注意力头或FFN层
7.2 多任务扩展
DETR框架可轻松扩展其他视觉任务:
- 实例分割:在解码器输出添加掩码头
- 姿态估计:预测关键点坐标
- 全景分割:统一语义分割和实例分割
- 视频目标检测:加入时序建模模块
7.3 最新改进版本
| 改进版本 | 核心创新 | 优势 |
|---|---|---|
| Deformable DETR | 可变形注意力机制 | 提升小目标检测,加速收敛 |
| DAB-DETR | 动态锚框查询 | 更稳定的训练过程 |
| DN-DETR | 去噪训练策略 | 解决二分图匹配不稳定性 |
| DINO-DETR | 对比学习预训练 | 更高检测精度 |
实际项目中,根据具体需求选择合适的DETR变体往往能获得更好的效果。例如,对于实时性要求高的场景,可选用Conditional DETR;对于小目标密集场景,Deformable DETR是更好的选择。
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐

所有评论(0)