手把手带你复现DETR:用PyTorch从零搭建你的第一个Transformer检测模型
手把手带你复现DETR:用PyTorch从零搭建你的第一个Transformer检测模型
在计算机视觉领域,目标检测一直是一个核心任务。传统方法如Faster R-CNN、YOLO等基于卷积神经网络(CNN)的检测器虽然效果显著,但往往需要复杂的后处理和非极大值抑制(NMS)操作。2020年,Facebook AI提出的DETR(DEtection TRansformer)彻底改变了这一局面,首次将Transformer架构成功应用于目标检测任务,实现了端到端的检测流程。
本文将带你从零开始,用PyTorch实现一个完整的DETR模型。不同于简单的API调用,我们会深入每个模块的实现细节,包括:
- 如何构建高效的ResNet特征提取器
- 位置编码(Positional Encoding)的设计与实现
- Transformer编码器-解码器架构的搭建
- Object Queries的初始化与优化
- 匈牙利匹配损失函数的实现
- 训练技巧与调试方法
通过这个实践过程,你不仅能掌握DETR的核心思想,还能深入理解Transformer在视觉任务中的应用方式。让我们开始这段代码之旅吧!
1. 环境准备与数据加载
在开始构建模型前,我们需要准备好开发环境。推荐使用Python 3.8+和PyTorch 1.10+版本。可以通过以下命令安装必要的依赖:
pip install torch torchvision torchaudio
pip install opencv-python matplotlib tqdm
对于数据集,我们将使用COCO格式的数据。如果你没有现成的数据集,可以从COCO官网下载或使用torchvision自带的简化版本:
from torchvision.datasets import CocoDetection
class CocoDetectionWithTransform(CocoDetection):
def __init__(self, root, annFile, transform=None):
super().__init__(root, annFile)
self.transform = transform
def __getitem__(self, idx):
img, target = super().__getitem__(idx)
if self.transform is not None:
img = self.transform(img)
return img, target
数据预处理是目标检测的关键环节。我们需要定义一组标准的转换操作:
from torchvision.transforms import Compose, Resize, ToTensor, Normalize
def get_transform(train=True):
transforms = []
transforms.append(Resize((800, 800)))
transforms.append(ToTensor())
transforms.append(Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]))
return Compose(transforms)
注意:在实际应用中,你可能需要添加更多的数据增强技术,如随机裁剪、颜色抖动等,特别是在训练数据有限的情况下。
2. 构建特征提取Backbone
DETR使用CNN作为特征提取的backbone,通常选择ResNet架构。我们将实现一个简化版的ResNet,重点关注与DETR配合的关键部分:
import torch
import torch.nn as nn
import torchvision.models as models
class Backbone(nn.Module):
def __init__(self, backbone_name='resnet50', pretrained=True):
super().__init__()
backbone = getattr(models, backbone_name)(pretrained=pretrained)
self.conv1 = backbone.conv1
self.bn1 = backbone.bn1
self.relu = backbone.relu
self.maxpool = backbone.maxpool
self.layer1 = backbone.layer1
self.layer2 = backbone.layer2
self.layer3 = backbone.layer3
self.layer4 = backbone.layer4
def forward(self, x):
x = self.conv1(x)
x = self.bn1(x)
x = self.relu(x)
x = self.maxpool(x)
x1 = self.layer1(x)
x2 = self.layer2(x1)
x3 = self.layer3(x2)
x4 = self.layer4(x3)
return x4 # 返回最高层特征图
在实际DETR实现中,我们通常需要从backbone中提取多尺度特征。这里我们简化处理,只使用最后一层特征。完整的实现可以参考官方代码,添加特征金字塔网络(FPN)结构。
3. 位置编码的实现
Transformer本身不具备位置感知能力,因此需要显式地添加位置信息。DETR采用了与原始Transformer类似的正弦位置编码:
import math
class PositionEmbeddingSine(nn.Module):
def __init__(self, num_pos_feats=64, temperature=10000, normalize=False, scale=None):
super().__init__()
self.num_pos_feats = num_pos_feats
self.temperature = temperature
self.normalize = normalize
if scale is not None and normalize is False:
raise ValueError("normalize should be True if scale is passed")
if scale is None:
scale = 2 * math.pi
self.scale = scale
def forward(self, x):
# x: [batch, channels, height, width]
batch, _, height, width = x.shape
mask = torch.zeros((batch, height, width), dtype=torch.bool, device=x.device)
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) * self.scale
x_embed = x_embed / (x_embed[:, :, -1:] + eps) * self.scale
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
这个位置编码模块会生成与输入特征图相同空间维度的位置信息,可以方便地与CNN特征相加融合。
4. Transformer架构实现
DETR的核心是Transformer架构。我们将分别实现编码器和解码器部分:
4.1 多头注意力机制
首先实现Transformer的基础模块——多头注意力:
class MultiHeadAttention(nn.Module):
def __init__(self, embed_dim, num_heads, dropout=0.0):
super().__init__()
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
assert self.head_dim * num_heads == embed_dim, "embed_dim must be divisible by num_heads"
self.q_proj = nn.Linear(embed_dim, embed_dim)
self.k_proj = nn.Linear(embed_dim, embed_dim)
self.v_proj = nn.Linear(embed_dim, embed_dim)
self.out_proj = nn.Linear(embed_dim, embed_dim)
self.dropout = nn.Dropout(dropout)
def forward(self, query, key, value, key_padding_mask=None):
batch_size = query.size(0)
# 线性变换并分头
q = self.q_proj(query).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(key).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(value).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
# 计算注意力分数
attn_scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim)
# 应用mask(如果有)
if key_padding_mask is not None:
attn_scores = attn_scores.masked_fill(
key_padding_mask.unsqueeze(1).unsqueeze(2),
float('-inf'))
# 计算注意力权重
attn_weights = torch.softmax(attn_scores, dim=-1)
attn_weights = self.dropout(attn_weights)
# 加权求和
output = torch.matmul(attn_weights, v)
# 合并多头
output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.embed_dim)
output = self.out_proj(output)
return output, attn_weights
4.2 Transformer编码器
基于多头注意力,我们可以构建Transformer编码器层:
class TransformerEncoderLayer(nn.Module):
def __init__(self, embed_dim, num_heads, dim_feedforward=2048, dropout=0.1):
super().__init__()
self.self_attn = MultiHeadAttention(embed_dim, num_heads, dropout)
self.linear1 = nn.Linear(embed_dim, dim_feedforward)
self.dropout = nn.Dropout(dropout)
self.linear2 = nn.Linear(dim_feedforward, embed_dim)
self.norm1 = nn.LayerNorm(embed_dim)
self.norm2 = nn.LayerNorm(embed_dim)
self.dropout1 = nn.Dropout(dropout)
self.dropout2 = nn.Dropout(dropout)
self.activation = nn.ReLU()
def forward(self, src, src_key_padding_mask=None):
# 自注意力
src2, attn_weights = self.self_attn(src, src, src, src_key_padding_mask)
src = src + self.dropout1(src2)
src = self.norm1(src)
# 前馈网络
src2 = self.linear2(self.dropout(self.activation(self.linear1(src))))
src = src + self.dropout2(src2)
src = self.norm2(src)
return src, attn_weights
完整的编码器由多个这样的层堆叠而成:
class TransformerEncoder(nn.Module):
def __init__(self, encoder_layer, num_layers):
super().__init__()
self.layers = nn.ModuleList([copy.deepcopy(encoder_layer) for _ in range(num_layers)])
def forward(self, src, src_key_padding_mask=None):
output = src
attn_weights = []
for layer in self.layers:
output, weights = layer(output, src_key_padding_mask)
attn_weights.append(weights)
return output, attn_weights
4.3 Transformer解码器
解码器部分稍微复杂一些,因为它需要处理两种注意力机制:
class TransformerDecoderLayer(nn.Module):
def __init__(self, embed_dim, num_heads, dim_feedforward=2048, dropout=0.1):
super().__init__()
self.self_attn = MultiHeadAttention(embed_dim, num_heads, dropout)
self.multihead_attn = MultiHeadAttention(embed_dim, num_heads, dropout)
self.linear1 = nn.Linear(embed_dim, dim_feedforward)
self.dropout = nn.Dropout(dropout)
self.linear2 = nn.Linear(dim_feedforward, embed_dim)
self.norm1 = nn.LayerNorm(embed_dim)
self.norm2 = nn.LayerNorm(embed_dim)
self.norm3 = nn.LayerNorm(embed_dim)
self.dropout1 = nn.Dropout(dropout)
self.dropout2 = nn.Dropout(dropout)
self.dropout3 = nn.Dropout(dropout)
self.activation = nn.ReLU()
def forward(self, tgt, memory, tgt_mask=None, memory_mask=None,
tgt_key_padding_mask=None, memory_key_padding_mask=None):
# 自注意力
tgt2, self_attn_weights = self.self_attn(
tgt, tgt, tgt, tgt_key_padding_mask)
tgt = tgt + self.dropout1(tgt2)
tgt = self.norm1(tgt)
# 编码器-解码器注意力
tgt2, cross_attn_weights = self.multihead_attn(
tgt, memory, memory, memory_key_padding_mask)
tgt = tgt + self.dropout2(tgt2)
tgt = self.norm2(tgt)
# 前馈网络
tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt))))
tgt = tgt + self.dropout3(tgt2)
tgt = self.norm3(tgt)
return tgt, self_attn_weights, cross_attn_weights
完整的解码器同样由多个这样的层组成:
class TransformerDecoder(nn.Module):
def __init__(self, decoder_layer, num_layers):
super().__init__()
self.layers = nn.ModuleList([copy.deepcopy(decoder_layer) for _ in range(num_layers)])
def forward(self, tgt, memory, tgt_mask=None, memory_mask=None,
tgt_key_padding_mask=None, memory_key_padding_mask=None):
output = tgt
self_attn_weights = []
cross_attn_weights = []
for layer in self.layers:
output, self_attn, cross_attn = layer(
output, memory, tgt_mask, memory_mask,
tgt_key_padding_mask, memory_key_padding_mask)
self_attn_weights.append(self_attn)
cross_attn_weights.append(cross_attn)
return output, self_attn_weights, cross_attn_weights
5. Object Queries与预测头
DETR的一个关键创新是使用可学习的Object Queries来替代传统检测器中的anchor:
class ObjectQueries(nn.Module):
def __init__(self, num_queries=100, embed_dim=256):
super().__init__()
self.queries = nn.Parameter(torch.randn(num_queries, embed_dim))
def forward(self, batch_size):
return self.queries.unsqueeze(0).expand(batch_size, -1, -1)
预测头负责将解码器的输出转换为最终的检测结果:
class DetectionHead(nn.Module):
def __init__(self, embed_dim, num_classes):
super().__init__()
self.class_embed = nn.Linear(embed_dim, num_classes + 1) # +1 for background
self.bbox_embed = MLP(embed_dim, embed_dim, 4, 3)
def forward(self, x):
class_logits = self.class_embed(x)
bbox_coords = self.bbox_embed(x).sigmoid() # 归一化到[0,1]
return {'pred_logits': class_logits, 'pred_boxes': bbox_coords}
class MLP(nn.Module):
def __init__(self, input_dim, hidden_dim, output_dim, num_layers):
super().__init__()
layers = []
for i in range(num_layers):
layers.append(nn.Linear(
hidden_dim if i > 0 else input_dim,
hidden_dim if i < num_layers - 1 else output_dim))
if i < num_layers - 1:
layers.append(nn.ReLU())
self.layers = nn.Sequential(*layers)
def forward(self, x):
return self.layers(x)
6. 匈牙利匹配损失实现
DETR使用二分图匹配来确定预测框与真实框的对应关系,然后计算损失:
def hungarian_matcher(pred_logits, pred_boxes, targets):
"""
pred_logits: [batch_size, num_queries, num_classes+1]
pred_boxes: [batch_size, num_queries, 4]
targets: list of dict with keys 'labels' and 'boxes'
"""
batch_size = pred_logits.size(0)
num_queries = pred_logits.size(1)
indices = []
for i in range(batch_size):
# 计算分类损失
cost_class = -pred_logits[i].softmax(-1)[:, targets[i]['labels']]
# 计算L1和IoU损失
cost_bbox = torch.cdist(pred_boxes[i], targets[i]['boxes'], p=1)
cost_giou = -generalized_box_iou(box_cxcywh_to_xyxy(pred_boxes[i]),
box_cxcywh_to_xyxy(targets[i]['boxes']))
# 总成本
C = 1 * cost_class + 5 * cost_bbox + 2 * cost_giou
C = C.reshape(num_queries, -1).cpu()
# 匈牙利算法匹配
with torch.no_grad():
indices_i = linear_sum_assignment(C)
indices.append((torch.as_tensor(indices_i[0], dtype=torch.int64),
torch.as_tensor(indices_i[1], dtype=torch.int64)))
return indices
def box_cxcywh_to_xyxy(x):
x_c, y_c, w, h = x.unbind(-1)
b = [(x_c - 0.5 * w), (y_c - 0.5 * h),
(x_c + 0.5 * w), (y_c + 0.5 * h)]
return torch.stack(b, dim=-1)
def generalized_box_iou(boxes1, boxes2):
"""
Generalized IoU from https://giou.stanford.edu/
boxes1: [N,4]
boxes2: [M,4]
"""
# 计算标准IoU
inter = box_intersection(boxes1, boxes2)
area1 = box_area(boxes1)
area2 = box_area(boxes2)
union = area1.unsqueeze(1) + area2.unsqueeze(0) - inter
iou = inter / union
# 计算最小闭合框面积
lt = torch.min(boxes1[:, None, :2], boxes2[:, :2])
rb = torch.max(boxes1[:, None, 2:], boxes2[:, 2:])
wh = (rb - lt).clamp(min=0)
area = wh[:, :, 0] * wh[:, :, 1]
return iou - (area - union) / area
7. 完整DETR模型组装
现在我们可以将所有组件组合成完整的DETR模型:
class DETR(nn.Module):
def __init__(self, backbone, transformer, num_classes, num_queries):
super().__init__()
self.backbone = backbone
self.transformer = transformer
self.num_queries = num_queries
# 将CNN特征映射到Transformer的embed_dim
hidden_dim = transformer.embed_dim
self.conv = nn.Conv2d(backbone.num_channels, hidden_dim, 1)
# 位置编码
self.position_embedding = PositionEmbeddingSine(hidden_dim // 2)
# Object Queries
self.query_embed = ObjectQueries(num_queries, hidden_dim)
# 预测头
self.head = DetectionHead(hidden_dim, num_classes)
def forward(self, x):
# 特征提取
features = self.backbone(x)
# 调整特征维度
features = self.conv(features)
batch_size = features.size(0)
# 位置编码
pos_embed = self.position_embedding(features)
# 准备Transformer输入
features = features.flatten(2).permute(0, 2, 1)
pos_embed = pos_embed.flatten(2).permute(0, 2, 1)
# Object Queries
query_embed = self.query_embed(batch_size)
# Transformer前向传播
hs = self.transformer(features, query_embed, pos_embed)
# 预测
outputs = self.head(hs)
return outputs
8. 训练策略与技巧
训练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},
]
return torch.optim.AdamW(param_dicts, lr=lr, weight_decay=weight_decay)
def train_one_epoch(model, criterion, data_loader, optimizer, device, epoch, max_norm=0.1):
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.values())
optimizer.zero_grad()
losses.backward()
if max_norm > 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
optimizer.step()
提示:DETR训练初期loss下降较慢是正常现象,通常需要50-100个epoch才能看到明显效果。可以使用预训练权重加速收敛。
9. 可视化与调试
理解模型行为的关键是可视化注意力权重:
def plot_attention_weights(image, attn_weights, query_idx=0, head_idx=0):
"""
image: 原始图像 [3, H, W]
attn_weights: 注意力权重 [num_layers, batch, num_heads, num_queries, H*W]
"""
fig, ax = plt.subplots(1, 2, figsize=(10, 5))
# 显示原始图像
ax[0].imshow(image.permute(1, 2, 0))
ax[0].axis('off')
ax[0].set_title('Original Image')
# 显示注意力热图
h, w = image.shape[1], image.shape[2]
attn = attn_weights[-1][0, head_idx, query_idx].view(h, w)
ax[1].imshow(attn.detach().cpu(), cmap='hot')
ax[1].axis('off')
ax[1].set_title(f'Attention Head {head_idx} Query {query_idx}')
plt.show()
在实际项目中,你可能还需要实现以下调试工具:
- 预测框可视化
- 损失曲线监控
- 学习率调度可视化
- 梯度流动分析
10. 常见问题与解决方案
在实现和训练DETR过程中,你可能会遇到以下问题:
-
训练初期loss不下降
- 检查学习率是否合适
- 确保数据预处理正确
- 尝试使用预训练backbone
-
模型收敛后性能不佳
- 增加训练epoch
- 调整匈牙利匹配的损失权重
- 尝试更大的模型或更多queries
-
GPU内存不足
- 减小batch size
- 使用混合精度训练
- 降低输入图像分辨率
-
训练不稳定
- 添加梯度裁剪
- 调整学习率调度策略
- 检查数据中是否存在异常样本
通过本文的实现,你应该已经掌握了DETR的核心思想和实现细节。虽然完整复现原论文结果需要更多工程优化,但这个基础版本已经包含了所有关键组件。建议你在理解这个实现后,尝试添加以下改进:
- 多尺度特征融合
- 可变形注意力机制
- 更高效的位置编码
- 知识蒸馏等训练技巧
更多推荐




所有评论(0)