1. 从图像分类到目标检测:ViT的跨界之旅

大家好,我是老张,在AI和硬件这块摸爬滚打了十几年。今天咱们不聊那些老生常谈的CNN,来聊聊一个真正改变了游戏规则的家伙——Vision Transformer (ViT)。我知道,一听到Transformer,很多人第一反应是“那不是搞自然语言处理的吗?”。没错,但正是这种跨界思维,让ViT在图像领域掀起了一场革命。你可能已经在各种论文里见过它的身影,但今天,我们不只讲原理,我要带你亲手把它“塞”进目标检测的任务里,看看这个原本为分类而生的模型,如何摇身一变成为检测任务的利器。

ViT的核心思想其实非常直观,甚至有点“暴力美学”的味道。它不像CNN那样,通过卷积核一层层地、局部地感受图像。相反,它把一张图片,比如224x224的,直接切成一个个16x16的小方块(我们称之为Patch)。每个小方块被拉平成向量,这就好比把一篇文章拆分成一个个单词。然后,这些“视觉单词”被送入标准的Transformer编码器里,让模型通过自注意力机制(Self-Attention)自己去学习这些“单词”之间的关系。这种全局建模的能力,是CNN那种局部感受野难以比拟的。想象一下,在目标检测中,要判断一个物体,往往需要结合它周围甚至很远处的上下文信息,ViT这种全局“视野”天生就具有优势。

那么,一个很自然的问题就来了:ViT在ImageNet分类上大放异彩,我们怎么把它用到目标检测上呢?目标检测不仅要识别“是什么”,还要定位出“在哪里”。这就像从“这张图里有猫”升级到“猫在图里的左上角,大小是100x150像素”。ViT本身输出的是一个全局的、用于分类的特征表示,我们得想办法从这个表示中,解码出一个个边界框(Bounding Box)和对应的类别。这中间需要一些巧妙的“桥梁”设计,也是我们今天实战和源码解析的重点。别担心,我会把每一步都掰开揉碎了讲,保证你能跟上。

2. 实战准备:环境、数据与核心库

工欲善其事,必先利其器。在开始敲代码之前,咱们先把环境搭好。我强烈建议使用Anaconda来管理环境,这样可以避免各种依赖冲突的“玄学”问题。

2.1 创建并配置Python环境

首先,我们创建一个新的Python环境,我习惯用Python 3.8,比较稳定。

conda create -n vit-detection python=3.8 -y
conda activate vit-detection

接下来安装核心的深度学习框架PyTorch。去PyTorch官网根据你的CUDA版本选择安装命令。如果你没有GPU,就用CPU版本,不过训练会慢很多。我这里有CUDA 11.3,所以我的命令是这样的:

pip install torch==1.12.1+cu113 torchvision==0.13.1+cu113 --extra-index-url https://download.pytorch.org/whl/cu113

然后,安装一些我们必需的辅助库。timm库是“神”库,里面集成了大量预训练的视觉模型,包括我们要用的ViT。albumentations是一个强大的图像增强库,比torchvision自带的更灵活。opencv-pythonpycocotools是处理数据和评估模型必备的。

pip install timm albumentations opencv-python pycocotools matplotlib pandas scikit-learn

2.2 理解目标检测的数据格式

目标检测的数据标注比分类复杂。最常用的格式是COCO格式。一个标注文件(通常是JSON)里,会包含图片信息、类别信息,以及最重要的——标注列表。每个标注都包含一个边界框([x_min, y_min, width, height])和对应的类别ID。

为了让大家有直观感受,我准备了一个超小型的数据集示例结构,并写一个简单的脚本来加载和可视化它。我们假设数据放在 ./data/coco 目录下。

import json
import cv2
import matplotlib.pyplot as plt
import matplotlib.patches as patches
from pathlib import Path

# 加载COCO格式的标注文件
data_dir = Path('./data/coco')
ann_file = data_dir / 'annotations' / 'instances_train2017.json'

with open(ann_file, 'r') as f:
    coco_data = json.load(f)

# 建立图片ID到文件名的映射
img_id_to_info = {img['id']: img for img in coco_data['images']}
# 建立类别ID到类别名的映射
cat_id_to_name = {cat['id']: cat['name'] for cat in coco_data['categories']}

# 获取第一张图片的ID和其对应的所有标注
first_img_id = coco_data['images'][0]['id']
annotations_for_first_img = [ann for ann in coco_data['annotations'] if ann['image_id'] == first_img_id]

# 读取图片并绘制标注
img_info = img_id_to_info[first_img_id]
img_path = data_dir / 'train2017' / img_info['file_name']
img = cv2.imread(str(img_path))
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) # OpenCV默认是BGR,转成RGB

fig, ax = plt.subplots(1, figsize=(12, 8))
ax.imshow(img)

for ann in annotations_for_first_img:
    # COCO格式的bbox是 [x, y, width, height]
    x, y, w, h = ann['bbox']
    # 创建矩形框
    rect = patches.Rectangle((x, y), w, h, linewidth=2, edgecolor='r', facecolor='none')
    ax.add_patch(rect)
    # 添加类别标签
    cat_name = cat_id_to_name[ann['category_id']]
    ax.text(x, y-5, cat_name, color='red', fontsize=12, weight='bold')

plt.axis('off')
plt.show()

运行这段代码,你就能看到一张带标注框的图片。理解数据是第一步,也是至关重要的一步。很多新手在模型跑不通时,问题往往就出在数据预处理环节。

3. ViT目标检测的核心架构:DETR与它的朋友们

直接把ViT拿来做检测是行不通的,我们需要一个“检测头”。目前最主流、也最能与ViT哲学契合的方法,是Facebook AI Research提出的 DETR(DEtection TRansformer)。它彻底抛弃了传统检测器中锚框(Anchor)、非极大值抑制(NMS)这些复杂的手工设计,用一套极其简洁的端到端框架实现了目标检测。

3.1 DETR框架精讲

DETR的流程可以概括为三步:特征提取 -> Transformer编解码 -> 预测头

  1. 特征提取:用CNN(如ResNet)或ViT作为骨干网络(Backbone),从输入图像中提取一个2D的特征图。如果我们用ViT,需要一点小改动,这个后面细说。
  2. Transformer编解码:这是DETR的灵魂。编码器对骨干网络提取的特征进行增强。解码器则接收一组固定数量的“对象查询”(Object Queries,比如100个),这些查询是可学习的参数,可以理解为模型学习到的、用于询问图像中“哪里有物体”的探针。解码器让这些对象查询与编码器输出的特征进行交互,最终每个查询输出一个包含物体信息的嵌入向量。
  3. 预测头:一个简单的全连接网络,接在每个对象查询的输出后面。它预测两个东西:类别(包括一个“无物体”类)和边界框(中心点坐标、宽高)。

DETR的损失函数也很有特色,它使用二分图匹配损失(Hungarian Loss)。简单说,就是把模型预测的100个框和图片里真实的GT框进行一对一匹配,找到代价最小的配对方式,然后只对匹配上的预测框计算分类和回归损失。这保证了预测框和真实框之间有序的、一对一的对应关系,从而避免了NMS。

3.2 将ViT作为DETR的骨干网络

原始的DETR论文用的是ResNet-50作为骨干。但ViT的输出是一个序列(比如197个Token),而不是一个2D特征图。怎么适配呢?一个常见的方法是使用ViT中间层的特征

ViT的blocks输出的是一个[Batch, Num_Tokens, Embed_Dim]的序列。其中第一个Token是[CLS],用于分类。剩下的196个Token对应原始的14x14个图像块。我们可以把这196个Token reshape回 [Batch, Embed_Dim, 14, 14] 的2D格式,这样它就变成了一个特征图,可以喂给DETR的Transformer编码器了!

下面我们来看一个简化版的代码,展示如何从timm中加载一个预训练的ViT,并提取中间层特征。

import torch
import torch.nn as nn
import timm

class ViTBackboneForDETR(nn.Module):
    def __init__(self, model_name='vit_base_patch16_224', pretrained=True, feature_layer=-2):
        super().__init__()
        # 加载预训练的ViT模型
        self.vit = timm.create_model(model_name, pretrained=pretrained, num_classes=0) # num_classes=0去掉分类头
        self.feature_layer = feature_layer # 指定取哪一层的输出,-2通常取倒数第二层

        # 获取一些模型配置信息
        self.patch_size = self.vit.patch_embed.patch_size[0]
        self.embed_dim = self.vit.embed_dim
        self.num_patches = self.vit.patch_embed.num_patches

    def forward(self, x):
        # x: [B, C, H, W]
        B = x.shape[0]
        # 通过ViT的patch embedding和pos embedding
        x = self.vit.patch_embed(x) # [B, num_patches, embed_dim]
        cls_token = self.vit.cls_token.expand(B, -1, -1)
        x = torch.cat((cls_token, x), dim=1)
        x = x + self.vit.pos_embed

        # 通过指定的Transformer blocks
        # 注意:timm的ViT通常把blocks放在一个ModuleList里
        for i, blk in enumerate(self.vit.blocks):
            x = blk(x)
            if i == self.feature_layer:
                features = x # 保存指定层的输出

        # features的形状是 [B, num_patches+1, embed_dim]
        # 我们需要去掉[CLS] token,并把序列reshape成2D特征图
        # 假设输入是224x224,patch是16,那么num_patches = 14*14=196
        patch_features = features[:, 1:, :] # 去掉第一个[CLS] token, [B, 196, 768]
        # Reshape成 [B, embed_dim, height, width]
        feature_map = patch_features.permute(0, 2, 1).reshape(B, self.embed_dim, 14, 14)

        return feature_map # 输出一个2D特征图,形状如 [B, 768, 14, 14]

# 测试一下
if __name__ == '__main__':
    model = ViTBackboneForDETR('vit_base_patch16_224', pretrained=True)
    dummy_input = torch.randn(2, 3, 224, 224)
    output = model(dummy_input)
    print(f"输入形状: {dummy_input.shape}")
    print(f"ViT骨干输出特征图形状: {output.shape}") # 应该是 torch.Size([2, 768, 14, 14])

这段代码是关键。我们创建了一个包装类,它加载预训练的ViT,并截取倒数第二层Block的输出(经验上这一层的特征兼顾了语义和位置信息)。然后,我们丢弃[CLS] token,将剩余的196个token(对应14x14网格)重新排列成 [Batch, 768, 14, 14] 的特征图。这个特征图就可以无缝输入到DETR后续的模块中了。

4. 手把手实现:构建ViT-DETR检测模型

理解了原理,我们现在来动手搭建一个完整的、简化版的ViT-DETR模型。我们会尽量保持代码清晰,并加上详细的注释。

4.1 定义Transformer编解码器

PyTorch已经提供了nn.Transformer模块,但为了更清晰地理解DETR的设计,我们这里参考其实现,自己构建一个简化的版本。重点是理解对象查询(Object Queries)和交叉注意力(Cross-Attention)的作用。

import copy
import math
import torch.nn.functional as F

def _get_clones(module, N):
    """克隆N个相同的模块"""
    return nn.ModuleList([copy.deepcopy(module) for _ in range(N)])

class TransformerEncoderLayer(nn.Module):
    """简化版的Transformer编码器层"""
    def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1):
        super().__init__()
        self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True)
        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 = F.relu

    def forward(self, src):
        # src: [Batch, Seq_Len, D_Model]
        # 自注意力
        src2 = self.self_attn(src, src, 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

class TransformerDecoderLayer(nn.Module):
    """简化版的Transformer解码器层,包含自注意力和交叉注意力"""
    def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1):
        super().__init__()
        self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True)
        self.multihead_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True) # 交叉注意力
        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.norm3 = nn.LayerNorm(d_model)
        self.dropout1 = nn.Dropout(dropout)
        self.dropout2 = nn.Dropout(dropout)
        self.dropout3 = nn.Dropout(dropout)
        self.activation = F.relu

    def forward(self, tgt, memory):
        # tgt: 对象查询 [Batch, Num_Queries, D_Model]
        # memory: 编码器输出的特征 [Batch, Seq_Len, D_Model]
        # 自注意力
        tgt2 = self.self_attn(tgt, tgt, tgt)[0]
        tgt = tgt + self.dropout1(tgt2)
        tgt = self.norm1(tgt)
        # 交叉注意力:查询是tgt,键和值是memory
        tgt2 = self.multihead_attn(tgt, memory, memory)[0]
        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

4.2 组装完整的ViT-DETR模型

现在,我们把ViT骨干、Transformer编解码器和预测头组合起来。

class ViTDETR(nn.Module):
    def __init__(self, backbone, hidden_dim=256, num_queries=100, num_classes=91, num_encoder_layers=6, num_decoder_layers=6):
        super().__init__()
        self.backbone = backbone # 我们之前定义的ViTBackboneForDETR
        self.num_queries = num_queries

        # 1. 将骨干网络输出的特征图投影到Transformer的隐藏维度
        backbone_output_channels = backbone.embed_dim # 例如768
        self.input_proj = nn.Conv2d(backbone_output_channels, hidden_dim, kernel_size=1)

        # 2. 位置编码:为特征图上的每个位置生成一个编码
        # DETR中使用的是正弦位置编码,但这里我们简化为可学习的位置编码
        self.position_embedding = nn.Parameter(torch.zeros(1, hidden_dim, 14, 14)) # 假设特征图是14x14

        # 3. Transformer编码器
        encoder_layer = TransformerEncoderLayer(d_model=hidden_dim, nhead=8)
        self.encoder = nn.ModuleList([copy.deepcopy(encoder_layer) for _ in range(num_encoder_layers)])

        # 4. 对象查询:可学习的参数
        self.query_embed = nn.Embedding(num_queries, hidden_dim)

        # 5. Transformer解码器
        decoder_layer = TransformerDecoderLayer(d_model=hidden_dim, nhead=8)
        self.decoder = nn.ModuleList([copy.deepcopy(decoder_layer) for _ in range(num_decoder_layers)])

        # 6. 预测头
        # 分类头:预测每个查询的类别(包括“无物体”)
        self.class_embed = nn.Linear(hidden_dim, num_classes + 1) # +1 for “no object”
        # 边界框回归头:预测中心点坐标、宽高(归一化到0-1)
        self.bbox_embed = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, 4) # (cx, cy, w, h)
        )

    def forward(self, images):
        # 提取特征
        features = self.backbone(images) # [B, C, H, W] e.g., [B, 768, 14, 14]
        # 投影到隐藏维度
        features_proj = self.input_proj(features) # [B, hidden_dim, H, W]
        # 加上位置编码
        features_proj = features_proj + self.position_embedding

        # 将2D特征图展平为序列,以适配Transformer
        B, C, H, W = features_proj.shape
        features_flat = features_proj.flatten(2).permute(0, 2, 1) # [B, H*W, hidden_dim]

        # 通过编码器
        memory = features_flat
        for enc_layer in self.encoder:
            memory = enc_layer(memory)

        # 准备对象查询
        query_embed = self.query_embed.weight.unsqueeze(0).repeat(B, 1, 1) # [B, num_queries, hidden_dim]
        tgt = torch.zeros_like(query_embed) # 解码器的初始输入是零向量

        # 通过解码器
        for dec_layer in self.decoder:
            tgt = dec_layer(tgt, memory)

        # 预测
        outputs_class = self.class_embed(tgt) # [B, num_queries, num_classes+1]
        outputs_coord = self.bbox_embed(tgt).sigmoid() # 用sigmoid将坐标约束在(0,1)

        return {'pred_logits': outputs_class, 'pred_boxes': outputs_coord}

# 实例化模型
backbone = ViTBackboneForDETR('vit_base_patch16_224', pretrained=True)
model = ViTDETR(backbone, hidden_dim=256, num_queries=100, num_classes=80) # COCO有80类
print(model)

这个模型结构已经具备了ViT-DETR的核心要素。query_embed是精髓,这100个可学习的向量,在训练过程中会逐渐学会去“询问”图像中不同位置、不同尺度的物体信息。预测时,每个查询输出一个分类概率和一个边界框。

5. 训练技巧与源码调试实战

模型搭好了,但直接训练很可能效果不佳甚至不收敛。这里我分享几个在实际项目中踩过坑才总结出来的关键技巧。

5.1 损失函数:匈牙利匹配的实现

DETR的损失函数实现是难点。我们需要计算预测集合和真实标注集合之间的最优二分图匹配。这里给出一个高度简化但有助于理解的核心匹配逻辑。

def hungarian_matcher(pred_logits, pred_boxes, targets):
    """
    简化的匈牙利匹配器。
    pred_logits: [B, N, C] 预测的分类分数
    pred_boxes: [B, N, 4] 预测的边界框 (cx, cy, w, h)
    targets: list of dict,每个dict包含'labels'和'boxes'
    """
    bs, num_queries = pred_logits.shape[:2]
    indices = []

    for i in range(bs):
        # 获取第i张图的真实标注
        tgt_labels = targets[i]['labels'] # [num_gt]
        tgt_boxes = targets[i]['boxes']   # [num_gt, 4]

        # 计算分类成本:负的概率(我们希望匹配上的预测概率高)
        cost_class = -pred_logits[i, :, tgt_labels].softmax(dim=0) # [N, num_gt]

        # 计算边界框L1成本(简化版,实际DETR用L1+GIoU)
        # 将预测框和真实框都扩展到相同维度以广播计算
        pred_boxes_i = pred_boxes[i].unsqueeze(1) # [N, 1, 4]
        tgt_boxes_i = tgt_boxes.unsqueeze(0)      # [1, num_gt, 4]
        cost_bbox = torch.cdist(pred_boxes_i, tgt_boxes_i, p=1).squeeze(1) # [N, num_gt]

        # 总成本 = 分类成本 + 边界框成本
        C = cost_class + cost_bbox

        # 为每个真实框找到成本最低的预测框(这里简化了,实际是全局最优匹配)
        # 使用一个简单的贪心匹配作为示例
        matched_row = []
        matched_col = []
        C_np = C.cpu().detach().numpy()
        num_gt = len(tgt_labels)

        for gt_idx in range(num_gt):
            # 找到对于当前gt_idx,成本最小的预测框
            min_cost_idx = C_np[:, gt_idx].argmin()
            matched_row.append(min_cost_idx)
            matched_col.append(gt_idx)

        indices.append((torch.as_tensor(matched_row), torch.as_tensor(matched_col)))

    return indices

在实际的DETR源码中,使用的是scipy.optimize.linear_sum_assignment函数来求解这个最优分配问题。匹配完成后,只有匹配上的预测框才会参与损失计算,未匹配的预测框则被归类为“无物体”。

5.2 学习率策略与梯度裁剪

ViT和Transformer模型对学习率非常敏感。我常用的策略是线性预热(Warmup) 后接余弦退火(Cosine Annealing)

  • Warmup:在训练初期(如前1000步),学习率从0线性增长到初始学习率。这有助于模型在训练初期稳定下来。
  • Cosine Annealing:之后,学习率按照余弦函数从初始学习率衰减到接近0。

此外,Transformer的梯度可能在某些层变得很大,导致训练不稳定。使用梯度裁剪(Gradient Clipping) 可以有效地缓解这个问题。

# 优化器设置示例
import torch.optim as optim
from torch.optim.lr_scheduler import LambdaLR, CosineAnnealingLR

def get_optimizer_and_scheduler(model, base_lr=1e-4, warmup_steps=1000, total_steps=90000):
    # 为不同部分设置不同的学习率是常见做法
    param_dicts = [
        {"params": [p for n, p in model.named_parameters() if "backbone" in n and p.requires_grad], "lr": base_lr * 0.1}, # 骨干网络用较小的学习率
        {"params": [p for n, p in model.named_parameters() if "backbone" not in n and p.requires_grad], "lr": base_lr},
    ]
    optimizer = optim.AdamW(param_dicts, weight_decay=1e-4)

    # 定义warmup函数
    def lr_lambda(current_step):
        if current_step < warmup_steps:
            return float(current_step) / float(max(1, warmup_steps))
        # 之后使用余弦退火
        progress = float(current_step - warmup_steps) / float(max(1, total_steps - warmup_steps))
        return max(0.0, 0.5 * (1.0 + math.cos(math.pi * progress)))

    scheduler = LambdaLR(optimizer, lr_lambda)
    return optimizer, scheduler

# 在训练循环中应用梯度裁剪
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.1) # 梯度裁剪
optimizer.step()
scheduler.step()

5.3 数据增强与调试技巧

对于目标检测,恰当的数据增强能极大提升模型泛化能力。除了常规的随机裁剪、翻转、颜色抖动,大规模抖动(Large Scale Jittering, LSJ) 是DETR论文中强调的有效技巧。它随机将图片缩放到一个很大范围的尺寸(比如0.1到2.0倍),然后进行随机裁剪。

在调试时,最怕的就是模型不输出损失或者输出NaN。我的经验是:

  1. 前向传播检查:在训练开始前,用一批数据跑一次前向传播,确保输出形状正确,没有NaN或Inf。
  2. 损失值监控:在训练初期,密切监控各个损失分量(分类损失、L1框损失、GIoU损失)。如果某个损失突然飙升或变为NaN,立即停止训练,检查数据或模型。
  3. 可视化中间特征:有时可以可视化Transformer编码器输出的特征图,或者解码器注意力图,看看模型是否在关注正确的区域。这能帮你直观理解模型在学什么。
# 一个简单的前向检查
model.eval()
with torch.no_grad():
    dummy_img = torch.randn(1, 3, 224, 224)
    outputs = model(dummy_img)
    print(f"预测logits形状: {outputs['pred_logits'].shape}") # 应为 [1, 100, num_classes+1]
    print(f"预测boxes形状: {outputs['pred_boxes'].shape}")   # 应为 [1, 100, 4]
    print(f"预测boxes值范围: {outputs['pred_boxes'].min():.3f}, {outputs['pred_boxes'].max():.3f}") # 应在0-1之间

6. 模型评估与结果分析

模型训练完成后,我们需要客观地评估它的性能。目标检测最常用的评估指标是平均精度(Average Precision, AP),尤其是在COCO数据集上定义的AP@[0.5:0.95](即IoU阈值从0.5到0.95,步长0.05,计算的平均AP)。

6.1 使用COCO评估工具

pycocotools库提供了官方的评估接口。我们需要将模型的预测结果转换成COCO评估格式。

from pycocotools.coco import COCO
from pycocotools.cocoeval import COCOeval
import numpy as np

def evaluate_coco(model, data_loader, device, threshold=0.7):
    model.eval()
    results = []
    image_ids = []

    with torch.no_grad():
        for images, targets in data_loader:
            images = images.to(device)
            outputs = model(images)

            pred_logits = outputs['pred_logits'].cpu().softmax(-1) # [B, N, C]
            pred_boxes = outputs['pred_boxes'].cpu() # [B, N, 4]

            for i in range(len(images)):
                # 获取当前图片的预测
                scores = pred_logits[i, :, :-1] # 去掉“无物体”类
                boxes = pred_boxes[i]
                # 过滤低置信度的预测
                max_scores, labels = scores.max(dim=1)
                keep = max_scores > threshold

                filtered_scores = max_scores[keep]
                filtered_labels = labels[keep]
                filtered_boxes = boxes[keep]

                # 将归一化的(cx, cy, w, h)转换为COCO格式的(x_min, y_min, w, h)
                img_h, img_w = targets[i]['orig_size']
                scale_fct = torch.stack([img_w, img_h, img_w, img_h], dim=0)
                # 注意:pred_boxes是中心点坐标,需要转换
                cx, cy, w, h = filtered_boxes.unbind(1)
                x_min = (cx - 0.5 * w) * img_w
                y_min = (cy - 0.5 * h) * img_h
                w = w * img_w
                h = h * img_h
                coco_boxes = torch.stack([x_min, y_min, w, h], dim=1).tolist()

                # 构建结果列表
                image_id = targets[i]['image_id'].item()
                for box, score, label in zip(coco_boxes, filtered_scores, filtered_labels):
                    results.append({
                        'image_id': image_id,
                        'category_id': label.item() + 1, # COCO类别ID从1开始
                        'bbox': [round(coord, 3) for coord in box],
                        'score': round(score.item(), 5)
                    })
                image_ids.append(image_id)

    # 加载标注文件
    coco_gt = COCO('path/to/annotations/instances_val2017.json')
    # 将预测结果加载为COCO格式
    coco_dt = coco_gt.loadRes(results)

    # 执行评估
    coco_eval = COCOeval(coco_gt, coco_dt, 'bbox')
    coco_eval.params.imgIds = image_ids
    coco_eval.evaluate()
    coco_eval.accumulate()
    coco_eval.summarize()

    # 打印主要指标
    stats = coco_eval.stats
    print(f"AP @[IoU=0.50:0.95]: {stats[0]:.3f}")
    print(f"AP @[IoU=0.50]: {stats[1]:.3f}")
    print(f"AP @[IoU=0.75]: {stats[2]:.3f}")
    print(f"AP for small objects: {stats[3]:.3f}")
    print(f"AP for medium objects: {stats[4]:.3f}")
    print(f"AP for large objects: {stats[5]:.3f}")

6.2 结果可视化与问题诊断

评估指标是冷冰冰的数字,可视化能给你更热的反馈。把模型预测的框和真实框画在一起,你能立刻看出问题所在:是框不准?还是漏检?或者是把背景误检成了物体?

def visualize_predictions(image, target, prediction, score_threshold=0.5):
    """
    image: 原始图片 (H, W, 3)
    target: 真实标注,包含'boxes', 'labels'
    prediction: 模型输出,包含'pred_boxes', 'pred_logits'
    """
    import matplotlib.patches as patches
    fig, ax = plt.subplots(1, 2, figsize=(20, 10))

    # 绘制真实框
    ax[0].imshow(image)
    ax[0].set_title('Ground Truth')
    for box, label in zip(target['boxes'], target['labels']):
        x, y, w, h = box
        rect = patches.Rectangle((x, y), w, h, linewidth=2, edgecolor='g', facecolor='none')
        ax[0].add_patch(rect)
        ax[0].text(x, y-5, f'GT:{label}', color='green', fontsize=10, weight='bold')

    # 绘制预测框
    ax[1].imshow(image)
    ax[1].set_title('Predictions')
    pred_boxes = prediction['pred_boxes'][0] # 取batch第一个
    pred_logits = prediction['pred_logits'][0].softmax(-1)
    scores, labels = pred_logits[:, :-1].max(dim=1) # 去掉“无物体”类

    for box, score, label in zip(pred_boxes, scores, labels):
        if score < score_threshold:
            continue
        cx, cy, w_n, h_n = box
        # 转换到像素坐标
        H, W, _ = image.shape
        x = (cx - 0.5 * w_n) * W
        y = (cy - 0.5 * h_n) * H
        w = w_n * W
        h = h_n * H
        rect = patches.Rectangle((x, y), w, h, linewidth=2, edgecolor='r', facecolor='none')
        ax[1].add_patch(rect)
        ax[1].text(x, y-5, f'Pred:{label}({score:.2f})', color='red', fontsize=10, weight='bold')

    plt.show()

通过可视化,你可能会发现ViT-DETR初期的一些典型问题,比如对小物体检测效果差(因为ViT的patch划分可能丢失小物体细节),或者对密集物体的区分能力不足。这为你后续的模型改进(比如使用多尺度特征、更精细的patch划分)指明了方向。

7. 进阶探索与性能优化

当你跑通了一个基础版本的ViT-DETR后,可以尝试一些进阶的优化策略来提升性能。

7.1 多尺度特征与特征金字塔

原始的ViT输出单一尺度的特征图(如14x14),这对于检测不同大小的物体是不利的。一个改进思路是借鉴FPN(特征金字塔网络)的思想,利用ViT中间不同深度的特征。较浅层的特征分辨率高,适合检测小物体;较深层的特征语义信息强,适合检测大物体。你可以修改ViTBackboneForDETR,让它返回多个尺度的特征图,然后在DETR的编码器或解码器中融合它们。

7.2 更高效的注意力机制

标准的Transformer自注意力计算复杂度是序列长度的平方。对于高分辨率图像,序列长度(patch数量)会很大,导致计算量爆炸。可以采用滑动窗口注意力(Swin Transformer)轴向注意力(Axial Attention) 等稀疏注意力机制来降低计算成本。timm库中也提供了Swin Transformer的预训练模型,你可以尝试用它作为骨干网络,这就是Swin Transformer-DETR,在很多检测榜单上表现优异。

7.3 知识蒸馏与模型压缩

ViT模型通常参数量巨大(ViT-Base约8600万参数),部署到端侧设备有困难。你可以使用知识蒸馏技术,用一个大的ViT模型(教师)去指导一个小的CNN或更小的ViT模型(学生)进行训练,让学生在保持精度的同时大幅减少参数量和计算量。另外,模型剪枝量化也是常用的压缩手段。PyTorch提供了相关的工具包,如torch.quantization,可以将FP32模型转换为INT8模型,显著提升推理速度。

7.4 在自己的数据集上微调

如果你有自己的标注数据,想用ViT-DETR来解决特定领域的问题(比如工业缺陷检测、医疗影像分析),微调是最快的路径。步骤通常是:

  1. 准备数据:将自己的数据转换成COCO格式。
  2. 修改类别数:将模型预测头中的num_classes改为你自己的类别数(记得+1给“无物体”类)。
  3. 加载预训练权重:加载在COCO等大数据集上预训练好的ViT-DETR权重。注意:分类头的权重因为类别数变了不能直接加载,需要小心处理。
  4. 调整学习率:对骨干网络使用更小的学习率(如base_lr * 0.01),对新加的或修改的层使用较大的学习率。
  5. 训练:在自己的数据上训练较少的轮数(Epoch)。

我自己的经验是,在数据量不是特别大的领域任务上,使用在ImageNet-21k或COCO上预训练过的ViT骨干进行微调,往往比从头训练一个CNN效果要好,收敛也更快,这得益于Transformer强大的迁移学习能力。

整个过程走下来,从理解ViT的Patch Embedding,到将其适配为检测骨干,再到实现DETR的匹配损失和训练技巧,最后进行评估和优化,这其实就是一个完整的AI项目研发流程。代码虽然看起来不少,但核心思想就是那么几条:用Transformer做全局关系建模,用可学习的查询去做物体探测,用二分图匹配来分配标签。希望这篇结合了实战和源码解析的长文,能帮你真正打通ViT目标检测的任督二脉。在实际动手时,多看看官方源码(如Detectron2中DETR的实现),多调试,多可视化,遇到问题别怕,那正是你深入理解模型的好机会。

Logo

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

更多推荐