摘要

2020 年,Google 团队提出Vision Transformer(ViT,视觉 Transformer),首次将纯 Transformer 架构直接应用于计算机视觉的图像分类任务,打破了卷积神经网络(CNN)在视觉领域十余年的统治地位。传统 ViT特指原始版本的 Vision Transformer,无任何 CNN 归纳偏置、无复杂改进,仅通过将图像拆分为序列令牌(Token),直接复用 NLP 领域的 Transformer Encoder 完成视觉建模。本文将从背景、前置知识、核心原理、网络结构、数学推导、完整代码实现、训练策略、优缺点等维度,对传统 ViT 进行全方位深度解析,帮助读者从零掌握这一视觉 Transformer 的奠基性模型。

关键词:Vision Transformer;ViT;图像分类;Transformer Encoder;自注意力机制


1. 引言

1.1 视觉领域的 CNN 统治时代

在 ViT 提出之前,计算机视觉任务(图像分类、检测、分割等)长期由卷积神经网络(CNN)主导。从 LeNet、AlexNet 到 ResNet、EfficientNet,CNN 依靠局部感受野、权重共享、平移不变性三大核心归纳偏置,成为视觉建模的标准范式。CNN 通过分层卷积提取局部特征(边缘→纹理→语义),但存在天然缺陷:卷积核感受野有限,难以建模图像中长距离的全局依赖关系(如远处物体的关联、全局上下文信息)。

1.2 Transformer 在 NLP 领域的颠覆性成功

2017 年,Google 提出Transformer架构,完全基于自注意力机制替代循环神经网络(RNN),成为 NLP 领域的基石。Transformer 能高效建模序列的全局依赖关系,且支持并行计算,后续 BERT、GPT 等大模型均基于 Transformer 构建。

研究者自然产生疑问:既然 Transformer 能处理文本序列,能否直接处理图像序列? 早期尝试将图像像素展平为序列,但计算量爆炸;直到 ViT 提出,通过图像分块(Patch) 替代像素序列,完美解决了计算效率问题。

1.3 传统 ViT 的提出与核心意义

2020 年,ICLR 论文《An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale》正式发布Vision Transformer(ViT)

  • 核心思想:将图像视为 “视觉单词序列”,类比 NLP 中的文本令牌,用纯 Transformer Encoder 完成图像分类;
  • 传统 ViT 定义:无 CNN、无额外改进、仅使用 Transformer Encoder 的原始视觉模型,是所有视觉 Transformer 的基础;
  • 核心结论:在大规模数据集上预训练后,ViT 性能超越 SOTA CNN,且计算效率更高。

本文聚焦传统 ViT,不涉及 Swin Transformer、DEiT 等后续改进版本,严格还原原始模型的原理与实现。


2. 前置知识:Transformer Encoder 核心原理

ViT仅使用 Transformer 的 Encoder 部分,无 Decoder,因此只需掌握 Transformer Encoder 的核心模块即可理解 ViT。

2.1 Transformer 整体架构(极简版)

原始 Transformer 包含Encoder(编码器)+ Decoder(解码器),ViT 仅保留堆叠的 Encoder Block,结构如下:

  1. 输入序列 → 嵌入层 → 位置编码;
  2. 堆叠 N 层 Encoder Block;
  3. 输出序列 → 任务头(分类 / 回归)。

2.2 自注意力机制(Self-Attention)

自注意力是 Transformer 的核心,作用是计算序列中每个令牌与所有令牌的相关性权重,实现全局依赖建模。

2.2.1 计算流程
  1. 对输入特征线性投影,生成三个向量:查询 Q(Query)、键 K(Key)、值 V(Value)
  2. 计算 Q 与 K 的点积,得到注意力分数(表征令牌间相关性);
  3. 对分数做 Softmax 归一化,得到注意力权重;
  4. 用权重对 V 加权求和,输出注意力特征。
2.2.2 数学公式

\text{Attention}(Q,K,V) = \text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V

  • d_k:K 向量的维度,除以\(\sqrt{d_k}\)是为了防止点积数值过大,导致 Softmax 梯度消失;
  • 输出:每个令牌融合了全局所有令牌的特征信息。

2.3 多头自注意力(MSA)

单一自注意力只能学习一种全局依赖模式,多头自注意力(Multi-Head Self-Attention, MSA) 将 Q/K/V 切分为 h 个头,并行学习不同的依赖关系,最后拼接输出:

\text{MSA}(Q,K,V) = \text{Concat}(\text{head}_1,...,\text{head}_h)W^O

ViT-Base 默认使用12 个注意力头

2.4 前馈神经网络(FFN)

Encoder 中,注意力模块后接两层全连接网络,中间用 GELU 激活函数:

\text{FFN}(x) = \max(0, xW_1+b_1)W_2+b_2

2.5 层归一化与残差连接

为解决深度网络训练退化问题,Transformer 采用Pre-LN 结构(层归一化放在模块前)+ 残差连接

x = x + \text{Module}(\text{LayerNorm}(x))

其中 Module 可以是 MSA 或 FFN。

2.6 Transformer Encoder Block

单个 Encoder Block 的固定结构: 层归一化 → 多头自注意力 → 残差连接 → 层归一化 → 前馈网络 → 残差连接

传统 ViT 就是将图像序列化后,输入堆叠的 Encoder Block 完成建模。


3. 传统 ViT 核心原理与创新点

CNN 的核心是局部特征提取,而 ViT 的核心是全局序列建模。传统 ViT 的所有创新都围绕如何将二维图像适配一维 Transformer 序列输入展开,三大核心创新:

3.1 核心创新 1:图像分块嵌入(Patch Embedding)

直接将图像像素展平会导致序列过长(如 224×224 图像 = 50176 个像素),计算量无法承受。ViT 将图像划分为固定大小的图像块(Patch),将每个 Patch 视为一个 “视觉令牌”,大幅缩短序列长度。

3.2 核心创新 2:分类令牌(Class Token)

基于NLP 中分类任务用 的token,ViT 借鉴该设计,在 Patch 序列前添加一个可学习的分类令牌,最终用该令牌的特征输出分类结果。

3.3 核心创新 3:可学习 1D 位置编码

Transformer 无天然的位置感知能力,ViT 放弃 NLP 的正弦余弦位置编码,使用可学习的一维位置编码,直接与 Patch Embedding 相加,注入位置信息。


4. 传统 ViT 网络结构逐模块详解

传统 ViT 的完整流程:图像输入 → Patch Embedding → 添加 Class Token → 加位置编码 → 堆叠 Transformer Encoder → 分类头输出。 以ViT-Base/16(最常用版本)为例,输入图像尺寸 224×224×3,Patch 大小 16×16,详细拆解如下:

4.1 模块 1:图像分块与 Patch Embedding

4.1.1 分块规则

4.1.2 特征映射

每个 Patch 展平为一维向量:P×P×C=16×16×3=768维; 通过线性层将向量投影到模型隐藏维度d(ViT-Base 中\(d=768\)),得到 Patch Embedding:N*d(196×768)。

4.1.3 实现方式

两种等价实现:

  1. 展平 Patch + 线性层;
  2. 卷积层(kernel=P, stride=P) 直接生成,效率更高。

4.2 模块 2:添加分类令牌(Class Token)

在 Patch Embedding 序列的最前方,拼接一个可学习的 Class Token(维度 1×d),序列长度变为\(N+1\)(197×768)。

  • 作用:聚合全局特征,最终仅用该 Token 的特征做分类,简化任务输出。

4.3 模块 3:位置编码(Positional Embedding)

生成可学习的位置编码:维度\((N+1) \times d\)(197×768),直接与 Patch+Class Token 的特征逐元素相加

  • 关键:ViT 使用1D 可学习位置编码,而非 2D 位置编码,证明纯 1D 编码已足够学习视觉位置信息。

4.4 模块 4:堆叠 Transformer Encoder

ViT-Base 堆叠12 层 Encoder Block,每层严格遵循 Pre-LN 结构:

  1. 层归一化 → 多头自注意力(12 头)→ 残差连接;
  2. 层归一化 → 前馈网络(隐藏层维度 4d=3072)→ 残差连接。

所有层共享相同结构,无卷积、无池化,纯注意力建模。

4.5 模块 5:分类头(MLP Head)

  1. 训练阶段:LayerNorm → Linear → GELU → Linear(映射到类别数);
  2. 推理阶段:仅用 Class Token 的特征,通过 Linear 层输出分类概率。

5. 传统 ViT 标准规格参数

原始 ViT 提供了 3 种规格,核心区别在于模型宽度、深度、注意力头数:

表格

模型规格层数 L隐藏维度 d注意力头数 hMLP 维度Patch 大小参数量
ViT-Base/161276812307216×1686M
ViT-Large/1624102416409616×16307M
ViT-Huge/1432128016512014×14632M

本文代码实现ViT-Base/16,是最常用的基础版本。


6. 传统 ViT PyTorch 完整代码实现

本代码严格还原原始传统 ViT,无任何改进,纯 PyTorch 实现,逐行注释,可直接运行。

6.1 环境依赖

import torch
import torch.nn as nn
import torch.nn.functional as F

6.2 核心模块实现

6.2.1 Patch Embedding 模块

将图像转换为 Patch 序列,支持卷积实现(高效版):

class PatchEmbedding(nn.Module):
    """
    图像分块嵌入层:将2D图像转为1D Patch序列
    输入:[batch_size, 3, H, W]
    输出:[batch_size, num_patches + 1, embed_dim] (+1是Class Token)
    """
    def __init__(self, img_size=224, patch_size=16, in_channels=3, embed_dim=768):
        super().__init__()
        self.img_size = img_size
        self.patch_size = patch_size
        # 计算Patch数量:(224/16)^2 = 196
        self.num_patches = (img_size // patch_size) ** 2
        
        # 用卷积实现分块+线性投影,等价于展平+线性层,效率更高
        self.proj = nn.Conv2d(
            in_channels=in_channels,
            out_channels=embed_dim,
            kernel_size=patch_size,
            stride=patch_size
        )

    def forward(self, x):
        # x: [B, 3, 224, 224] -> [B, embed_dim, 14, 14]
        x = self.proj(x)
        # 展平为序列:[B, embed_dim, 14, 14] -> [B, embed_dim, 196]
        x = x.flatten(2)
        # 转置维度:[B, embed_dim, 196] -> [B, 196, embed_dim]
        x = x.transpose(1, 2)
        return x
6.2.2 多头自注意力模块(MSA)
class MultiHeadAttention(nn.Module):
    """多头自注意力机制,严格遵循原始ViT实现"""
    def __init__(self, embed_dim=768, num_heads=12, qkv_bias=True, attn_drop=0., proj_drop=0.):
        super().__init__()
        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads  # 每个头的维度:768/12=64
        assert self.head_dim * num_heads == embed_dim, "嵌入维度必须能被头数整除"

        # QKV线性投影层
        self.qkv = nn.Linear(embed_dim, embed_dim * 3, bias=qkv_bias)
        self.attn_drop = nn.Dropout(attn_drop)
        self.proj = nn.Linear(embed_dim, embed_dim)
        self.proj_drop = nn.Dropout(proj_drop)

    def forward(self, x):
        B, N, C = x.shape  # B:批次, N:序列长度, C:嵌入维度

        # 生成QKV:[B, N, 3*embed_dim] -> 拆分后3个[B, num_heads, N, head_dim]
        qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4)
        q, k, v = qkv.unbind(0)

        # 自注意力计算:Q*K^T / sqrt(d_k)
        attn_score = (q @ k.transpose(-2, -1)) * (self.head_dim ** -0.5)
        attn_weight = attn_score.softmax(dim=-1)
        attn_weight = self.attn_drop(attn_weight)

        # 加权求和 + 投影
        x = (attn_weight @ v).transpose(1, 2).reshape(B, N, C)
        x = self.proj(x)
        x = self.proj_drop(x)
        return x
6.2.3 Transformer Encoder Block
class TransformerBlock(nn.Module):
    """单个Transformer Encoder块:Pre-LN + MSA + FFN + 残差"""
    def __init__(self, embed_dim=768, num_heads=12, mlp_ratio=4., qkv_bias=True, drop=0., attn_drop=0.):
        super().__init__()
        self.norm1 = nn.LayerNorm(embed_dim)  # 层归一化
        self.attn = MultiHeadAttention(embed_dim, num_heads, qkv_bias, attn_drop, drop)
        self.norm2 = nn.LayerNorm(embed_dim)
        # MLP:隐藏层维度 = embed_dim * 4
        mlp_hidden_dim = int(embed_dim * mlp_ratio)
        self.mlp = nn.Sequential(
            nn.Linear(embed_dim, mlp_hidden_dim),
            nn.GELU(),  # ViT使用GELU激活函数
            nn.Dropout(drop),
            nn.Linear(mlp_hidden_dim, embed_dim),
            nn.Dropout(drop)
        )

    def forward(self, x):
        # 残差连接1:MSA分支
        x = x + self.attn(self.norm1(x))
        # 残差连接2:FFN分支
        x = x + self.mlp(self.norm2(x))
        return x
6.2.4 传统 ViT 完整模型
class ViT(nn.Module):
    """
    传统Vision Transformer(ViT-Base/16)
    img_size: 输入图像尺寸
    patch_size: 分块大小
    num_classes: 分类类别数
    embed_dim: 嵌入维度
    depth: Encoder层数
    num_heads: 注意力头数
    """
    def __init__(
        self,
        img_size=224,
        patch_size=16,
        num_classes=1000,
        embed_dim=768,
        depth=12,
        num_heads=12,
        mlp_ratio=4.,
        qkv_bias=True,
        drop_rate=0.,
        attn_drop_rate=0.
    ):
        super().__init__()
        # 1. Patch嵌入层
        self.patch_embed = PatchEmbedding(img_size, patch_size, embed_dim=embed_dim)
        num_patches = self.patch_embed.num_patches

        # 2. 可学习分类令牌Class Token
        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
        # 3. 可学习位置编码
        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
        self.pos_drop = nn.Dropout(p=drop_rate)

        # 4. 堆叠Transformer Encoder
        self.blocks = nn.Sequential(*[
            TransformerBlock(embed_dim, num_heads, mlp_ratio, qkv_bias, drop_rate, attn_drop_rate)
            for _ in range(depth)
        ])

        # 5. 分类头
        self.norm = nn.LayerNorm(embed_dim)
        self.head = nn.Linear(embed_dim, num_classes)

        # 权重初始化
        nn.init.trunc_normal_(self.pos_embed, std=0.02)
        nn.init.trunc_normal_(self.cls_token, std=0.02)
        self.apply(self._init_weights)

    def _init_weights(self, m):
        if isinstance(m, nn.Linear):
            nn.init.trunc_normal_(m.weight, std=0.02)
            if isinstance(m, nn.Linear) and m.bias is not None:
                nn.init.constant_(m.bias, 0)
        elif isinstance(m, nn.LayerNorm):
            nn.init.constant_(m.bias, 0)
            nn.init.constant_(m.weight, 1.0)

    def forward(self, x):
        # 步骤1:Patch嵌入
        x = self.patch_embed(x)  # [B, 196, 768]

        # 步骤2:添加Class Token
        cls_token = self.cls_token.expand(x.shape[0], -1, -1)  # [B, 1, 768]
        x = torch.cat((cls_token, x), dim=1)  # [B, 197, 768]

        # 步骤3:加位置编码 + Dropout
        x = x + self.pos_embed
        x = self.pos_drop(x)

        # 步骤4:通过Transformer Encoder
        x = self.blocks(x)

        # 步骤5:提取Class Token特征 + 分类
        x = self.norm(x)
        cls_token_final = x[:, 0]  # 取第一个令牌(Class Token)
        x = self.head(cls_token_final)

        return x

6.3 代码测试

验证模型输入输出是否正确:

if __name__ == "__main__":
    # 初始化ViT-Base/16
    model = ViT(
        img_size=224,
        patch_size=16,
        num_classes=1000,
        embed_dim=768,
        depth=12,
        num_heads=12
    )

    # 构造随机输入:[批次, 通道, 高, 宽]
    dummy_input = torch.randn(2, 3, 224, 224)
    # 前向传播
    output = model(dummy_input)
    print(f"输入形状: {dummy_input.shape}")
    print(f"输出形状: {output.shape}")  # 预期输出:[2, 1000]

6.4 代码说明

  1. 严格还原传统 ViT结构,无任何改进;
  2. 输入:224×224×3 的图像,输出:1000 类分类概率;
  3. 可通过修改num_classes适配自定义数据集;
  4. 参数量与原始 ViT-Base/16 完全一致(86M)。

7. 传统 ViT 训练策略与实验结果

7.1 数据集要求

传统 ViT 的核心特性:对数据集规模敏感

  • 小数据集(如 CIFAR-10):ViT 性能不如 ResNet,缺乏 CNN 的归纳偏置;
  • 大规模数据集(ImageNet-1K/21K、JFT-300M):ViT 性能超越所有 CNN,全局建模优势凸显。

7.2 预训练与微调

  1. 预训练:在 JFT-300M(3 亿图像)上预训练,学习通用视觉特征;
  2. 微调:在目标数据集(ImageNet-1K)上微调,仅修改分类头。

7.3 关键超参数

  • 优化器:AdamW;
  • 学习率:1e-3;
  • 数据增强:RandAugment、MixUp;
  • 批次大小:4096。

7.4 实验性能

ImageNet-1K 数据集上:

  • ViT-Base/16:81.2% Top-1 准确率;
  • ViT-Large/16:84.9% Top-1 准确率,超越 ResNet50(76.1%)。

8. 传统 ViT 的优势与局限性

8.1 核心优势

  1. 全局感受野:自注意力直接建模全局依赖,远超 CNN 的局部感受野;
  2. 结构简单:无卷积、无池化,纯全连接 + 注意力,易于扩展;
  3. 迁移能力强:大规模预训练后,泛化性优于 CNN;
  4. 可扩展性好:模型越大(Large/Huge),性能越高,符合大模型规律。

8.2 核心局限性

  1. 小数据集性能差:缺乏 CNN 的局部归纳偏置,小数据易过拟合;
  2. 计算量大:自注意力复杂度为\(O(N^2)\)(N 为序列长度),高分辨率图像效率低;
  3. 缺乏层级特征:ViT 输出单一尺度特征,不适合检测 / 分割等密集预测任务;
  4. 位置编码简单:1D 位置编码无法充分利用图像 2D 空间结构。

8.3 与 CNN 的核心区别

特性CNN传统 ViT
归纳偏置局部性、平移不变性
感受野局部、逐层扩大全局
计算复杂度\(O(N)\)\(O(N^2)\)
数据依赖性高(需大规模数据)

9. 传统 ViT 的行业影响

传统 ViT 是视觉 Transformer 的开山之作,彻底改变了计算机视觉的发展方向:

  1. 开启了视觉 Transformer 时代,后续 Swin Transformer、DEiT、BEiT、MAE 等模型均基于 ViT 改进;
  2. 证明了纯注意力机制可替代 CNN 完成视觉建模,为多模态大模型(文生图、图文理解)奠定基础;
  3. 统一了 NLP 与视觉的架构设计,推动了通用人工智能的发展。

10. 总结

传统 Vision Transformer(ViT)是视觉领域的里程碑式模型,其核心创新是将图像转化为 Patch 序列,直接复用 Transformer Encoder 实现全局视觉建模。本文完整解析了传统 ViT 的背景、原理、网络结构、数学公式,并提供了可直接运行的 PyTorch 原生代码,严格还原了原始模型的设计思想。

传统 ViT 的核心价值不在于极致性能,而在于证明了 Transformer 在视觉领域的可行性,为后续所有视觉 Transformer 提供了基础框架。尽管存在小数据集性能差、计算量大等缺陷,但它彻底打破了 CNN 的垄断,成为现代计算机视觉的基石。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐