手把手带你复现DETR:用PyTorch从零搭建你的第一个Transformer检测模型

在计算机视觉领域,目标检测一直是一个核心任务。传统方法如Faster R-CNN、YOLO等基于卷积神经网络(CNN)的检测器虽然效果显著,但往往需要复杂的后处理和非极大值抑制(NMS)操作。2020年,Facebook AI提出的DETR(DEtection TRansformer)彻底改变了这一局面,首次将Transformer架构成功应用于目标检测任务,实现了端到端的检测流程。

本文将带你从零开始,用PyTorch实现一个完整的DETR模型。不同于简单的API调用,我们会深入每个模块的实现细节,包括:

  1. 如何构建高效的ResNet特征提取器
  2. 位置编码(Positional Encoding)的设计与实现
  3. Transformer编码器-解码器架构的搭建
  4. Object Queries的初始化与优化
  5. 匈牙利匹配损失函数的实现
  6. 训练技巧与调试方法

通过这个实践过程,你不仅能掌握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过程中,你可能会遇到以下问题:

  1. 训练初期loss不下降

    • 检查学习率是否合适
    • 确保数据预处理正确
    • 尝试使用预训练backbone
  2. 模型收敛后性能不佳

    • 增加训练epoch
    • 调整匈牙利匹配的损失权重
    • 尝试更大的模型或更多queries
  3. GPU内存不足

    • 减小batch size
    • 使用混合精度训练
    • 降低输入图像分辨率
  4. 训练不稳定

    • 添加梯度裁剪
    • 调整学习率调度策略
    • 检查数据中是否存在异常样本

通过本文的实现,你应该已经掌握了DETR的核心思想和实现细节。虽然完整复现原论文结果需要更多工程优化,但这个基础版本已经包含了所有关键组件。建议你在理解这个实现后,尝试添加以下改进:

  • 多尺度特征融合
  • 可变形注意力机制
  • 更高效的位置编码
  • 知识蒸馏等训练技巧
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐