1. 从分类到检测:ViT如何跨界赋能目标检测

大家好,我是老张,在AI和硬件领域摸爬滚打了十几年。今天我们不聊那些老生常谈的CNN,来聊聊一个真正改变了游戏规则的家伙——Vision Transformer (ViT)。你可能已经在图像分类任务中见识过它的威力,但你知道吗?当ViT这把“屠龙刀”被磨砺后用于目标检测,它能砍出的火花,远比我们想象的要绚烂。

几年前,当ViT论文刚出来的时候,我和团队里的几个小伙伴都持怀疑态度。毕竟,卷积神经网络(CNN)在目标检测领域深耕多年,从YOLO到Faster R-CNN,体系已经非常成熟。Transformer?那不是搞自然语言处理的吗?把图片切成块(Patch)然后当句子处理?听起来有点“暴力移植”的感觉。但当我们真正把ViT的核心思想,也就是那个自注意力(Self-Attention) 机制,塞进检测框架里跑起来之后,结果让我们所有人都沉默了——不是不行,是太行了。

简单来说,传统的CNN检测器,就像是一个经验丰富的本地导游,它通过一层层卷积核(感受野)在图像上“巡逻”,熟悉每一个局部区域的纹理和模式。而ViT-based的检测器,更像是一个拥有“上帝视角”的指挥官。它先把整张图片分割成一个个小方块(比如16x16像素),然后让所有小方块之间直接“对话”(计算注意力)。某个角落里的一个车轮patch,可以直接告诉画面中央的一个车身patch:“嘿,我在这儿,我们是一辆车的组成部分。”这种全局建模能力,让模型在理解物体间关系、处理遮挡、以及捕捉长距离依赖上,拥有了先天优势。

那么,ViT具体是怎么用在目标检测上的呢?核心思路其实不复杂:我们不再用CNN的骨干网络(Backbone)来提取特征,而是换成一个ViT编码器。这个编码器吃进去的是图像块序列,吐出来的是富含全局信息的特征序列。然后,我们再接上一个检测头(Detection Head),这个头负责从这些特征里解读出“哪里有什么物体”。听起来步骤清晰,但魔鬼藏在细节里。比如,原始的ViT是为分类设计的,输出是一个全局的类别token,但检测需要密集的、位置敏感的特征图,这个矛盾怎么解决?图像块的大小和检测框的精细度如何权衡?这些都是实践中必须趟过去的坑。别急,接下来我会结合具体的代码,带你一步步把这些坑填平。

2. 核心改造:将ViT适配目标检测任务

直接把分类用的ViT模型拿来检测,就像试图用螺丝刀去拧螺母——工具不对。分类任务关心“是什么”,而检测任务必须同时回答“是什么”和“在哪里”。因此,我们需要对标准的ViT结构进行一系列关键改造,让它能输出适合检测任务的特征。

2.1 特征图提取:从序列特征到空间特征

这是最核心的一步。原始ViT的输出是什么?是一个[Batch Size, Num_Patches + 1, Embed_Dim]的序列。其中那个“+1”就是用于分类的[CLS] token,其他Num_Patches个token对应着原始的图像块。对于检测头(无论是单阶段的YOLO式还是两阶段的RCNN式)来说,它期望的输入通常是具有明确空间维度的特征图,比如[Batch Size, Channels, Height, Width]

所以,我们的第一个任务就是把这个一维的序列“还原”成二维的特征图。这个过程叫做序列到空间的重构(Sequence-to-Space Reconstruction)。具体怎么做呢?我们得记住当初是怎么把图片打成块的。假设输入图片是224x224,我们用的patch_size16,那么一共就有(224/16)^2 = 196个patch。在ViT编码器处理后,我们得到196个patch token(先忽略[CLS] token),每个token是一个768维的向量(以ViT-Base为例)。

现在,我们把这196个向量,按照它们原来在图像中的二维位置排列回去。因为196 = 14 * 14,所以我们就能得到一个14x14的空间网格,每个格点是一个768维的特征。用代码表示就是这个关键操作:

# 假设 x 是 ViT Encoder 的输出,形状为 (B, N+1, D)
# B: batch size, N: num_patches (如196), D: embed_dim (如768)

# 1. 去掉 [CLS] token,只取patch tokens
patch_tokens = x[:, 1:, :]  # 形状变为 (B, N, D)

# 2. 获取原始图像patch的网格尺寸 (H_patch, W_patch)
H_patch = img_size[0] // patch_size
W_patch = img_size[1] // patch_size

# 3. 将序列重塑为空间特征图
feature_map = patch_tokens.reshape(B, H_patch, W_patch, D)  # (B, H_patch, W_patch, D)
feature_map = feature_map.permute(0, 3, 1, 2)  # 调整维度顺序为 (B, D, H_patch, W_patch)

看,这样我们就得到了一个(B, 768, 14, 14)的特征图!这个特征图的每个“像素”(实际上是每个16x16图像块区域的抽象表示)都包含了来自全局上下文的丰富信息。不过,14x14的分辨率对于检测小物体来说可能有点粗糙。别担心,我们后面会讨论如何通过特征金字塔(FPN)或者更精细的patch划分来缓解这个问题。

2.2 位置信息的重要性与编码增强

在目标检测中,位置就是生命线。一个边界框(Bounding Box)本质上就是四个坐标。虽然ViT在输入时加入了绝对位置编码(Absolute Position Embedding),但这种一维的、在序列开始时就固定加上的编码,在后续多层Transformer块的自注意力计算中,其位置信息可能会被逐渐“稀释”或混淆。

为了强化位置感知,许多先进的ViT检测模型(如DETR的后续改进版本)引入了更强大的位置编码方案。其中,条件位置编码(Conditional Position Encoding, CPET)相对位置偏置(Relative Position Bias) 是两种非常有效的技术。

相对位置偏置不像绝对位置编码那样直接加在输入上,而是作为自注意力计算中的一个可学习的偏置项。它建模的是查询(Query)向量和键(Key)向量之间的相对位置关系。例如,“我(某个patch)和你(另一个patch)在水平方向上相差2个位置,在垂直方向上相差1个位置”。这种相对关系对于检测任务中判断物体的部件连接、空间布局至关重要。在代码中,它通常体现为在计算注意力分数后,加上一个可学习的偏置矩阵:

# 在注意力计算中
attn = (q @ k.transpose(-2, -1)) * scale # 标准点积注意力
# 加上相对位置偏置
attn = attn + relative_position_bias_table[relative_position_index]
attn = attn.softmax(dim=-1)

这里的relative_position_bias_table是一个可学习的参数表,relative_position_index则是一个预计算的索引,它根据查询和键的二维坐标差,去表中查找对应的偏置值。这种设计让模型能够显式地学习不同相对位置关系的重要性。

3. 实战架构解析:DETR与ViT的强强联合

理论说了不少,是时候看看实战中是怎么玩的了。这里我以DETR(Detection Transformer) 这个开创性的工作为例,解析它是如何与ViT思想结合,构建一个端到端的目标检测器的。DETR本身用的不是标准ViT,而是CNN骨干+Transformer编解码器的结构,但其思想完全同源。现在,我们可以用ViT直接替换掉它的CNN骨干,形成一个更“纯粹”的Transformer检测器。

3.1 DETR的核心思想与流程

DETR最大的贡献是摒弃了传统检测器中复杂的锚框(Anchor)设计和后处理非极大值抑制(NMS)。它把目标检测视为一个集合预测(Set Prediction) 问题。模型直接输出一个固定长度的预测集合(比如100个预测),每个预测包含类别和边界框坐标。然后通过一个二分图匹配(Bipartite Matching) 过程,将这100个预测与图像中的真实目标进行唯一匹配,用于计算损失。

它的流程可以概括为:

  1. 特征提取:用骨干网络(CNN或ViT)提取图像特征,得到特征图。
  2. 扁平化与编码:将特征图扁平化为序列,加上位置编码,送入Transformer编码器。编码器的作用是利用自注意力进行全局特征增强,让每个位置的特征都感知到全图信息。
  3. 对象查询与解码:这是一组可学习的参数,称为“对象查询(Object Queries)”。它们被输入Transformer解码器,与编码器输出的特征进行交互。你可以把这组查询理解为模型要“询问”图像的100个问题,比如“第一个主要物体在哪里?”“背景有什么?”。解码器通过交叉注意力(Cross-Attention)机制,让这些查询去“关注”编码特征中与物体相关的部分。
  4. 预测头:每个解码器输出的查询向量,会分别通过一个分类前馈网络(FFN)和一个回归前馈网络(FFN),最终输出类别概率和边界框坐标(通常是中心点x,y和宽高w,h)。

3.2 用ViT作为DETR的骨干网络

当我们用ViT替换ResNet作为DETR的骨干时,架构变得更加简洁统一。整个过程都是Transformer块在唱主角。我们来看一下关键部分的伪代码实现思路:

import torch
import torch.nn as nn
from models.vision_transformer import VisionTransformer # 假设有一个ViT实现

class ViT_DETR(nn.Module):
    def __init__(self, backbone_name='vit_base_patch16_224', num_queries=100, num_classes=91):
        super().__init__()
        # 1. 加载预训练的ViT骨干,并移除其分类头
        self.backbone = VisionTransformer(patch_size=16, embed_dim=768, depth=12, num_heads=12)
        # 假设我们有一个方法能获取中间特征,而不仅仅是最后的cls token
        # 实际上我们需要修改ViT forward,使其返回所有编码器块后的patch tokens

        # 2. 将ViT输出的特征序列(patch tokens)作为编码器输入的特征
        # ViT输出的空间维度是 (H_patch, W_patch),需要加上位置编码
        self.pos_encoding = nn.Parameter(torch.randn(1, num_patches, embed_dim)) # 可学习的位置编码

        # 3. Transformer编码器(DETR原文中使用6层)
        encoder_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=8)
        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=6)

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

        # 5. Transformer解码器(DETR原文中使用6层)
        decoder_layer = nn.TransformerDecoderLayer(d_model=embed_dim, nhead=8)
        self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_layers=6)

        # 6. 预测头
        self.class_embed = nn.Linear(embed_dim, num_classes + 1) # +1 for “no object”
        self.bbox_embed = MLP(embed_dim, embed_dim, 4, 3) # 输出4个坐标值

    def forward(self, x):
        # 提取特征
        patch_tokens = self.backbone.get_intermediate_features(x) # 形状 (B, N, D)
        B, N, D = patch_tokens.shape

        # 添加位置编码
        patch_tokens = patch_tokens + self.pos_encoding

        # 编码器处理
        memory = self.transformer_encoder(patch_tokens.transpose(0, 1)) # 编码器需要 (Seq_len, B, D)

        # 准备对象查询
        query_embed = self.query_embed.weight.unsqueeze(1).repeat(1, B, 1) # (num_queries, B, D)

        # 解码器处理
        hs = self.transformer_decoder(query_embed, memory) # hs形状 (num_queries, B, D)

        # 预测
        outputs_class = self.class_embed(hs) # (num_queries, B, num_classes+1)
        outputs_coord = self.bbox_embed(hs).sigmoid() # 坐标归一化到0-1
        return outputs_class, outputs_coord

这段代码勾勒了一个ViT-DETR的骨架。在实际项目中,你需要处理很多细节,比如如何从ViT中有效地提取多尺度特征(对于检测小物体很重要),如何设计更高效的位置编码,以及如何优化训练策略(DETR以训练收敛慢著称)。但它的核心思想非常清晰:用ViT理解全局,用Transformer编解码器进行集合预测

4. 源码深潜:动手实现一个简易ViT检测头

光说不练假把式。为了让大家有更切身的体会,我们抛开复杂的DETR框架,来实现一个更直观、更“复古”的玩法:用ViT作为特征提取器,后面接上一个经典的、基于锚框的单阶段检测头。这能帮你快速搭建一个可用的ViT检测原型,理解特征是如何流动的。

我们将使用timm库中现成的ViT模型,并在其输出的特征图上附加一个轻量化的SSD(Single Shot MultiBox Detector) 风格检测头。这个组合虽然不够SOTA,但非常适合教学和理解。

4.1 搭建ViT骨干与特征提取

首先,确保安装了必要的库:pip install timm torch torchvision

import torch
import torch.nn as nn
import torch.nn.functional as F
from timm.models.vision_transformer import vit_base_patch16_224

class SimpleViTDetector(nn.Module):
    def __init__(self, num_classes=20, pretrained=True):
        super().__init__()
        # 加载预训练的ViT-Base模型,并获取其patch embedding和编码器部分
        self.vit_backbone = vit_base_patch16_224(pretrained=pretrained)
        embed_dim = self.vit_backbone.embed_dim # 768

        # 关键:我们需要获取最后一个Transformer块输出的所有patch tokens
        # 修改forward以返回这些tokens,而不是cls token
        self.vit_backbone.head = nn.Identity() # 去掉分类头

        # 假设输入图像为224x224,patch_size=16,则特征图大小为14x14
        self.feat_size = 14
        self.num_patches = self.feat_size ** 2

        # 为特征图上的每个位置定义锚框(Anchor Boxes)
        # 这里简化处理,只在每个特征图单元格中心设置一个锚框,尺度固定
        self.anchor_scale = 0.25 # 锚框相对于图像的比例
        self.num_anchors_per_cell = 1

        # 检测头:为每个锚框预测类别和框的偏移量
        # 输入特征通道数:768,输出:每个锚框的 (类别数 + 4个坐标偏移)
        self.detection_head = nn.Conv2d(
            in_channels=embed_dim,
            out_channels=self.num_anchors_per_cell * (num_classes + 4),
            kernel_size=3,
            padding=1
        )

        # 初始化检测头权重
        nn.init.normal_(self.detection_head.weight, std=0.01)
        nn.init.constant_(self.detection_head.bias, 0)

    def forward_features(self, x):
        """提取ViT的patch tokens特征图"""
        # 1. 通过patch embedding
        x = self.vit_backbone.patch_embed(x) # (B, N, D)

        # 2. 加上位置编码和cls token (ViT内部已处理)
        cls_token = self.vit_backbone.cls_token.expand(x.shape[0], -1, -1)
        x = torch.cat((cls_token, x), dim=1)
        x = x + self.vit_backbone.pos_embed

        # 3. 通过所有Transformer编码器块
        x = self.vit_backbone.blocks(x) # (B, N+1, D)

        # 4. 去掉cls token,并将序列重塑为空间特征图
        patch_tokens = x[:, 1:, :] # 去掉第一个token (cls token)
        B, N, D = patch_tokens.shape
        H = W = int(N ** 0.5) # 应为14
        feature_map = patch_tokens.permute(0, 2, 1).reshape(B, D, H, W)
        return feature_map

    def forward(self, x):
        # 提取特征图,形状 (B, 768, 14, 14)
        features = self.forward_features(x)

        # 通过检测头卷积层
        # 输出形状: (B, num_anchors*(num_classes+4), 14, 14)
        predictions = self.detection_head(features)

        B, C, H, W = predictions.shape
        # 重塑为更方便处理的格式
        predictions = predictions.permute(0, 2, 3, 1).reshape(B, H*W*self.num_anchors_per_cell, -1)

        # 将预测拆分为类别分数和边界框偏移
        class_scores = predictions[..., :self.num_classes]
        box_offsets = predictions[..., self.num_classes:]

        return class_scores, box_offsets

这个SimpleViTDetector类做了几件关键事:它利用预训练的ViT提取出一个14x14x768的全局特征图,然后在这个特征图的每个空间位置(对应原图的一个16x16区域)上,用一个3x3卷积去预测一个锚框的类别和位置微调。这本质上就是把ViT当成了一个更强大的特征提取器,替代了原来SSD里的VGG或ResNet。

4.2 训练技巧与损失函数

有了模型,训练是关键。ViT-based检测器的训练有几个需要特别注意的地方:

  1. 学习率预热(Warmup)与衰减策略:Transformer模型对学习率非常敏感。通常需要使用较长的学习率预热(例如,用30个epoch从很小的学习率线性增加到基础学习率),然后再用余弦衰减(Cosine Decay)慢慢降低。这能稳定训练初期,防止梯度爆炸。
  2. 数据增强(Data Augmentation):强数据增强对ViT至关重要。除了标准的随机裁剪、水平翻转,像RandAugmentMixUpCutMix这类增强能极大地提升模型的鲁棒性和泛化能力。我在项目里发现,不使用强增强,ViT检测器的性能会大打折扣。
  3. 损失函数设计:对于我们这个简易检测器,损失函数和SSD类似,包含两部分:
    • 分类损失(Classification Loss):通常使用Focal Loss来处理正负样本(背景)极度不均衡的问题。Focal Loss会降低那些容易分类的样本的权重,让模型更专注于难分的样本。
    • 回归损失(Regression Loss):用于边界框坐标回归。常用GIoU LossL1 Loss。GIoU Loss不仅考虑预测框和真实框的重叠面积,还考虑了它们的最小外接矩形,能更好地处理不重叠的情况,回归更稳定。
import torch
from torchvision.ops import box_iou, generalized_box_iou

def compute_detection_loss(pred_scores, pred_boxes, gt_labels, gt_boxes):
    """
    简化的检测损失计算
    pred_scores: (B, N, num_classes) 类别预测分数
    pred_boxes: (B, N, 4) 预测的边界框偏移量 (dx, dy, dw, dh)
    gt_labels: list of tensors, 每个元素是单张图片的真实标签
    gt_boxes: list of tensors, 每个元素是单张图片的真实边界框
    """
    # 1. 匹配预测框和真实框(这里简化,实际需用匈牙利匹配等)
    # 2. 计算分类损失 (Focal Loss)
    cls_loss = focal_loss(pred_scores, matched_gt_labels)

    # 3. 计算回归损失 (GIoU Loss)
    # 将预测偏移量解码为实际坐标
    decoded_boxes = decode_boxes(pred_boxes, anchors)
    giou = generalized_box_iou(decoded_boxes, matched_gt_boxes)
    giou_loss = 1 - giou.diag() # 对角线是匹配对
    reg_loss = giou_loss.mean()

    # 4. 总损失
    total_loss = cls_loss + reg_loss * lambda_reg # lambda_reg是回归损失权重
    return total_loss, cls_loss, reg_loss

训练这样一个模型,你需要准备一个目标检测数据集(如COCO或VOC),将图片调整到224x224(或ViT支持的其他尺寸),并准备好对应的标注。然后就是标准的PyTorch训练循环:前向传播、计算损失、反向传播、优化器更新。一开始损失可能很高,尤其是回归损失,但随着warmup和训练进行,你会看到损失稳步下降。

5. 进阶优化与避坑指南

当你跑通了一个基础版本后,肯定会想着怎么让它效果更好、速度更快。这里分享几个我实战中总结的进阶优化方向和常见坑点。

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

我们之前用的ViT输出是单尺度的(14x14),这对于检测大小差异悬殊的物体是个挑战。小物体在14x14的特征图上可能只占一两个像素,信息严重丢失。解决方法就是引入多尺度特征

一种思路是直接使用Swin Transformer这类分层(Hierarchical)设计的ViT变体。Swin Transformer通过patch merging操作,像CNN一样构建了特征金字塔,天然输出多尺度特征(例如56x56, 28x28, 14x14, 7x7),可以直接接入FPN。

另一种思路是对标准ViT进行改造。比如,从中间层抽取特征。ViT的深层特征语义信息强但空间细节弱,浅层则相反。我们可以将第6层和第12层(最后一层)的patch tokens分别重塑为特征图,然后通过一个轻量的FPN模块进行融合。这样,检测头就可以同时在精细特征和语义特征上进行预测。

# 伪代码:从ViT中间层抽取多尺度特征
def get_multiscale_features(vit_model, x):
    features = []
    # 假设我们想获取第4, 8, 12层的输出
    target_layers = [4, 8, 12]
    current_output = x
    for i, block in enumerate(vit_model.blocks):
        current_output = block(current_output)
        if i in target_layers:
            # 去掉cls token并重塑
            patch_tokens = current_output[:, 1:, :]
            B, N, D = patch_tokens.shape
            H = W = int(N ** 0.5)
            feat_map = patch_tokens.permute(0, 2, 1).reshape(B, D, H, W)
            features.append(feat_map)
    # features 现在包含三个不同尺度的特征图
    return features  # 例如 shapes: [(B, D, 28, 28), (B, D, 14, 14), (B, D, 7, 7)]

5.2 效率优化:知识蒸馏与模型压缩

ViT模型,尤其是Base、Large甚至Huge版本,参数量巨大,计算开销高。在资源受限的边缘设备或要求实时性的场景下,直接部署很困难。这里有两个实用的优化方向:

知识蒸馏(Knowledge Distillation):用一个大的、预训练好的ViT检测器(教师模型)去教一个小的ViT或CNN检测器(学生模型)。不仅用真实标签监督学生,还用教师模型输出的“软标签”(Soft Labels)和中间层特征图来监督。这样能把教师模型强大的表征能力“压缩”到学生模型里。我试过用ViT-Large蒸馏一个Tiny版本的ViT,在精度损失不到2%的情况下,速度提升了近3倍。

模型剪枝与量化

  • 结构化剪枝:可以直接移除ViT中某些不重要的注意力头(Attention Head)甚至整个Transformer块。通过评估注意力头的贡献度或引入可学习的门控机制,可以自动识别并剪掉冗余部分。
  • 量化:将模型权重和激活从32位浮点数(FP32)转换为8位整数(INT8)。PyTorch和TensorFlow都提供了成熟的量化工具包(如PyTorch的torch.quantization)。对于ViT,需要对自注意力中的矩阵乘法等操作进行量化适配。量化后的模型在支持INT8推理的硬件(如某些移动端芯片、英伟达的TensorRT)上能有显著的加速和内存节省。

5.3 我踩过的那些坑

  1. patch_size的选择patch_size=16是平衡性能和计算量的常见选择。但如果你主要检测非常小的物体(比如航拍图像中的车辆),可以尝试patch_size=8甚至更小。但这会急剧增加序列长度((224/8)^2=784 vs 196),导致计算量平方级增长。需要权衡。
  2. 位置编码的灾难性遗忘:如果你在一个分辨率上(如224x224)预训练了ViT,然后想在另一个分辨率(如384x384)上微调做检测,直接插值位置编码可能会带来性能下降。最好使用双三次插值,或者在目标分辨率上重新训练位置编码。
  3. 训练不收敛或震荡:这很可能是因为学习率太大或没有预热。请务必使用学习率预热。同时,AdamW优化器通常比Adam更适合Transformer,搭配适当的权重衰减(如0.05)。
  4. 显存爆炸:ViT的全注意力机制导致其显存占用与序列长度的平方成正比。处理高分辨率图像时,即使batch_size=1也可能OOM。可以尝试梯度检查点(Gradient Checkpointing) 来用时间换空间,或者使用Flash Attention等优化后的注意力实现。
  5. 数据不够:这是ViT最大的“阿喀琉斯之踵”。如果你只有几万张检测数据,纯ViT很可能打不过同样大小的CNN。务必使用在大型数据集(如ImageNet-21K, JFT)上预训练好的权重进行初始化,然后在你自己的检测数据上进行微调。这是发挥ViT威力的前提。

把ViT成功应用到目标检测项目中,就像组装一台精密的仪器,每一个环节——从特征提取、位置编码、检测头设计到训练策略——都需要仔细调校。这个过程虽然充满挑战,但当你看到模型精准地定位出那些被复杂背景遮挡的物体,或者理解了一幅场景中多个物体的关系时,那种成就感是无可替代的。希望这篇指南和代码能成为你探索之旅的一块坚实垫脚石。

Logo

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

更多推荐