DETR目标检测实战:手把手教你用Transformer实现端到端检测(附代码)
DETR目标检测实战:手把手教你用Transformer实现端到端检测(附代码)
在计算机视觉领域,目标检测一直是核心任务之一。传统方法如Faster R-CNN、YOLO等虽然效果显著,但都依赖于复杂的后处理流程(如非极大值抑制NMS)和人工设计的锚框(anchor)。DETR(Detection Transformer)的出现彻底改变了这一局面,它首次将Transformer架构引入目标检测领域,实现了真正的端到端检测。
本文将带你从零开始搭建和训练DETR模型,涵盖环境配置、数据准备、模型训练和推理全流程。我们特别关注实际开发中可能遇到的坑点,并提供经过验证的解决方案。无论你是刚接触目标检测的新手,还是希望了解前沿技术的资深开发者,都能从中获得实用价值。
1. 环境配置与依赖安装
搭建DETR开发环境需要准备以下核心组件:
- PyTorch 1.5+:DETR完全基于PyTorch实现
- Torchvision 0.6+:用于图像预处理和数据增强
- OpenCV/Pillow:基础图像处理库
- pycocotools:COCO数据集评估工具
推荐使用conda创建虚拟环境:
conda create -n detr python=3.8
conda activate detr
pip install torch torchvision torchaudio
pip install opencv-python pillow
pip install pycocotools
对于GPU加速,需要额外安装对应版本的CUDA和cuDNN。以CUDA 11.3为例:
conda install cudatoolkit=11.3 -c pytorch
验证环境是否配置成功:
import torch
print(torch.__version__, torch.cuda.is_available()) # 应输出PyTorch版本和True
常见问题排查:
- 如果遇到
CUDA out of memory错误,尝试减小batch size pycocotools安装失败时,可先安装cython再重试- Windows用户可能需要从源码编译安装pycocotools
2. 数据准备与预处理
DETR官方使用COCO数据集进行训练和评估。我们可以通过以下步骤准备数据:
- 下载COCO数据集(约18GB训练集+1GB验证集)
- 解压到指定目录,保持如下结构:
coco/ ├── annotations │ ├── instances_train2017.json │ └── instances_val2017.json ├── train2017 └── val2017 - 实现自定义Dataset类:
from torch.utils.data import Dataset
from pycocotools.coco import COCO
import cv2
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):
img_id = self.ids[idx]
ann_ids = self.coco.getAnnIds(imgIds=img_id)
annotations = self.coco.loadAnns(ann_ids)
img_info = self.coco.loadImgs(img_id)[0]
img_path = os.path.join(self.img_folder, img_info['file_name'])
img = cv2.imread(img_path)
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
boxes = [ann['bbox'] for ann in annotations] # [x,y,w,h]
boxes = torch.as_tensor(boxes, dtype=torch.float32)
boxes[:, 2:] += boxes[:, :2] # 转换为[x1,y1,x2,y2]格式
labels = torch.tensor(
[ann['category_id'] for ann in annotations], dtype=torch.int64)
target = {
'boxes': boxes,
'labels': labels,
'image_id': torch.tensor([img_id]),
'area': torch.tensor([ann['area'] for ann in annotations]),
'iscrowd': torch.tensor([ann['iscrowd'] for ann in annotations])
}
if self.transforms is not None:
img, target = self.transforms(img, target)
return img, target
关键预处理步骤包括:
- 图像归一化(均值[0.485, 0.456, 0.406],标准差[0.229, 0.224, 0.225])
- 随机水平翻转(增强数据多样性)
- 调整图像大小(保持长宽比,短边至少800像素)
3. 模型构建详解
DETR的核心架构包含四个部分:CNN骨干网络、Transformer编码器、Transformer解码器和预测头。让我们逐步实现:
3.1 骨干网络(Backbone)
默认使用ResNet-50作为特征提取器:
import torchvision
from torch import nn
class Backbone(nn.Module):
def __init__(self, name='resnet50', train_backbone=True, dilation=False):
super().__init__()
backbone = getattr(torchvision.models, name)(
replace_stride_with_dilation=[False, False, dilation],
pretrained=True)
# 冻结部分层参数
for name, parameter in backbone.named_parameters():
if not train_backbone or 'layer2' not in name and 'layer3' not in name and 'layer4' not in name:
parameter.requires_grad_(False)
self.body = backbone
self.num_channels = 2048 if name in ('resnet50', 'resnet101') else 512
def forward(self, x):
# ResNet的前向传播
x = self.body.conv1(x)
x = self.body.bn1(x)
x = self.body.relu(x)
x = self.body.maxpool(x)
x = self.body.layer1(x)
x = self.body.layer2(x)
x = self.body.layer3(x)
x = self.body.layer4(x)
return x
3.2 Transformer架构
实现Transformer编码器-解码器结构:
from torch.nn import MultiheadAttention
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)
self._reset_parameters()
def _reset_parameters(self):
for p in self.parameters():
if p.dim() > 1:
nn.init.xavier_uniform_(p)
def forward(self, src, mask, query_embed, pos_embed):
# 展平空间维度 (batch_size, c, h, w) -> (h*w, batch_size, c)
bs, c, h, w = src.shape
src = src.flatten(2).permute(2, 0, 1)
pos_embed = pos_embed.flatten(2).permute(2, 0, 1)
# 初始化query嵌入
query_embed = query_embed.unsqueeze(1).repeat(1, bs, 1)
# Transformer编码器
memory = self.encoder(src, src_key_padding_mask=mask, pos=pos_embed)
# Transformer解码器
tgt = torch.zeros_like(query_embed)
hs = self.decoder(tgt, memory, memory_key_padding_mask=mask,
pos=pos_embed, query_pos=query_embed)
return hs.transpose(1, 2), memory.permute(1, 2, 0).view(bs, c, h, w)
3.3 位置编码
DETR使用可学习的位置编码:
class PositionEmbeddingSine(nn.Module):
def __init__(self, num_pos_feats=64, temperature=10000, normalize=False):
super().__init__()
self.num_pos_feats = num_pos_feats
self.temperature = temperature
self.normalize = normalize
def forward(self, x, mask):
not_mask = ~mask
y_embed = not_mask.cumsum(1, dtype=torch.float32)
x_embed = not_mask.cumsum(2, dtype=torch.float32)
if self.normalize:
eps = 1e-6
y_embed = y_embed / (y_embed[:, -1:, :] + eps) * 2 * math.pi
x_embed = x_embed / (x_embed[:, :, -1:] + eps) * 2 * math.pi
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
4. 训练流程与技巧
DETR的训练有几个关键点需要特别注意:
4.1 损失函数实现
匈牙利匹配损失是DETR的核心创新:
import torch.nn.functional as F
from scipy.optimize import linear_sum_assignment
class HungarianMatcher(nn.Module):
def __init__(self, cost_class=1, cost_bbox=5, cost_giou=2):
super().__init__()
self.cost_class = cost_class
self.cost_bbox = cost_bbox
self.cost_giou = cost_giou
@torch.no_grad()
def forward(self, outputs, targets):
bs, num_queries = outputs["pred_logits"].shape[:2]
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[:, tgt_ids]
# 计算L1成本
cost_bbox = torch.cdist(out_bbox, tgt_bbox, p=1)
# 计算GIoU成本
cost_giou = -generalized_box_iou(
box_cxcywh_to_xyxy(out_bbox),
box_cxcywh_to_xyxy(tgt_bbox))
# 最终成本矩阵
C = self.cost_bbox * cost_bbox + self.cost_class * cost_class + self.cost_giou * cost_giou
C = C.view(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]
4.2 训练循环示例
def train_one_epoch(model, criterion, optimizer, data_loader, device, epoch):
model.train()
criterion.train()
for images, targets in data_loader:
images = images.to(device)
targets = [{k: v.to(device) for k, v in t.items()} for t in targets]
outputs = model(images)
loss_dict = criterion(outputs, targets)
losses = sum(loss_dict[k] for k in loss_dict.keys())
optimizer.zero_grad()
losses.backward()
optimizer.step()
# 打印训练日志
if step % args.print_freq == 0:
print(f"Epoch {epoch} Step {step} Loss: {losses.item():.4f}")
4.3 学习率调度
DETR对学习率非常敏感,推荐使用以下调度策略:
def get_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.StepLR(optimizer, step_size=30, gamma=0.1)
return optimizer, lr_scheduler
5. 推理与可视化
训练完成后,我们可以使用模型进行预测并可视化结果:
@torch.no_grad()
def inference(model, image, transform, device):
model.eval()
img_tensor = transform(image).unsqueeze(0).to(device)
outputs = model(img_tensor)
# 过滤低置信度预测
probas = outputs['pred_logits'].softmax(-1)[0, :, :-1]
keep = probas.max(-1).values > 0.7
# 转换框格式
boxes = outputs['pred_boxes'][0, keep].cpu()
boxes = box_cxcywh_to_xyxy(boxes)
# 获取类别和分数
scores, labels = probas[keep].max(-1)
return boxes, labels, scores
def visualize(image, boxes, labels, scores, class_names):
plt.figure(figsize=(16,10))
plt.imshow(image)
ax = plt.gca()
for box, label, score in zip(boxes, labels, scores):
x1, y1, x2, y2 = box
ax.add_patch(plt.Rectangle((x1, y1), x2-x1, y2-y1,
fill=False, color='red', linewidth=2))
text = f"{class_names[label]}: {score:.2f}"
ax.text(x1, y1, text, fontsize=10,
bbox=dict(facecolor='yellow', alpha=0.5))
plt.axis('off')
plt.show()
6. 性能优化技巧
在实际应用中,我们可以通过以下方法提升DETR的性能和效率:
- 多尺度训练:在训练时随机调整图像尺寸(480-800像素),增强模型对不同尺度目标的检测能力
- 梯度裁剪:防止梯度爆炸,稳定训练过程
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.1) - 混合精度训练:减少显存占用,加速训练
scaler = torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): outputs = model(images) loss = criterion(outputs, targets) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() - 模型蒸馏:使用更大的DETR模型(如DETR-DC5)作为教师模型,指导基础模型训练
7. 常见问题与解决方案
在DETR的实际应用中,开发者常会遇到以下问题:
问题1:训练初期损失不下降
- 原因:Transformer需要更长的预热期
- 解决方案:使用更小的学习率(如1e-5)并延长训练周期
问题2:小物体检测效果差
- 原因:高维特征丢失了小物体信息
- 解决方案:
- 使用更高分辨率的输入图像
- 采用特征金字塔结构(如FPN)增强多尺度特征
问题3:训练显存不足
- 解决方案:
- 减小batch size(最低可设为1)
- 使用梯度累积技术
- 启用checkpointing减少内存占用
问题4:推理速度慢
- 优化方案:
- 量化模型权重(FP16或INT8)
- 使用TensorRT加速
- 剪枝去除冗余注意力头
8. 进阶改进方向
对于希望进一步提升DETR性能的开发者,可以考虑以下研究方向:
- Deformable DETR:引入可变形注意力机制,显著提升小物体检测性能
- Conditional DETR:改进query设计,加速模型收敛
- DAB-DETR:将query明确表示为动态锚框,提升定位精度
- DN-DETR:通过去噪训练解决二分图匹配不稳定的问题
- DINO-DETR:结合对比学习,提升模型表征能力
每个改进方向都有对应的开源实现,开发者可以根据具体需求选择适合的变体。例如,要实现Deformable DETR,只需替换标准Transformer中的注意力模块:
from deformable_detr import DeformableTransformerEncoderLayer
encoder_layer = DeformableTransformerEncoderLayer(
d_model=256, nhead=8, dim_feedforward=1024,
dropout=0.1, activation="relu",
n_levels=4, n_points=4)
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐
所有评论(0)