从零手写Vision Transformer:PyTorch实现ViT核心原理与代码
在深度学习模型层出不穷的今天,Transformer 早已不是 NLP 领域的专属名词。无论是 ChatGPT 背后的 GPT 系列,还是图像分类中表现强劲的 ViT(Vision Transformer),其核心架构都指向同一个基础模块——Transformer。很多初学者在看完 Attention Is All You Need 论文后仍然一脸茫然,代码更是不知道从何下手。本文将围绕 Transformer 最核心的原理展开,重点拆解 ViT 中的 Patch Embedding、Forward 前向传播流程,并用 PyTorch 手写一套完整的可运行代码,带大家从零实现一个 Vision Transformer。全文既有概念解释,也有逐步推导,还有完整实战代码和常见坑点盘点,适合算法初学者、CV 开发者以及对大模型底层机制感兴趣的读者。
1. 背景与核心概念:Transformer 为什么能通吃 NLP 和 CV
1.1 从 NLP 到 CV:Transformer 的跨界之路
Transformer 最初由 Vaswani 等人在 2017 年提出,核心创新是 Self-Attention(自注意力)机制。它与传统的 RNN、LSTM 最大的区别在于:RNN 必须按时间步顺序处理序列,而 Transformer 可以并行地处理整个序列中的所有元素。
这一特性使得 Transformer 在长文本建模上拥有天然优势。后来研究者发现,既然图像可以看作是由像素点组成的“序列”,那 Transformer 是不是也能用到图像上呢?
于是 2020 年,Google 团队提出了 Vision Transformer,简称 ViT。ViT 将图像切分成固定大小的 Patch(图像块),然后像处理文字 token 一样处理这些 Patch,并通过 Transformer Encoder 完成图像分类任务。这个思路打破了 CNN 在视觉领域长期垄断的局面。
1.2 核心概念拆解:Self-Attention、QKV、Multi-Head
要理解 Transformer,必须先理解 Self-Attention。Self-Attention 的作用是计算序列中每个元素与其他元素之间的相关程度,从而让模型动态地关注重要信息。
在 Self-Attention 中,每个输入 token 都会生成三个向量:
- Query(查询向量):表示当前元素“想找什么”。
- Key(键向量):表示当前元素“能提供什么”。
- Value(值向量):表示当前元素“实际的内容”是什么。
Attention 权重通过 Query 与 Key 的点积计算相似度,再经过 Softmax 归一化,最后与 Value 加权求和。
具体公式如下:
其中 d_k 是 Q 和 K 的向量维度。除以 sqrt(d_k) 是为了防止点积过大导致 Softmax 梯度消失。
Multi-Head Attention(多头注意力)则是将 Q、K、V 分别投影到多个子空间中进行注意力计算,最后拼接起来。这样做的好处是让模型从不同维度关注信息,例如一个头关注局部纹理,另一个头关注全局形状。
1.3 ViT 的整体架构
ViT 的结构可以分为以下几个关键模块:
- Patch Embedding:将图像切块并映射成向量序列。
- Position Embedding:为每个 Patch 添加位置信息。
- Class Token:分类令牌,最终用于图像分类。
- Transformer Encoder:由多层 Multi-Head Self-Attention 和 MLP 组成。
- Classification Head:输出类别概率。
ViT 的完整前向流程可以概括为:
- 输入图像形状为 (B, C, H, W)。
- 将图像划分为 P×P 大小的 Patch,得到 N 个 Patch。
- 对每个 Patch 做线性映射,得到 Patch Embedding。
- 拼接 Class Token,并叠加位置编码。
- 输入 Transformer Encoder 进行特征提取。
- 取出 Class Token 对应的输出,送入分类头得到结果。
2. 环境准备与版本说明
在开始写代码之前,先说明一下本文的示例环境。由于不同机器环境可能不同,这里给出的版本是本文验证时的常用组合,读者需要根据实际情况调整。
2.1 运行环境
- 操作系统:Windows 10 / Ubuntu 20.04 / macOS 均可
- Python:3.8 及以上
- PyTorch:1.10 及以上(推荐 2.0)
- torchvision:0.11 及以上
2.2 安装依赖
如果你有 GPU,建议安装 CUDA 版本的 PyTorch,具体安装命令请参考 PyTorch 官网。没有 GPU 也没关系,本文示例在 CPU 上也能运行,只是训练速度会慢一些。
2.3 示例项目结构
为了便于阅读,我们先规划一下项目结构:
本文的核心代码主要在 vit_model.py 中,通过逐块实现来拆解 ViT 的完整前向流程。
3. 手撕 Transformer 核心模块
现在进入正题。我们从零开始实现一个 Vision Transformer,不求代码精简,而是追求每一步都可解释、可运行。
3.1 Patch Embedding:图像如何变成 token
Patch Embedding 是 ViT 区别于传统 Transformer 的关键步骤。它的作用是把形状为 (B, C, H, W) 的图像转换为形状为 (B, N, D) 的向量序列,其中:
- B:Batch Size,批次大小。
- C:图像通道数,RGB 图像为 3。
- H、W:图像的高和宽。
- P:Patch 大小。
- N = (H/P) × (W/P):Patch 数量。
- D:Embedding 维度。
我们先把图像切成多个 Patch。假设输入图像为 224×224,Patch 大小为 16×16,那么一共得到:
也就是说,一张 224×224 的图像会被切成 196 个 Patch,每个 Patch 的大小是 3×16×16。
把每个 Patch 拉平后,就变成一个长度为 768(3×16×16)的向量。但我们需要的是一个固定维度的 Embedding,所以还要经过一层线性映射,把维度从 768 映射到 D,比如 768 或 512。
下面用 PyTorch 实现 Patch Embedding。
关键点解释:
- 我这里用卷积实现 Patch Embedding,原因是 Conv2d 的 kernel_size 和 stride 都等于 patch_size 时,每一个卷积位置正好对应一个 Patch,且互不重叠。
- flatten(2) 是从第 2 个维度开始展平,也就是把 H/P 和 W/P 两个维度合并成 N。
- 输出的形状为 (B, N, embed_dim),其中 N 就是 Patch 数量。
如果你不想用卷积,也可以用 unfold 或手动切片实现,但卷积方案最简洁,也是 PyTorch 官方 ViT 实现中使用的方式。
3.2 Class Token 和 Position Embedding
得到 Patch Embedding 后,我们还需要做两件事:添加 Class Token 和 Position Embedding。
为什么需要 Class Token?
在 BERT 中,输入序列开头会加一个特殊的 [CLS] token,用于汇聚整个序列的信息。ViT 借鉴了这个设计,在 Patch Embedding 序列开头也拼接一个可学习的 Class Token。它的作用相当于图像级别的全局表示,在最终输出时,我们只取 Class Token 对应的输出向量送入分类头。
Class Token 的初始化值是一组可学习的参数,初始值可以随机,也可以设置为零向量。
为什么需要 Position Embedding?
Transformer 本身不具备顺序感,它不像 CNN 那样通过卷积核天然感知局部空间位置。如果不加位置编码,模型会把所有 Patch 当作无序集合,这显然不符合图像的空间结构。因此,需要给每个 Patch 添加位置信息。
Position Embedding 有两种常见形式:
- 固定的正弦余弦位置编码。
- 可学习的位置编码。
ViT 原论文使用的是 1D 可学习位置编码,也就是直接初始化 N+1 个位置向量(N 个 Patch + 1 个 Class Token),作为可学习参数。实验证明,1D 可学习位置编码在 ViT 中已经足够,且实现简单。
代码实现如下:
在实际实现中,位置编码通常直接集成在 ViT 主类中,不需要单独抽取成一个模块。上面拆出来是为了方便讲解。
3.3 Multi-Head Self-Attention 的实现
接下来是 Transformer 的核心组件:多头自注意力。我们先从单头自注意力入手,再扩展到多头。
自注意力的计算流程:
- 对输入 x 做三个线性变换,得到 Q、K、V。
- 将 Q、K、V 按头数拆分。
- 计算缩放点积注意力。
- 拼接所有头的结果。
- 经过输出线性投影。
代码如下:
这里有几个容易出错的地方:
- reshape 的顺序很关键。我们先把 [B, N, D] 变成 [B, N, num_heads, head_dim],再用 transpose(1, 2) 交换维度,变成 [B, num_heads, N, head_dim]。如果顺序搞反,注意力计算会出错。
- 除以 head_dim 的平方根是为了控制分数尺度。如果 head_dim 是 64,那么点积最大值能达到 64 这个量级,除以 8 后可以稳定 Softmax 的梯度。
- 最后 reshape 时,因为我们已经把维度调整成了 [B, N, num_heads, head_dim],所以用 reshape(B, N, D) 可以直接拼接。
3.4 MLP 和 Dropout 层
Transformer Encoder 中的 MLP 模块通常由两个全连接层组成,中间夹一个 GELU 激活函数。原论文使用 GELU,不过在简单实现中也可以使用 ReLU 替换。
MLP 的作用是对注意力输出做进一步的非线性特征变换,增强模型的表达能力。
3.5 Transformer Encoder 层
有了 MSA 和 MLP,我们就可以组装一个完整的 Transformer Encoder Block 了。每个 Encoder Block 的结构为:
- Layer Norm
- Multi-Head Self-Attention
- 残差连接
- Layer Norm
- MLP
- 残差连接
为什么要有残差连接和 Layer Norm?
残差连接(Residual Connection)可以缓解深层网络的梯度消失问题。在 Transformer 中,每一层子层的输出都加上它的输入,形成跳跃连接。即使网络很深,梯度也能通过捷径传播到浅层。
Layer Norm 则是将每个 token 的特征向量归一化到均值为 0、方差为 1 的分布,使得训练更加稳定。与 Batch Norm 不同,Layer Norm 是在特征维度上归一化,而不是在批次维度上。
为什么 Layer Norm 要放在注意力前面?
“Pre-LN” 结构指在子层之前先做 Layer Norm,这是近年来 ViT 和 GPT 等模型的主流选择。相比原始 Transformer 的 Post-LN,Pre-LN 在深层网络中更稳定,收敛也更快。
代码实现如下:
3.6 分类头
分类头(Classification Head)的作用是将 Class Token 的输出向量映射为类别概率。通常由一个 Layer Norm 和一个线性层组成。如果是预训练任务,也可能是两层 MLP。
4. 完整实战:从零构建 Vision Transformer
现在,我们将上述模块组装成一个完整的 ViT 模型。为了便于读者阅读理解,这里不使用 nn.Sequential 堆层,而是显式地写出每个步骤。
4.1 完整 ViT 模型定义
4.2 前向传播流程详解
上面的 forward 函数是整个 ViT 的 Forward 全流程。我们来逐步梳理一下维度变化:
- 输入形状:(B, 3, 224, 224)
- Patch Embedding 后:(B, 196, 768)
- 拼接 Class Token 后:(B, 197, 768)
- 添加位置编码后:(B, 197, 768)
- 经过 12 层 Encoder 后:(B, 197, 768)
- 取 Class Token 并分类后:(B, num_classes)
每一步的维度变化清晰明了。其中 Class Token 的输出位置始终是 0 号位置,所以 x[:, 0] 取出的就是 Class Token 对应的特征。
4.3 模型测试
定义好模型后,我们先创建一个随机输入来测试前向传播是否正常:
预期输出:
输出形状为 (4, 10),表示 4 张图像分别得到 10 个类别的 logits 分数。
4.4 在小型数据集上验证训练流程
为了让读者能够完整地跑通训练流程,我们使用 torchvision 自带的 CIFAR-10 数据集做一个小型分类任务。CIFAR-10 图像大小为 32×32,因此需要调整 patch_size。
为了加快训练速度,这里我们使用一个小型 ViT 配置。注意,这个配置不是原论文的 ViT-Base,而是适用于小数据集的精简版。
4.5 推理脚本示例
训练完成后,我们可以编写一个简单的推理脚本,对单张图片进行分类。
5. 常见问题与排查思路
在实际编写和运行 ViT 代码时,经常会遇到一些报错或效果不佳的情况。下面整理几个高频问题。
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
| 维度不匹配报错 | image_size 不能被 patch_size 整除 | 调整 patch_size,或添加 padding |
| 注意力输出 NaN | 学习率过大或初始化不当 | 降低学习率,使用 warmup |
| 模型不收敛 | 数据集太小,Transformer 缺少归纳偏置 | 使用预训练权重或数据增强 |
| 训练速度慢 | ViT 计算量大且无局部性 | 减小 patch_size、embed_dim、depth |
| 精度远低于 CNN | 超参数未调优 | 参考原论文 lr、weight decay、warmup |
| GPU 显存不足 | Batch Size 过大 | 减小 batch_size 或降低分辨率 |
5.1 维度不匹配问题
如果你修改了 image_size 或 patch_size,最容易出现的问题就是维度对不上。例如:
224 / 32 = 7,7×7 = 49 个 Patch,可以整除,没问题。但如果设置 image_size=224,patch_size=48,224 / 48 不是整数,就会报错。代码里已经加了 assert 断言,因此你会在运行时立刻发现问题。
另外,如果你输入的是灰度图(单通道),需要将 in_channels 参数改为 1,或者提前将灰度图转成三通道。
5.2 模型不收敛的排查顺序
很多读者第一次训练 ViT 时发现 accuracy 一直在 10% 附近徘徊,这通常意味着模型没有学到任何有效信息。建议按以下顺序排查:
- 检查数据归一化是否正确。
- 检查标签是否从 0 开始连续编号。
- 检查学习率是否过大或过小。
- 检查损失函数是否在下降。
- 尝试减小 depth 和 embed_dim,用一个小模型跑通流程。
- 加入 warmup,因为 ViT 对学习率比较敏感。
5.3 Overfitting 问题
ViT 在小数据集上非常容易过拟合,因为 Transformer 参数多且缺少 CNN 的归纳偏置。解决过拟合的常用方法包括:
- 增加数据增强:RandomCrop、RandomHorizontalFlip、CutMix、Mixup。
- 增大 Dropout 率。
- 使用 Weight Decay 正则化。
- 使用预训练的 ViT 权重做微调。
- 减小模型规模。
6. 最佳实践与工程建议
6.1 从预训练模型开始
对于实际任务,强烈建议不要从头训练 ViT。ViT 需要海量数据才能发挥出与 CNN 相当甚至更好的效果。Google 原论文中,ViT-Large 是在 JFT-300M 数据集上预训练的。如果没有足够的数据,从头训练 ViT 往往效果不如 ResNet。
因此在工程落地时,优先使用 torchvision 或 timm 库提供的预训练权重,然后在自己的数据集上微调。
6.2 使用 timm 库
timm(PyTorch Image Models)是一个非常优秀的视觉模型库,内置了大量预训练 ViT 变体,例如 ViT-Tiny、ViT-Small、ViT-Base、DeiT、Swin Transformer 等。
使用 timm 加载 ViT 非常简单:
timm 内部已经实现了最优的模型结构、默认初始化和数据预处理细节,我们只需要关注业务逻辑即可。
6.3 学习率设置
ViT 对优化器配置比较敏感,推荐使用 AdamW 优化器。在微调场景下,学习率通常设置在 1e-5 到 5e-4 之间。如果是从头训练,学习率可以适当增大,但必须配合 warmup。
Warmup 的基本思路是让学习率在前几个 epoch 从 0 线性增长到目标值,然后按余弦曲线衰减。这样可以避免模型在初始阶段因梯度异常而震荡。
6.4 推理优化
ViT 的推理速度相比同精度 CNN 可能不占优势,尤其是在 CPU 上。如果需要在生产中部署,可以考虑以下优化:
- 模型量化(PyTorch 的 quantization)。
- ONNX 导出。
- TensorRT 加速。
- 知识蒸馏到小模型。
6.5 日志与实验管理
在实际实验中,建议使用 tensorboard 或 wandb 记录训练曲线。尤其是当你在调参时,没有日志几乎不可能定位问题。至少需要记录:
- 训练 loss。
- 验证 accuracy。
- 学习率变化。
- 梯度范数。
- 验证集误分类样本。
7. 总结与学习路线
本文从 Transformer 的核心原理出发,详细拆解了 Self-Attention、Multi-Head Attention、Layer Norm、残差连接等关键概念,并针对视觉任务重点讲解了 ViT 中的 Patch Embedding、Class Token、Position Embedding 与完整 Forward 流程。通过手写代码,我们完成了从 Patch Embedding 到 Transformer Encoder 再到分类输出的全链路实现,并在 CIFAR-10 上给出了可运行的训练和推理示例。
如果你正在学习 Transformer 相关技术,建议按照下面的路线继续深入:
- 阅读原论文 Attention Is All You Need,重点理解 Multi-Head Attention 的公式推导。
- 阅读 ViT 原论文 An Image is Worth 16x16 Words,理解视觉模型如何“序列化”。
- 复现完基础 ViT 后,学习 Swin Transformer、DeiT、ConvNeXt 等改进模型,理解它们是如何解决 ViT 计算量大、需要海量数据等问题的。
- 深入学习位置编码的各种变体,理解为什么 2D Position Embedding 在多数任务中不如 1D 简单直接。
- 结合源码阅读建议使用 PyTorch 官方 timm 库,对比自己的实现和官方实现的差异,找出可以优化的点。
希望本篇文章对大家有所帮助。如果你在实现过程中遇到问题,可以先回顾一下 forward 的维度变化,看看是不是在某个 reshape 或 transpose 环节出现了维度顺序问题。代码里的每一步维度变化都是一条完整的链路,搞清楚每个维度的含义,ViT 便不会再是一个看不懂的黑盒。