摘要

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 注意力头数 h MLP 维度 Patch 大小 参数量
ViT-Base/16 12 768 12 3072 16×16 86M
ViT-Large/16 24 1024 16 4096 16×16 307M
ViT-Huge/14 32 1280 16 5120 14×14 632M

本文代码实现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

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐