DETR目标检测实战:从零构建基于Transformer的端到端检测系统

在计算机视觉领域,目标检测一直是核心技术难题之一。传统方法如Faster R-CNN、YOLO等虽然效果显著,但都依赖于复杂的后处理流程(如非极大值抑制NMS)和人工设计的锚框(anchor)。2020年,Facebook AI团队提出的DETR(DEtection TRansformer)彻底改变了这一局面,首次将Transformer架构成功应用于目标检测任务,实现了真正的端到端检测。本文将带您从零开始,完整实现一个基于PyTorch的DETR模型,涵盖数据准备、模型构建、训练技巧和推理部署全流程。

1. 环境准备与数据加载

1.1 安装依赖库

首先确保您的Python环境为3.7+,并安装以下核心依赖:

pip install torch==1.10.0+cu113 torchvision==0.11.1+cu113 -f https://download.pytorch.org/whl/torch_stable.html
pip install pycocotools matplotlib scikit-image

1.2 COCO数据集准备

DETR默认使用COCO数据集进行训练和评估。下载并解压数据集到data/coco目录:

data/coco/
├── annotations/  # 存放instances_train2017.json等标注文件
├── train2017/    # 训练集图片
└── val2017/      # 验证集图片

创建自定义数据集加载器,继承PyTorch的Dataset类:

from torch.utils.data import Dataset
from pycocotools.coco import COCO

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):
        coco = self.coco
        img_id = self.ids[idx]
        ann_ids = coco.getAnnIds(imgIds=img_id)
        target = coco.loadAnns(ann_ids)
        
        path = coco.loadImgs(img_id)[0]['file_name']
        img = Image.open(os.path.join(self.img_folder, path)).convert('RGB')
        
        if self.transforms is not None:
            img, target = self.transforms(img, target)
            
        return img, target

1.3 数据增强策略

DETR对数据增强相对敏感,建议采用以下组合:

from torchvision.transforms import Compose, RandomResizedCrop, ToTensor, Normalize

def make_transforms(image_set):
    normalize = Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
    
    if image_set == 'train':
        return Compose([
            RandomHorizontalFlip(),
            RandomResizedCrop(800, scale=(0.8, 1.0)),
            ToTensor(),
            normalize
        ])
    else:
        return Compose([
            Resize(800),
            ToTensor(),
            normalize
        ])

2. DETR模型架构实现

2.1 主干网络(Backbone)

DETR使用标准的CNN(如ResNet50)提取图像特征:

import torch
from torch import nn
from torchvision.models import resnet50

class Backbone(nn.Module):
    def __init__(self, hidden_dim=256):
        super().__init__()
        self.backbone = resnet50(pretrained=True)
        self.conv = nn.Conv2d(2048, hidden_dim, 1)
        
    def forward(self, x):
        x = self.backbone.conv1(x)
        x = self.backbone.bn1(x)
        x = self.backbone.relu(x)
        x = self.backbone.maxpool(x)
        
        x = self.backbone.layer1(x)
        x = self.backbone.layer2(x)
        x = self.backbone.layer3(x)
        x = self.backbone.layer4(x)  # [batch, 2048, h/32, w/32]
        
        x = self.conv(x)  # [batch, hidden_dim, h/32, w/32]
        return x

2.2 Transformer编码器-解码器

实现Transformer的核心组件:

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)
        
    def forward(self, src, mask, query_embed, pos_embed):
        # src: [h*w, batch, d_model]
        # mask: [batch, h, w]
        # query_embed: [num_queries, d_model]
        # pos_embed: [h*w, batch, d_model]
        
        memory = self.encoder(src, src_key_padding_mask=mask, pos=pos_embed)
        hs = self.decoder(query_embed.unsqueeze(1).repeat(1, memory.shape[1], 1),
                         memory, memory_key_padding_mask=mask,
                         pos=pos_embed, query_pos=query_embed)
        return hs.transpose(1, 2)  # [num_queries, batch, d_model]

2.3 位置编码与预测头

DETR采用可学习的位置编码和简单的MLP预测头:

class PositionEmbeddingSine(nn.Module):
    def __init__(self, num_pos_feats=64, temperature=10000):
        super().__init__()
        self.num_pos_feats = num_pos_feats
        self.temperature = temperature
        
    def forward(self, x):
        not_mask = torch.ones_like(x[:,0,:,:])
        y_embed = not_mask.cumsum(1, dtype=torch.float32)
        x_embed = not_mask.cumsum(2, dtype=torch.float32)
        
        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

class DETR(nn.Module):
    def __init__(self, num_classes=91, num_queries=100):
        super().__init__()
        self.backbone = Backbone()
        self.transformer = Transformer()
        self.query_embed = nn.Embedding(num_queries, 256)
        self.pos_encoder = PositionEmbeddingSine()
        
        # 预测头
        self.class_embed = nn.Linear(256, num_classes + 1)
        self.bbox_embed = MLP(256, 256, 4, 3)
        
    def forward(self, x):
        features = self.backbone(x)  # [batch, 256, h, w]
        mask = torch.zeros_like(x[:,0,:,:], dtype=torch.bool)
        pos = self.pos_encoder(features)
        
        # 展平特征图
        features = features.flatten(2).permute(2, 0, 1)  # [h*w, batch, 256]
        pos = pos.flatten(2).permute(2, 0, 1)
        
        hs = self.transformer(features, mask, self.query_embed.weight, pos)
        
        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. 匈牙利匹配与损失函数

3.1 二分图匹配实现

DETR使用匈牙利算法进行预测框与真实框的最优匹配:

from scipy.optimize import linear_sum_assignment

def hungarian_matcher(outputs, targets):
    bs, num_queries = outputs["pred_logits"].shape[:2]
    
    # 展平batch维度
    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[i*num_queries:(i+1)*num_queries, tgt_ids]
        
        # 计算框损失 (L1 + GIoU)
        cost_bbox = torch.cdist(out_bbox[i*num_queries:(i+1)*num_queries], tgt_bbox, p=1)
        cost_giou = -generalized_box_iou(box_cxcywh_to_xyxy(out_bbox[i*num_queries:(i+1)*num_queries]),
                                        box_cxcywh_to_xyxy(tgt_bbox))
        
        # 总成本矩阵
        C = 1 * cost_class + 5 * cost_bbox + 2 * cost_giou
        C = C.reshape(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]

3.2 损失函数设计

DETR的损失函数包含分类损失和框回归损失:

def detr_loss(outputs, targets, matcher):
    # 获取匹配结果
    indices = matcher(outputs, targets)
    
    # 分类损失 (交叉熵)
    src_logits = outputs['pred_logits']
    idx = self._get_src_permutation_idx(indices)
    target_classes = torch.full(src_logits.shape[:2], 91, 
                               dtype=torch.int64, device=src_logits.device)
    for i, (_, J) in enumerate(indices):
        target_classes[i, J] = targets[i]['labels']
    loss_ce = F.cross_entropy(src_logits.transpose(1, 2), target_classes)
    
    # 框回归损失 (L1 + GIoU)
    src_boxes = outputs['pred_boxes']
    target_boxes = torch.cat([t['boxes'][i] for t, (_, i) in zip(targets, indices)], dim=0)
    loss_bbox = F.l1_loss(src_boxes[idx], target_boxes, reduction='none')
    loss_giou = 1 - torch.diag(generalized_box_iou(
        box_cxcywh_to_xyxy(src_boxes[idx]),
        box_cxcywh_to_xyxy(target_boxes)))
    
    # 总损失
    losses = {
        'loss_ce': loss_ce,
        'loss_bbox': loss_bbox.sum() / num_boxes,
        'loss_giou': loss_giou.sum() / num_boxes
    }
    return losses

4. 训练技巧与优化策略

4.1 学习率调度

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},
    ]
    optimizer = torch.optim.AdamW(param_dicts, lr=lr, weight_decay=weight_decay)
    
    # 学习率预热
    lr_scheduler = torch.optim.lr_scheduler.SequentialLR(
        optimizer,
        [
            torch.optim.lr_scheduler.LinearLR(optimizer, 0.0001, 1.0, 100),
            torch.optim.lr_scheduler.MultiStepLR(optimizer, [40], gamma=0.1)
        ],
        [100]
    )
    return optimizer, lr_scheduler

4.2 梯度裁剪与混合精度训练

为稳定训练过程,建议启用梯度裁剪和混合精度:

scaler = torch.cuda.amp.GradScaler()

for epoch in range(epochs):
    for images, targets in dataloader:
        optimizer.zero_grad()
        
        with torch.cuda.amp.autocast():
            outputs = model(images)
            loss_dict = criterion(outputs, targets)
            losses = sum(loss_dict.values())
            
        scaler.scale(losses).backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), 0.1)
        scaler.step(optimizer)
        scaler.update()
        
        lr_scheduler.step()

4.3 关键训练参数配置

参数推荐值说明
batch_size16-32根据GPU显存调整
base_lr1e-4主干网络使用0.1倍
weight_decay1e-4防止过拟合
epochs50-300小数据集需更多epoch
warmup_epochs100学习率线性预热
dropout0.1Transformer层dropout率
num_queries100默认预测框数量

5. 模型推理与部署

5.1 预测后处理

DETR推理过程无需NMS,直接过滤低置信度预测:

def postprocess(outputs, threshold=0.7):
    probas = outputs['pred_logits'].softmax(-1)[:, :, :-1]  # 去除背景类
    keep = probas.max(-1).values > threshold
    
    results = []
    for p, b, k in zip(probas, outputs['pred_boxes'], keep):
        if k.sum() == 0:
            results.append({'scores': [], 'labels': [], 'boxes': []})
            continue
            
        scores = p[k]
        labels = scores.argmax(-1)
        boxes = b[k]
        results.append({
            'scores': scores.max(-1).values,
            'labels': labels,
            'boxes': boxes
        })
    return results

5.2 ONNX导出与优化

将模型导出为ONNX格式便于部署:

dummy_input = torch.randn(1, 3, 800, 800)
torch.onnx.export(
    model, 
    dummy_input,
    "detr.onnx",
    input_names=["input"],
    output_names=["logits", "boxes"],
    dynamic_axes={
        "input": {0: "batch"},
        "logits": {0: "batch"},
        "boxes": {0: "batch"}
    },
    opset_version=12
)

5.3 TensorRT加速

使用TensorRT进一步优化推理速度:

# 转换ONNX到TensorRT引擎
trtexec --onnx=detr.onnx --saveEngine=detr.engine \
        --fp16 --workspace=4096 --minShapes=input:1x3x800x800 \
        --optShapes=input:4x3x800x800 --maxShapes=input:8x3x800x800

# Python加载引擎
with open("detr.engine", "rb") as f:
    runtime = trt.Runtime(trt.Logger(trt.Logger.WARNING))
    engine = runtime.deserialize_cuda_engine(f.read())

6. 常见问题与解决方案

6.1 训练收敛慢问题

现象:模型在初期训练阶段损失下降缓慢,AP提升不明显。

解决方案:

  1. 确保正确实现了学习率预热(前100个step从1e-6线性增加到1e-4)
  2. 检查主干网络是否冻结了BatchNorm层的统计量
  3. 增加梯度裁剪(norm=0.1)防止初期梯度爆炸
  4. 验证匈牙利匹配实现是否正确,特别是GIoU计算部分

6.2 小目标检测效果差

现象:模型对大目标检测效果良好,但小目标漏检率高。

优化策略:

  1. 在主干网络后添加FPN结构,融合多尺度特征
  2. 减小输入图像的下采样倍数(如从32倍降到16倍)
  3. 增加小目标在训练数据中的采样权重
  4. 使用Deformable DETR等改进版本,专门优化小目标检测

6.3 显存不足问题

现象:训练时出现CUDA out of memory错误。

应对措施:

  1. 减小batch size(最低可到4,配合梯度累积)
  2. 使用混合精度训练(AMP)
  3. 启用checkpointing技术减少中间缓存
  4. 简化Transformer结构(如减少层数或隐藏维度)
# 梯度累积示例
accumulation_steps = 4
optimizer.zero_grad()

for i, (images, targets) in enumerate(dataloader):
    outputs = model(images)
    loss = criterion(outputs, targets)
    loss = loss / accumulation_steps
    loss.backward()
    
    if (i+1) % accumulation_steps == 0:
        optimizer.step()
        optimizer.zero_grad()

7. 进阶优化方向

7.1 模型轻量化

对于移动端部署,可考虑以下优化:

  1. 替换主干网络:使用MobileNetV3或EfficientNet-Lite替代ResNet
  2. 知识蒸馏:用大模型指导小模型训练
  3. 量化感知训练:直接训练8整型模型
  4. 剪枝:移除冗余的注意力头或FFN层

7.2 多任务扩展

DETR框架可轻松扩展其他视觉任务:

  1. 实例分割:在解码器输出添加掩码头
  2. 姿态估计:预测关键点坐标
  3. 全景分割:统一语义分割和实例分割
  4. 视频目标检测:加入时序建模模块

7.3 最新改进版本

改进版本核心创新优势
Deformable DETR可变形注意力机制提升小目标检测,加速收敛
DAB-DETR动态锚框查询更稳定的训练过程
DN-DETR去噪训练策略解决二分图匹配不稳定性
DINO-DETR对比学习预训练更高检测精度

实际项目中,根据具体需求选择合适的DETR变体往往能获得更好的效果。例如,对于实时性要求高的场景,可选用Conditional DETR;对于小目标密集场景,Deformable DETR是更好的选择。

Logo

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

更多推荐