DETR目标检测实战:从零搭建Transformer检测模型(附PyTorch代码)

在计算机视觉领域,目标检测一直是核心任务之一。传统方法如Faster R-CNN、YOLO系列虽然成熟,但依赖手工设计的锚框(Anchor)和非极大值抑制(NMS)后处理,流程复杂且超参数敏感。2020年Facebook提出的DETR(Detection Transformer)彻底改变了这一局面,首次将Transformer架构引入目标检测,实现了真正的端到端训练。本文将带您从零开始,用PyTorch实现一个精简版DETR模型,并深入解析关键技术的实现细节。

1. 环境准备与数据加载

实现DETR模型需要以下核心依赖:

# 基础环境配置
pip install torch==1.12.0+cu113 torchvision==0.13.0+cu113 -f https://download.pytorch.org/whl/torch_stable.html
pip install pycocotools matplotlib scipy

对于目标检测任务,COCO数据集是最常用的基准测试集。我们使用torchvision内置的数据加载器:

from torchvision.datasets import CocoDetection

class CocoDetectionWithTransform(CocoDetection):
    def __init__(self, root, annFile, transforms):
        super().__init__(root, annFile)
        self._transforms = transforms
        
    def __getitem__(self, idx):
        img, target = super().__getitem__(idx)
        image_id = self.ids[idx]
        target = {'image_id': image_id, 'annotations': target}
        if self._transforms is not None:
            img, target = self._transforms(img, target)
        return img, target

注意:COCO数据集的标注格式需要转换为DETR所需的格式。每个标注应包含类别标签和边界框坐标(cx, cy, w, h格式)

数据增强策略对DETR训练至关重要,推荐使用以下组合:

  • 随机水平翻转(p=0.5)
  • 随机缩放(短边在480-800像素之间)
  • 随机裁剪(确保目标完整)
  • 归一化(ImageNet均值方差)

2. 模型架构实现

DETR的核心由三部分组成:CNN骨干网络、Transformer编码器-解码器结构和预测头。我们先实现最关键的Transformer部分:

2.1 Transformer编码器

import torch
import torch.nn as nn
from torch.nn import MultiheadAttention

class TransformerEncoderLayer(nn.Module):
    def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1):
        super().__init__()
        self.self_attn = MultiheadAttention(d_model, nhead, dropout=dropout)
        self.linear1 = nn.Linear(d_model, dim_feedforward)
        self.dropout = nn.Dropout(dropout)
        self.linear2 = nn.Linear(dim_feedforward, d_model)
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)
        self.dropout1 = nn.Dropout(dropout)
        self.dropout2 = nn.Dropout(dropout)
        self.activation = nn.ReLU()

    def forward(self, src, pos_embed):
        q = k = src + pos_embed
        src2 = self.self_attn(q, k, value=src)[0]
        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

2.2 Transformer解码器

解码器的核心是对象查询(Object Queries)与编码器输出的交互:

class TransformerDecoderLayer(nn.Module):
    def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1):
        super().__init__()
        self.self_attn = MultiheadAttention(d_model, nhead, dropout=dropout)
        self.multihead_attn = MultiheadAttention(d_model, nhead, dropout=dropout)
        # 其余结构与编码器类似...
        
    def forward(self, tgt, memory, pos_embed, query_pos):
        q = k = tgt + query_pos
        tgt2 = self.self_attn(q, k, value=tgt)[0]
        tgt = tgt + self.dropout1(tgt2)
        tgt = self.norm1(tgt)
        tgt2 = self.multihead_attn(
            query=tgt + query_pos,
            key=memory + pos_embed,
            value=memory)[0]
        tgt = tgt + self.dropout2(tgt2)
        tgt = self.norm2(tgt)
        # FFN部分...
        return tgt

2.3 完整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
        self.query_embed = nn.Embedding(num_queries, transformer.d_model)
        # 预测头实现...
        
    def forward(self, images):
        # 1. 通过骨干网络提取特征
        features = self.backbone(images)
        
        # 2. 准备位置编码和查询向量
        pos_embed = self.position_encoding(features)
        query_embed = self.query_embed.weight.unsqueeze(1)
        
        # 3. 通过Transformer
        hs = self.transformer(features, None, query_embed, pos_embed)
        
        # 4. 预测输出
        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. 匈牙利匹配损失实现

DETR的核心创新是将目标检测视为集合预测问题,通过匈牙利算法实现预测与真值的最优匹配:

from scipy.optimize import linear_sum_assignment

def hungarian_matching(pred_boxes, pred_logits, targets):
    """
    pred_boxes: [batch_size, num_queries, 4]
    pred_logits: [batch_size, num_queries, num_classes]
    targets: list of dict with 'boxes' and 'labels'
    """
    batch_size = pred_logits.shape[0]
    indices = []
    
    for i in range(batch_size):
        # 计算分类损失矩阵
        cost_class = -pred_logits[i].softmax(-1)[:, targets[i]['labels']]
        
        # 计算边界框损失矩阵
        cost_bbox = torch.cdist(pred_boxes[i], targets[i]['boxes'], p=1)
        
        # 计算GIoU损失矩阵
        cost_giou = -generalized_box_iou(pred_boxes[i], targets[i]['boxes'])
        
        # 组合总成本矩阵
        C = 1.0 * cost_class + 1.0 * cost_bbox + 1.0 * cost_giou
        C = C.detach().cpu().numpy()
        
        # 匈牙利算法求解最优匹配
        row_ind, col_ind = linear_sum_assignment(C)
        indices.append((row_ind, col_ind))
    
    return indices

提示:实际实现中需要处理不同数量的目标(如填充到固定大小),并考虑背景类的特殊处理

匹配完成后,计算最终的损失函数:

def detr_loss(pred_logits, pred_boxes, targets, indices):
    # 分类损失(交叉熵)
    loss_ce = F.cross_entropy(pred_logits.transpose(1, 2), targets['labels'])
    
    # 边界框损失(L1)
    src_boxes = pred_boxes[indices]
    target_boxes = targets['boxes'][indices]
    loss_bbox = F.l1_loss(src_boxes, target_boxes, reduction='none')
    
    # GIoU损失
    loss_giou = 1 - torch.diag(generalized_box_iou(
        box_cxcywh_to_xyxy(src_boxes),
        box_cxcywh_to_xyxy(target_boxes)))
    
    # 组合损失
    loss = loss_ce + 5 * loss_bbox + 2 * loss_giou
    return loss

4. 训练技巧与优化

DETR原始版本需要500个epoch才能收敛,以下技巧可以显著加速训练:

4.1 学习率调度

def adjust_learning_rate(optimizer, epoch, args):
    """Decay the learning rate based on schedule"""
    lr = args.lr
    if epoch >= args.lr_drop:
        lr *= 0.1
    for param_group in optimizer.param_groups:
        param_group['lr'] = lr

optimizer = torch.optim.AdamW([
    {'params': [p for n, p in model.named_parameters() 
               if 'backbone' not in n], 'lr': 1e-4},
    {'params': [p for n, p in model.named_parameters() 
               if 'backbone' in n], 'lr': 1e-5},
])

4.2 梯度裁剪

torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.1)

4.3 关键参数设置

参数推荐值说明
batch_size16-64根据GPU内存调整
num_queries100默认足够覆盖COCO图像中的目标
lr_backbone1e-5骨干网络较小学习率
lr1e-4Transformer部分学习率
weight_decay1e-4AdamW优化器的权重衰减
epochs300原始DETR需要长训练周期

4.4 训练过程监控

建议监控以下指标:

  • 分类准确率
  • 边界框回归误差
  • GIoU值
  • 匈牙利匹配成本
  • 正样本匹配率

使用TensorBoard或WandB记录训练曲线,特别关注验证集指标是否同步提升。

5. 模型评估与推理

DETR的推理过程非常简洁,无需NMS后处理:

def inference(model, image, threshold=0.7):
    model.eval()
    with torch.no_grad():
        outputs = model(image)
    
    # 过滤低置信度预测
    probas = outputs['pred_logits'].softmax(-1)[0, :, :-1]
    keep = probas.max(-1).values > threshold
    
    # 转换框格式为[x0, y0, x1, y1]
    boxes = box_cxcywh_to_xyxy(outputs['pred_boxes'][0, keep])
    scores = probas[keep]
    labels = scores.argmax(-1)
    
    return boxes, labels, scores.max(-1).values

评估指标采用标准的COCO mAP:

from pycocotools.cocoeval import COCOeval

def evaluate(model, data_loader, device):
    model.eval()
    results = []
    for images, targets in data_loader:
        outputs = model(images.to(device))
        # 转换输出为COCO格式...
        results.extend(coco_format_outputs(outputs))
    
    coco_dt = data_loader.dataset.coco.loadRes(results)
    coco_eval = COCOeval(data_loader.dataset.coco, coco_dt, 'bbox')
    coco_eval.evaluate()
    coco_eval.accumulate()
    coco_eval.summarize()
    return coco_eval.stats[0]  # AP@[.5:.95]

6. 进阶优化方向

基础DETR实现后,可以考虑以下改进:

6.1 多尺度特征融合

class MultiScaleDETR(nn.Module):
    def __init__(self, backbone, transformer, num_classes, num_queries):
        super().__init__()
        # 使用FPN或类似结构融合多尺度特征
        self.fpn = FPN(backbone.out_channels, transformer.d_model)
        
    def forward(self, images):
        # 骨干网络输出多尺度特征
        features = self.backbone(images)  # 返回不同层级的特征图
        
        # 通过FPN融合特征
        fused_features = self.fpn(features)
        
        # 其余部分与标准DETR相同...

6.2 可变形注意力机制

from deformable_attention import DeformableAttention

class DeformableDecoderLayer(nn.Module):
    def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1):
        super().__init__()
        self.self_attn = MultiheadAttention(d_model, nhead, dropout=dropout)
        self.cross_attn = DeformableAttention(d_model, nhead)
        # 其余部分...

6.3 动态锚框改进

class DABDETR(nn.Module):
    def __init__(self, backbone, transformer, num_classes, num_queries):
        super().__init__()
        # 将查询向量显式表示为锚框
        self.query_embed = nn.Embedding(num_queries, 4)  # 直接学习(cx, cy, w, h)
        
    def forward(self, images):
        # 从查询向量解码初始锚框
        init_boxes = self.query_embed.weight.sigmoid()
        
        # 在解码过程中逐步优化锚框
        for layer in self.transformer.decoder.layers:
            # 将当前锚框信息注入注意力计算
            ...

在实现完整模型后,我发现匈牙利匹配的实现细节对模型性能影响最大。特别是如何处理不同数量的真实目标(包括背景类)以及损失权重平衡,需要多次实验调整。另一个关键点是位置编码的设计——在原始DETR中,2D正弦位置编码与查询向量的交互方式直接影响了模型对空间关系的理解能力。

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐