传统 Vision Transformer(ViT)原理深度解析与 PyTorch 代码实现
摘要
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,结构如下:
- 输入序列 → 嵌入层 → 位置编码;
- 堆叠 N 层 Encoder Block;
- 输出序列 → 任务头(分类 / 回归)。
2.2 自注意力机制(Self-Attention)
自注意力是 Transformer 的核心,作用是计算序列中每个令牌与所有令牌的相关性权重,实现全局依赖建模。
2.2.1 计算流程
- 对输入特征线性投影,生成三个向量:查询 Q(Query)、键 K(Key)、值 V(Value);
- 计算 Q 与 K 的点积,得到注意力分数(表征令牌间相关性);
- 对分数做 Softmax 归一化,得到注意力权重;
- 用权重对 V 加权求和,输出注意力特征。
2.2.2 数学公式
:K 向量的维度,除以\(\sqrt{d_k}\)是为了防止点积数值过大,导致 Softmax 梯度消失;
- 输出:每个令牌融合了全局所有令牌的特征信息。
2.3 多头自注意力(MSA)
单一自注意力只能学习一种全局依赖模式,多头自注意力(Multi-Head Self-Attention, MSA) 将 Q/K/V 切分为 h 个头,并行学习不同的依赖关系,最后拼接输出:
ViT-Base 默认使用12 个注意力头。
2.4 前馈神经网络(FFN)
Encoder 中,注意力模块后接两层全连接网络,中间用 GELU 激活函数:
2.5 层归一化与残差连接
为解决深度网络训练退化问题,Transformer 采用Pre-LN 结构(层归一化放在模块前)+ 残差连接:
其中 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 实现方式
两种等价实现:
- 展平 Patch + 线性层;
- 用卷积层(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 结构:
- 层归一化 → 多头自注意力(12 头)→ 残差连接;
- 层归一化 → 前馈网络(隐藏层维度 4d=3072)→ 残差连接。
所有层共享相同结构,无卷积、无池化,纯注意力建模。
4.5 模块 5:分类头(MLP Head)
- 训练阶段:LayerNorm → Linear → GELU → Linear(映射到类别数);
- 推理阶段:仅用 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 代码说明
- 严格还原传统 ViT结构,无任何改进;
- 输入:224×224×3 的图像,输出:1000 类分类概率;
- 可通过修改
num_classes适配自定义数据集; - 参数量与原始 ViT-Base/16 完全一致(86M)。
7. 传统 ViT 训练策略与实验结果
7.1 数据集要求
传统 ViT 的核心特性:对数据集规模敏感
- 小数据集(如 CIFAR-10):ViT 性能不如 ResNet,缺乏 CNN 的归纳偏置;
- 大规模数据集(ImageNet-1K/21K、JFT-300M):ViT 性能超越所有 CNN,全局建模优势凸显。
7.2 预训练与微调
- 预训练:在 JFT-300M(3 亿图像)上预训练,学习通用视觉特征;
- 微调:在目标数据集(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 核心优势
- 全局感受野:自注意力直接建模全局依赖,远超 CNN 的局部感受野;
- 结构简单:无卷积、无池化,纯全连接 + 注意力,易于扩展;
- 迁移能力强:大规模预训练后,泛化性优于 CNN;
- 可扩展性好:模型越大(Large/Huge),性能越高,符合大模型规律。
8.2 核心局限性
- 小数据集性能差:缺乏 CNN 的局部归纳偏置,小数据易过拟合;
- 计算量大:自注意力复杂度为\(O(N^2)\)(N 为序列长度),高分辨率图像效率低;
- 缺乏层级特征:ViT 输出单一尺度特征,不适合检测 / 分割等密集预测任务;
- 位置编码简单:1D 位置编码无法充分利用图像 2D 空间结构。
8.3 与 CNN 的核心区别
| 特性 | CNN | 传统 ViT |
|---|---|---|
| 归纳偏置 | 局部性、平移不变性 | 无 |
| 感受野 | 局部、逐层扩大 | 全局 |
| 计算复杂度 | \(O(N)\) | \(O(N^2)\) |
| 数据依赖性 | 低 | 高(需大规模数据) |
9. 传统 ViT 的行业影响
传统 ViT 是视觉 Transformer 的开山之作,彻底改变了计算机视觉的发展方向:
- 开启了视觉 Transformer 时代,后续 Swin Transformer、DEiT、BEiT、MAE 等模型均基于 ViT 改进;
- 证明了纯注意力机制可替代 CNN 完成视觉建模,为多模态大模型(文生图、图文理解)奠定基础;
- 统一了 NLP 与视觉的架构设计,推动了通用人工智能的发展。
10. 总结
传统 Vision Transformer(ViT)是视觉领域的里程碑式模型,其核心创新是将图像转化为 Patch 序列,直接复用 Transformer Encoder 实现全局视觉建模。本文完整解析了传统 ViT 的背景、原理、网络结构、数学公式,并提供了可直接运行的 PyTorch 原生代码,严格还原了原始模型的设计思想。
传统 ViT 的核心价值不在于极致性能,而在于证明了 Transformer 在视觉领域的可行性,为后续所有视觉 Transformer 提供了基础框架。尽管存在小数据集性能差、计算量大等缺陷,但它彻底打破了 CNN 的垄断,成为现代计算机视觉的基石。
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)