Vision Transformer (ViT) 在目标检测中的实战应用与源码解析
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-python和pycocotools是处理数据和评估模型必备的。
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编解码 -> 预测头。
- 特征提取:用CNN(如ResNet)或ViT作为骨干网络(Backbone),从输入图像中提取一个2D的特征图。如果我们用ViT,需要一点小改动,这个后面细说。
- Transformer编解码:这是DETR的灵魂。编码器对骨干网络提取的特征进行增强。解码器则接收一组固定数量的“对象查询”(Object Queries,比如100个),这些查询是可学习的参数,可以理解为模型学习到的、用于询问图像中“哪里有物体”的探针。解码器让这些对象查询与编码器输出的特征进行交互,最终每个查询输出一个包含物体信息的嵌入向量。
- 预测头:一个简单的全连接网络,接在每个对象查询的输出后面。它预测两个东西:类别(包括一个“无物体”类)和边界框(中心点坐标、宽高)。
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。我的经验是:
- 前向传播检查:在训练开始前,用一批数据跑一次前向传播,确保输出形状正确,没有NaN或Inf。
- 损失值监控:在训练初期,密切监控各个损失分量(分类损失、L1框损失、GIoU损失)。如果某个损失突然飙升或变为NaN,立即停止训练,检查数据或模型。
- 可视化中间特征:有时可以可视化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来解决特定领域的问题(比如工业缺陷检测、医疗影像分析),微调是最快的路径。步骤通常是:
- 准备数据:将自己的数据转换成COCO格式。
- 修改类别数:将模型预测头中的
num_classes改为你自己的类别数(记得+1给“无物体”类)。 - 加载预训练权重:加载在COCO等大数据集上预训练好的ViT-DETR权重。注意:分类头的权重因为类别数变了不能直接加载,需要小心处理。
- 调整学习率:对骨干网络使用更小的学习率(如
base_lr * 0.01),对新加的或修改的层使用较大的学习率。 - 训练:在自己的数据上训练较少的轮数(Epoch)。
我自己的经验是,在数据量不是特别大的领域任务上,使用在ImageNet-21k或COCO上预训练过的ViT骨干进行微调,往往比从头训练一个CNN效果要好,收敛也更快,这得益于Transformer强大的迁移学习能力。
整个过程走下来,从理解ViT的Patch Embedding,到将其适配为检测骨干,再到实现DETR的匹配损失和训练技巧,最后进行评估和优化,这其实就是一个完整的AI项目研发流程。代码虽然看起来不少,但核心思想就是那么几条:用Transformer做全局关系建模,用可学习的查询去做物体探测,用二分图匹配来分配标签。希望这篇结合了实战和源码解析的长文,能帮你真正打通ViT目标检测的任督二脉。在实际动手时,多看看官方源码(如Detectron2中DETR的实现),多调试,多可视化,遇到问题别怕,那正是你深入理解模型的好机会。
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐


所有评论(0)