手撕ViT:图像到序列的完整代码实现与原理拆解
不少初学者第一次接触 ViT 时,都会经历一个“看似懂了、一写就卡”的阶段。Transformer 论文里的公式读起来不复杂,无外乎是 Q、K、V 三个矩阵相乘,再做一次 softmax 归一化;可真要自己动手写代码,问题就出来了。尤其是从自然语言处理跑到视觉方向的 Transformer 之后,第一道坎几乎都是同一个:图像明明是二维的、像素是密集排列的,它到底怎么变成 Transformer 能够处理的序列?
这个问题背后,其实藏着整条理解链路。Patch Embedding 负责把图像翻译成序列,Transformer Encoder 负责在序列中建立全局依赖,最终的 Forward 则把数据流动的完整过程串起来。等你亲手把这三部分代码写一遍,会发现 Transformer 的核心原理并不像论文里写得那么抽象,它本质上就是在回答三个问题:数据以什么形状进来,中间经过了哪些形状变换,最终以什么形状输出。
这篇文章我用代码把这条链路完整走一遍。没有花哨的封装,只讲最直接的实现思路。
1. 先搞清楚 Transformer 真正改变的,是“建模距离”的方式
很多人把 Transformer 简单理解成“一个更厉害的神经网络层”,这种理解不能说错,但很容易让人忽略它真正改变的东西。
1.1 从 RNN 和 CNN 的局限说起
在 Transformer 出现之前,序列建模主要靠 RNN 一族,图像建模主要靠 CNN。RNN 的特点是逐步处理数据,当前时刻的输出依赖上一个时刻的隐藏状态。这种方式天然适合时间序列,但也有一个绕不开的问题:信息要一步步传递,距离越远,信息损耗越大。虽然 LSTM、GRU 通过门控机制缓解了长期依赖问题,但本质上仍是串行路径。
CNN 则是通过卷积核在局部区域滑动。卷积核越大,感受野越大,但大卷积核的计算成本会快速上升。即便通过堆叠层数来扩大感受野,底层的信息要传到高层,也要经过很多层。
换句话说,在 Transformer 出现以前,视觉和序列模型的核心矛盾都是同一个:如何用可控的计算代价,让不同位置的信息直接发生交互。
1.2 自注意力机制的实质:信息直接通信
Transformer 给出的答案是自注意力机制。它不再依赖“一步一步传递”,而是让序列中的每个位置,都能直接跟其他所有位置计算相关性。相关性高的信息被加权融合,相关性低的自然被忽略。
这个设计在视觉任务里的意义尤其明显。对 CNN 来说,一张图中相距很远的两个像素,要建立起联系,需要经过很多层卷积;对 Transformer 来说,这就是一次注意力计算的事。代价是计算复杂度会从 CNN 的线性级别上升到序列长度的平方级别,这也是为什么 ViT 后面会出现各种改进,比如 Swin Transformer 用窗口注意力来限制计算范围。
1.3 所以,Transformer 的核心不是“注意力公式”
我见过不少同学把注意力公式背得很熟,但问他“为什么图像要用 Patch 而不是像素”,就答不上来了。原因就在于,他把注意力公式当成了 Transformer 的核心,而忽略了注意力只是一个工具。
Transformer 真正核心的设计,是把数据表示成序列,然后让序列中每个元素都能动态地聚合全局信息。注意力公式只是实现这个目标的手段。理解了这一点,再看 ViT 的 Patch Embedding,你就能明白为什么它会是整个模型的第一块拼图。
与其把注意力公式背下来,不如先想清楚一个问题:你的输入数据是什么形状,你的输出希望是什么形状,中间每一步有没有把信息保留下来。
2. Patch Embedding:图像是怎么被翻译成序列的
ViT 的第一步,就是 Patch Embedding。它的输入是一张图片,输出是一个序列。
2.1 为什么不能直接逐像素做序列
从纯理论角度看,图像变成序列最简单的方式,是把每个像素当作序列中的一个元素。比如一张 224×224 的彩色图片,展开后会得到 150528 个元素。Transformer 的自注意力计算复杂度是序列长度的平方,也就是大约 226 亿次相关性计算。这个数字在当前的硬件条件下,几乎不可能落地。
于是 ViT 的作者想到一个折中方案:把图片切成一个个小块。每个小块叫一个 Patch。经典配置下,把 224×224 的图片切成 16×16 的 Patch,会得到 14×14=196 个 Patch,序列长度直接从 15 万级别降到了 196。每个 Patch 内部再用线性变换或卷积映射成一个向量。
这就是 Patch Embedding 的核心逻辑:先降序列长度,再保留局部信息。
2.2 用卷积实现 Patch Embedding,才是正确姿势
关于 Patch 的实现,有一个很容易踩的误区。很多人看到“切块”,第一反应是写循环,把图片按坐标切成小块,再逐个映射。这种写法不是不行,但效率很低,而且不容易利用 GPU 的并行能力。
更常见的做法,是用一个卷积核大小等于 Patch 大小、步长也等于 Patch 大小的卷积层。比如 Patch 大小是 16,就用 kernel_size=16、stride=16 的卷积。这样做的好处是:
- 卷积本身就是“局部区域映射到向量”的天然实现。
- 每个 Patch 之间不会重叠。
- 一次前向传播就能完成所有 Patch 的映射。
用代码写出来是这样的:
import torch import torch.nn as nn class PatchEmbedding(nn.Module): 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 self.num_patches = (img_size // patch_size) ** 2 self.proj = nn.Conv2d( in_channels, embed_dim, kernel_size=patch_size, stride=patch_size, ) self.norm = nn.LayerNorm(embed_dim) def forward(self, x): # x 形状: [B, 3, 224, 224] B, C, H, W = x.shape x = self.proj(x) # [B, embed_dim, 14, 14] x = x.flatten(2) # [B, embed_dim, 196] x = x.transpose(1, 2) # [B, 196, embed_dim] x = self.norm(x) return x注意:卷积输出形状先变成[B, embed_dim, H/patch, W/patch],再通过flatten(2)和transpose(1, 2)变成[B, num_patches, embed_dim]。这个[B, 序列长度, 特征维度]的形状,才是 Transformer Encoder 需要的输入格式。
2.3 class token 和位置编码:两个容易被忽略的细节
Patch Embedding 之后,ViT 还会做两件事:拼上一个 class token,再加一个位置编码。
class token 是一个可学习的向量,形状是[1, 1, embed_dim]。它会被拼到 Patch 序列的最前面。为什么要加它?因为分类任务最后需要从序列中提取一个全局表示。虽然也可以对所有 Patch 的表示做平均池化,但 ViT 选择了单独学一个 class token,让模型自己决定它应该聚合哪些信息。实践证明,这种方式比简单平均池化效果更好。
位置编码则是给每个 Patch 加上位置信息。Transformer 本身不像 RNN 那样有天然的顺序概念,如果不加位置编码,两个 Patch 互换位置后模型的输出不会变,这在图像任务里是不合理的。位置编码可以是固定的正弦余弦编码,也可以是可学习参数。ViT 使用的是可学习位置编码,矩阵形状是[1, num_patches + 1, embed_dim],加号是因为还有一个 class token 的位置。
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))初始化用全零向量,后续在训练中会不断更新。这里有个细节:class token 和位置编码的大小维度必须和 Patch Embedding 的输出一致,否则后面相加会报错。
2.4 形状变化的完整顺序
从输入到进入 Transformer Block 之前,数据的形状变化可以用一张表总结。
| 阶段 | 张量形状 | 说明 |
|---|---|---|
| 原始图像 | [B, 3, 224, 224] | B 为 batch size |
| Patch Embedding | [B, 196, 768] | 196 个 Patch,每个映射成 768 维向量 |
| 拼接 class token | [B, 197, 768] | 序列开头多一个全局表示向量 |
| 加位置编码 | [B, 197, 768] | 每个位置加上可学习位置向量 |
到这一步,图像已经变成 Transformer 能处理的序列了。接下来就进入 Transformer Encoder 部分。
3. 手撕 Transformer Encoder:自注意力、残差和 MLP
Transformer Encoder 是整个模型的特征提取主体。ViT 通常会堆叠 12 层 Transformer Block,每一层的结构相同,但参数独立。
3.1 多头自注意力:为什么是“多头”
自注意力最核心的公式是:
Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V多头的意思,是把特征维度拆成多个子空间,每个子空间独立计算注意力,最后再合并。比如 768 维的特征,拆成 12 个 head,每个 head 处理 64 维。
多头的好处在于,不同的 head 可以关注不同粒度的关系。有的 head 可能关注颜色相近的区域,有的 head 可能关注位置相邻的区域,有的 head 可能关注语义相关的区域。如果只有一个 head,这些模式只能混合在一起,表达力会受限。
3.2 手撕多头自注意力代码
class MultiHeadSelfAttention(nn.Module): def __init__(self, embed_dim, num_heads=8, dropout=0.1): super().__init__() self.num_heads = num_heads self.head_dim = embed_dim // num_heads self.scale = self.head_dim ** -0.5 self.qkv = nn.Linear(embed_dim, embed_dim * 3, bias=True) self.attn_drop = nn.Dropout(dropout) self.proj = nn.Linear(embed_dim, embed_dim) self.proj_drop = nn.Dropout(dropout) def forward(self, x): # x: [B, N, C] B, N, C = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim) qkv = qkv.permute(2, 0, 3, 1, 4) q, k, v = qkv[0], qkv[1], qkv[2] attn = (q @ k.transpose(-2, -1)) * self.scale attn = attn.softmax(dim=-1) attn = self.attn_drop(attn) x = (attn @ v).transpose(1, 2).reshape(B, N, C) x = self.proj(x) x = self.proj_drop(x) return x, attn这段代码里有几个关键步骤,值得拆开看。
第一步,self.qkv(x)把输入映射成 Query、Key、Value 三个向量。因为后面还要拆成多头,所以输出维度是embed_dim * 3,一次完成三个映射。
第二步,reshape和permute是整段代码中比较容易绕晕的地方。原始形状是[B, N, 3, num_heads, head_dim],通过permute(2, 0, 3, 1, 4)后变成[3, B, num_heads, N, head_dim]。这样拆解后,q、k、v 就各自独立了。
第三步,q @ k.transpose(-2, -1)计算每个 token 与其他 token 的相似度。除以sqrt(head_dim)是为了防止点积结果过大导致 softmax 落到饱和区。这一步在论文里叫缩放点积注意力。
第四步,softmax把相似度变成和为 1 的权重,再通过attn @ v聚合信息。
3.3 Transformer Block:残差是标配,不是可选项
多头自注意力完成后,还需要经过一个 MLP 层,并在每层前后加上残差连接和 LayerNorm。整个 Transformer Block 的结构如下:
class TransformerEncoderBlock(nn.Module): def __init__(self, embed_dim, num_heads, mlp_ratio=4.0, dropout=0.1): super().__init__() self.norm1 = nn.LayerNorm(embed_dim) self.attn = MultiHeadSelfAttention(embed_dim, num_heads, dropout) self.norm2 = nn.LayerNorm(embed_dim) hidden_dim = int(embed_dim * mlp_ratio) self.mlp = nn.Sequential( nn.Linear(embed_dim, hidden_dim), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden_dim, embed_dim), nn.Dropout(dropout), ) def forward(self, x): x = x + self.attn(self.norm1(x))[0] x = x + self.mlp(self.norm2(x)) return x为什么用残差连接?深层网络在训练时会出现梯度消失或退化问题,残差连接为梯度提供了一条从输出直接回到输入的捷径。没有残差连接的 Transformer,在层数加深时训练难度会明显增大。
为什么 LayerNorm 放在注意力之前?这是 Transformer 在后续实践中的一个重要调整。相比原始论文的后置 LayerNorm,前置 LayerNorm 在训练稳定性上更好。ViT 使用的是前置 LayerNorm,也就是 Pre-Norm 结构。
MLP 的作用也不可忽视。自注意力主要做信息交互和聚合,MLP 则是对每个位置单独做非线性变换。交互和非线性变换交替进行,模型才能表达更复杂的特征。
4. 完整 Forward:数据在 ViT 里到底是怎么流动的
有了前面的模块,现在可以把整条链路组装起来了。这个阶段的目标,不是再新增一个复杂模块,而是把所有组件拼成一个完整的 ViT 模型。
4.1 组装完整 ViT 模型
class VisionTransformer(nn.Module): def __init__(self, img_size=224, patch_size=16, in_channels=3, num_classes=1000, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4.0, dropout=0.1): super().__init__() self.patch_embed = PatchEmbedding(img_size, patch_size, in_channels, embed_dim) num_patches = self.patch_embed.num_patches self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim)) self.pos_drop = nn.Dropout(dropout) self.blocks = nn.Sequential(*[ TransformerEncoderBlock(embed_dim, num_heads, mlp_ratio, dropout) for _ in range(depth) ]) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, num_classes) def forward(self, x): B = x.shape[0] x = self.patch_embed(x) cls_tokens = self.cls_token.expand(B, -1, -1) x = torch.cat((cls_tokens, x), dim=1) x = x + self.pos_embed x = self.pos_drop(x) x = self.blocks(x) x = self.norm(x) cls_out = x[:, 0] logits = self.head(cls_out) return logits这段代码的 forward 流程,就是 ViT 的完整“内容”:
- 输入
[B, 3, 224, 224]。 - Patch Embedding 转为
[B, 196, 768]。 - 拼上 class token,变成
[B, 197, 768]。 - 加上位置编码,形状不变。
- 经过 12 层 Transformer Block,每层的输入输出形状都是
[B, 197, 768]。 - 经过 LayerNorm,取 class token 位置的输出。
- 送入分类头,得到
[B, num_classes]的 logits。
4.2 为什么分类只用 class token 而不是全部 token
这是初学 ViT 时最容易产生疑惑的地方。既然所有 Patch 都经过多层 Transformer 计算,为什么分类时只取第一个位置(class token)的输出?
关键在于,class token 在自注意力计算中,可以关注到所有 Patch 的信息。经过多层堆叠后,它的表示已经聚合了全局信息。取它作为分类特征,相当于让模型自己决定如何汇总全场信息。
另一种做法是取所有 token 的平均池化,但 ViT 的实验中,class token 的表现更好。这也是为什么代码里x[:, 0]而不是x.mean(dim=1)。
4.3 一次前向传播的形状变化
如果只看形状变化,ViT 的 forward 流程可以概括成:
[B, 3, 224, 224] → Patch Embedding → [B, 197, 768] → Transformer Block × 12 → [B, 197, 768] → LayerNorm → [B, 197, 768] → 取 class token → [B, 768] → Linear → [B, num_classes]这里有一个很值得体会的设计:Transformer Block 不会改变张量的形状。它的作用不是让特征变小,而是在保持形状的前提下逐层优化特征表示。降维只发生在最后的分类头。
这种设计带来的好处是,网络可以设计得很深而不必担心特征丢失;坏处是计算量会比较大。尤其是序列长度较长时,自注意力的平方复杂度会让训练变得很慢。
5. 跑通之后,真正麻烦的是调试和排查
手撕完代码,只是第一步。真正到了训练和部署阶段,你会遇到各种问题。这里我列一份在 ViT 调试中比较常见的排查链路。
5.1 最常见的错误集中在形状不匹配
第一次跑模型时,报错最多的位置几乎都在形状变换上。常见的有:
permute之后维度顺序搞混,导致q、k、v形状错误。torch.cat拼接 class token 时,维度没对齐。pos_embed的序列长度和输入 Patch 数量不一致。- 图片尺寸不是 Patch 大小的整数倍,导致
num_patches计算错误。
我通常建议的做法是,先构造一个极小样本,比如batch_size=2, 3×224×224,把每个模块的输出形状打印出来,一步步核对。
x = torch.randn(2, 3, 224, 224) model = VisionTransformer(img_size=224, patch_size=16, num_classes=10) out = model(x) print(out.shape) # 期望是 [2, 10]如果输出形状不对,从 Patch Embedding 开始逐层打印,定位问题只需要几分钟。
5.2 排查顺序:输入、形状、参数、资源
如果训练时 loss 不下降、精度异常,或者出现显存溢出,我建议按这个顺序排查:
第一步,检查输入。图片是否正确做了归一化、resize、通道顺序对不对。视觉任务的很多“玄学”问题,最后都出在预处理上。
第二步,检查形状。用一张小图跑一遍前向传播,确认每一层的输出形状符合预期。
第三步,检查参数。学习率、batch size、warmup 策略。ViT 通常需要较小的学习率和较长的训练轮数,这和 CNN 的直觉不太一样。如果没有预训练权重,从零训练一个小规模的 ViT 往往不容易收敛,这是正常现象。
第四步,检查资源。显存溢出时,先把 batch size 调小,或把图片尺寸调小。ViT 对显存的消耗通常比同规模 CNN 更大,因为自注意力的中间计算矩阵会被保存。
5.3 可视化是最直接的理解工具
很多论文里会用注意力图来展示模型关注到了哪些区域,这是检验模型是否学到有效特征的好办法。实现上,只需要在前向传播时把每层的attn矩阵保存下来,然后对 class token 的注意力权重做可视化。
# 在 MultiHeadSelfAttention 的 forward 里返回 attn # 在前向传播时收集每一层的 attention attn_maps = [] for block in model.blocks: output, attn = block(x) attn_maps.append(attn)如果模型训练正常,class token 对前景区域的注意力权重通常会更高;如果注意力图很散乱,说明模型还没有学到有效的全局特征。
注意:不要一开始就在完整数据集上跑 ViT。先用几十张图片过拟合一个小批量,确认 loss 能降下去,再去放大规模。这能帮你快速区分“模型写错了”和“训练不够充分”这两类问题。
6. 从“手撕”到工程化,别忘了适用边界
代码手撕最大的价值,不是让你背住 ViT 的结构,而是建立起对数据流动的直觉。但从“能跑通”到“能在项目中用好”,中间还隔着一段工程化的距离。
6.1 教学版和生产版的差距
上面给的代码是教学简化版,突出的是核心链路。真正要应用到生产环境,还需要补上许多细节:
- 预训练权重:从零训练 ViT 通常很慢,也更难收敛。实践中更常见的做法,是加载在 ImageNet 或更大数据集上预训练好的权重,再做迁移学习。
- 学习率策略:ViT 对优化器很敏感。AdamW 是常用选择,学习率一般从 1e-4 到 1e-3 之间开始调,配合 warmup 和余弦退火。
- 数据增强:ViT 相比 CNN 需要更强的正则化。Mixup、CutMix、RandAugment 在 ViT 训练中都是常见配置。
- 推理优化:如果要做部署,ONNX 导出、TensorRT、INT8 量化都需要额外处理。自注意力的动态形状问题,在做推理优化时会比较麻烦。
这也是很多初学同学容易产生的误解:以为手撕完代码就掌握了 ViT。实际上,手撕只是建立了模型骨架的直觉,工程化才是让模型真正落地的关键。
6.2 什么时候该用 ViT,什么时候别用
ViT 最适合的场景是数据量足够大、任务复杂度较高、需要建模全局依赖的视觉任务。比如大规模图像分类、目标检测、语义分割等。在这些任务上,ViT 的全局感受野优势能充分发挥。
但如果你面临的是这样几种场景,可能需要重新考虑:
- 数据量很小:几千张图片之下,ViT 很容易过拟合,最直接的表现是训练集精度很高、验证集精度很低。此时轻量 CNN 或 Swin 这类窗口注意力模型可能更合适。
- 需要在移动端或边缘设备部署:ViT 的参数量和计算量都比较大,尤其自注意力的内存占用很高。虽然有小体积 ViT 变体,但整体优化难度还是比 CNN 高不少。
- 强实时性任务:单帧推理时间要求极低时,CNN 的成熟推理方案通常更容易满足要求。
6.3 一个可复用的学习框架:手撕 → 改造 → 工程化
如果你想把 Transformer 的学习延伸到更多场景,我建议遵循一个三层路径:
第一层:手撕核心链路。从 Patch Embedding、注意力机制、Transformer Block 到完整前向传播,把最基础的形状变化搞明白。这一层只要求能跑通,不追求性能。
第二层:做改造实验。把 patch_size 从 16 改成 8,看看参数量和计算量怎么变化;把 num_heads 从 8 改成 4,看看精度和速度的差异;把位置编码从可学习改成正弦编码,对比训练曲线。改造实验的价值,是让你理解每个超参数的影响,而这比背结构更有效。
第三层:接入工程化框架。使用成熟的深度学习库或模型库,加载预训练权重,做数据增强、分布式训练、模型导出。这一层的目的是把模型放进真实项目,让它稳定产出结果。
这三层不是替代关系,而是递进关系。跳过第一层直接进入工程化,遇到问题会缺少拆解能力;一直停留在第一层,又会在真实项目里寸步难行。
如果能把“手撕”当成理解工具,而不是最终目的,那这篇文章带来的价值就不会停留在“我照着代码敲了一遍”,而是沉淀成一种拆解复杂模型的方法。下次再遇到新的网络结构,你也能更快地摸清它的数据流、形状变化和关键设计。
