FEATURED · 精选文章

手撕ViT:图像到序列的完整代码实现与原理拆解

发布时间 / 2026/8/30 17:10:13
来源 / 创域科博编辑部
栏目 / 资讯中心
手撕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×14196 个 Patch序列长度直接从 15 万级别降到了 196。每个 Patch 内部再用线性变换或卷积映射成一个向量。这就是 Patch Embedding 的核心逻辑先降序列长度再保留局部信息。2.2 用卷积实现 Patch Embedding才是正确姿势关于 Patch 的实现有一个很容易踩的误区。很多人看到“切块”第一反应是写循环把图片按坐标切成小块再逐个映射。这种写法不是不行但效率很低而且不容易利用 GPU 的并行能力。更常见的做法是用一个卷积核大小等于 Patch 大小、步长也等于 Patch 大小的卷积层。比如 Patch 大小是 16就用 kernel_size16、stride16 的卷积。这样做的好处是卷积本身就是“局部区域映射到向量”的天然实现。每个 Patch 之间不会重叠。一次前向传播就能完成所有 Patch 的映射。用代码写出来是这样的import torch import torch.nn as nn class PatchEmbedding(nn.Module): def __init__(self, img_size224, patch_size16, in_channels3, embed_dim768): 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_sizepatch_size, stridepatch_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 sizePatch Embedding[B, 196, 768]196 个 Patch每个映射成 768 维向量拼接 class token[B, 197, 768]序列开头多一个全局表示向量加位置编码[B, 197, 768]每个位置加上可学习位置向量到这一步图像已经变成 Transformer 能处理的序列了。接下来就进入 Transformer Encoder 部分。3. 手撕 Transformer Encoder自注意力、残差和 MLPTransformer 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_heads8, dropout0.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, biasTrue) 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_ratio4.0, dropout0.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_size224, patch_size16, in_channels3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4.0, dropout0.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), dim1) 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(dim1)。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_size2, 3×224×224把每个模块的输出形状打印出来一步步核对。x torch.randn(2, 3, 224, 224) model VisionTransformer(img_size224, patch_size16, num_classes10) 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看看精度和速度的差异把位置编码从可学习改成正弦编码对比训练曲线。改造实验的价值是让你理解每个超参数的影响而这比背结构更有效。第三层接入工程化框架。使用成熟的深度学习库或模型库加载预训练权重做数据增强、分布式训练、模型导出。这一层的目的是把模型放进真实项目让它稳定产出结果。这三层不是替代关系而是递进关系。跳过第一层直接进入工程化遇到问题会缺少拆解能力一直停留在第一层又会在真实项目里寸步难行。如果能把“手撕”当成理解工具而不是最终目的那这篇文章带来的价值就不会停留在“我照着代码敲了一遍”而是沉淀成一种拆解复杂模型的方法。下次再遇到新的网络结构你也能更快地摸清它的数据流、形状变化和关键设计。
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻