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数据集进行训练和评估。我们可以通过以下步骤准备数据:

  1. 下载COCO数据集(约18GB训练集+1GB验证集)
  2. 解压到指定目录,保持如下结构:
    coco/
    ├── annotations
    │   ├── instances_train2017.json
    │   └── instances_val2017.json
    ├── train2017
    └── val2017
    
  3. 实现自定义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的性能和效率:

  1. 多尺度训练:在训练时随机调整图像尺寸(480-800像素),增强模型对不同尺度目标的检测能力
  2. 梯度裁剪:防止梯度爆炸,稳定训练过程
    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.1)
    
  3. 混合精度训练:减少显存占用,加速训练
    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()
    
  4. 模型蒸馏:使用更大的DETR模型(如DETR-DC5)作为教师模型,指导基础模型训练

7. 常见问题与解决方案

在DETR的实际应用中,开发者常会遇到以下问题:

问题1:训练初期损失不下降

  • 原因:Transformer需要更长的预热期
  • 解决方案:使用更小的学习率(如1e-5)并延长训练周期

问题2:小物体检测效果差

  • 原因:高维特征丢失了小物体信息
  • 解决方案:
    • 使用更高分辨率的输入图像
    • 采用特征金字塔结构(如FPN)增强多尺度特征

问题3:训练显存不足

  • 解决方案:
    • 减小batch size(最低可设为1)
    • 使用梯度累积技术
    • 启用checkpointing减少内存占用

问题4:推理速度慢

  • 优化方案:
    • 量化模型权重(FP16或INT8)
    • 使用TensorRT加速
    • 剪枝去除冗余注意力头

8. 进阶改进方向

对于希望进一步提升DETR性能的开发者,可以考虑以下研究方向:

  1. Deformable DETR:引入可变形注意力机制,显著提升小物体检测性能
  2. Conditional DETR:改进query设计,加速模型收敛
  3. DAB-DETR:将query明确表示为动态锚框,提升定位精度
  4. DN-DETR:通过去噪训练解决二分图匹配不稳定的问题
  5. 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)
Logo

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

更多推荐