十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Transformer与ViT手写实现:从Attention机制到图像分类的完整指南

Transformer与ViT手写实现:从Attention机制到图像分类的完整指南 Day 34。今天终于把 Transformer 和 Vision TransformerViT这条线完整啃下来了。从 Attention 机制一路推到 ViT 的 patch embedding这个过程比我想象中复杂但也比想象中有意思。这篇笔记我边读边写把整个理解链路、手写代码的过程以及踩过的坑都整理出来希望能给后面学到这里的朋友省点时间。我和很多人一样最开始接触深度学习是从 CNN 入门的习惯了卷积的滑动窗口思维突然切到 Transformer 这种全局建模的架构其实很不适应。但等你真正理解了 Attention 在做的事情再回头看 ViT 的设计就会觉得这是一套非常优雅且自洽的方案。今天这篇内容适合两类人一是刚学完基础神经网络、想往 NLP/CV 前沿架构深入的同学二是已经用过 PyTorch 但一直对 Transformer 内部细节模模糊虎的实践者。我尽量把 Attention 的数学直觉、Transformer Encoder 的结构、ViT 的改动点都讲清楚再附上可以跑通的代码。1. 今天为什么要把 Transformer 和 ViT 放在一起学1.1 从 Attention 到 Transformer 的核心动机学习 Transformer 之前必须先把 RNN 和 CNN 的痛点想明白。RNN 的优势是天然按顺序处理序列能够记住前文信息但问题也很明显随着序列变长前面时刻的信息会被逐步稀释这叫长程依赖困境。LSTM 和 GRU 某种程度上缓解了这个问题但它们依旧是顺序执行的每一步依赖上一步的输出既慢又难并行化。CNN 在并行性上没问题但感受野是瓶颈你得叠加很多层才能让某个位置的输出覆盖到图像远端的信息而且这个覆盖是隐式的不是专门设计的。Transformer 的思路是直接把所有位置之间的关系一次性算出来不依赖顺序执行。它用了一个叫 Self-Attention 的机制让序列中任意两个位置之间可以直接建立依赖。这就等于说不管两个 token 离得有多远网络都能在一层之内就让它们的特征产生交互。这种设计让并行性大幅提升也使得长程依赖变成默认能力而不是额外能力。2017 年那篇《Attention Is All You Need》提出这个架构后NLP 领域迅速从 RNN 转向 Transformer后来的 BERT、GPT 都是在这个底座上长出来的。1.2 ViT 解决的是图像领域怎么用 Transformer 的问题图片本质上是像素点阵不能直接扔给标准的 Transformer 去处理。ViT 的贡献在于提出了一种非常简洁的桥接方式把图片切成固定大小的 patch比如 16×16 像素一块每个 patch 展平之后经过一个线性映射变成一个 token 向量。这样一张 224×224 的图就变成了 196 个 token再加上一个用于分类的 cls token总共 197 个 token 序列。接下来就是用标准的 Transformer Encoder 去处理这段序列最后用 cls token 的输出接一个分类头。很多人第一次看到 ViT 会有疑虑把图片切块再线性映射是不是丢掉了太多空间结构CNN 明明可以通过卷积核天然感知局部纹理和边缘Token 化会不会让模型什么都学不到其实这种担心是合理的ViT 论文里也承认了这一点所以在 Transformer 序列前加了可学习的位置编码并且用大量数据做预训练。论文中有一个非常直观的对比在 ImageNet-1k 这种中等规模的数据集上从头训练ViT 的效果略逊于当时的 SOTA CNN但在 JFT-300M 这种超大数据集上预训练之后ViT 再迁移到下游任务效果反超了 CNN。这个结果传达了一个关键信息Transformer 的强表达能力需要数据量来喂养一旦喂饱了它的上限比带归纳偏置的 CNN 更高。2. Self-Attention 的数学逻辑与 Multi-Head 的意义2.1 Attention 到底在算什么要读懂 ViT先得把 Self-Attention 的计算过程彻底弄明白。假设输入序列是 X形状是 [B, N, C]B 是 batch sizeN 是 token 数量C 是每个 token 的特征维度。Self-Attention 的第一步是把每个 token 分别映射成三个向量Query查询、Key键、Value值。用生活化的类比来解释假设你在教室里找座位Query 是你对自己需求的描述Key 是每个座位的特点标签Value 是这个座位的实际舒适程度。你需要先用自己的 Query 去跟所有座位的 Key 做匹配看看哪些座位更符合你的需求然后用匹配的权重去加权汇总所有座位的 Value最终得到一个融合了全局信息的表示。在 Transformer 里这个过程被形式化为Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V这里有个细节容易被忽略为什么要除以根号 d_k当特征维度比较大的时候Q 和 K 的点积结果会变得很大softmax 的梯度会趋近于零导致训练困难。除以根号 d_k 是为了把点积的方差缩放到一个合适范围让 softmax 的输出不至于过于极端。d_k 是每个 head 的维度实际代码里 C / num_heads。计算一次 Attention 之后每个 token 的输出其实是所有 token 的 Value 按注意力权重加权的结果。权重越大说明当前 token 越关注那个位置。这也是为什么 Attention 被翻译成“注意力”的原因——网络会自动学习关注哪里。在 ViT 的代码里通常用矩阵乘法实现Q 和 K 的转置相乘得到 [B, num_heads, N, N] 的注意力矩阵这个矩阵是可视化的常用工具可以看到 cls token 关注了哪些 patch。2.2 为什么需要多头注意力单头 Attention 有一个潜在问题它只能学习一种“关注模式”。可实际上一个输入序列里的词或 patch 之间的关系是多维度的有的需要关注相邻区域有的需要关注远处相似结构有的需要关注颜色或纹理。一个头学不过来那就多开几个头。Multi-Head Attention 的做法是把 C 维特征切成 num_heads 份每个头独立计算一份 Q、K、V得到各自独立的注意力矩阵。每个头可以关注到不同方面的关系。最后把所有头的输出拼接起来经过一个线性层映射回 C 维。这样模型就能同时捕捉多种特征交互模式表达能力自然上去了。代码里最常见的实现是把 QKV 投影在一个大线性层里做然后 reshape 成多头形状这样做效率更高也更符合 PyTorch 的习惯。2.3 位置编码Transformer 与 ViT 的关键差异Self-Attention 本身是置换不变的交换任意两个 token 的输入顺序输出并不会改变顺序信息。这对语言和图像都是不可接受的句子里的词序决定语义图片里的 patch 位置决定结构。所以必须额外把位置信息注入进去。原版 Transformer 用的是三角函数式的位置编码根据位置的奇偶分别用不同频率的 sin 和 cos 函数生成固定向量。这种编码的好处是不需要学习能够外推到比训练时更长的序列。ViT 那边则更直接用可学习的位置编码初始化一个形状是 [1, N1, C] 的 Parameter和 patch embedding 相加。在训练过程中位置编码会跟着整个网络一起更新相当于让模型自己学会“每个位置应该长什么样”。那么问题来了ViT 为什么不用三角函数位置编码这个问题也困扰过我一阵后来看了一些消融实验才明白。对于图像 patch 序列来说位置数量是固定的不像 NLP 需要拼长文本所以可学习编码完全够用而且可学习编码对每个 patch 独立建模位置信息不像 2D 相对位置编码那样显式编码空间邻接关系但它保留的绝对位置信号足以让 attention 学到空间关系。当然后续的 Swin Transformer 等改进模型又引入了相对位置编码因为它在局部窗口内显式建模了 patch 之间的相对偏移对小目标和密集预测任务更友好。3. Vision Transformer 的核心架构拆解3.1 Patch Embedding图像如何变成 tokenViT 的前处理可以理解为“图像分词”。给定一张 224×224×3 的彩色图片设定 patch_size16那么一个 patch 就是 16×16×3 的小方块总共切成 (224/16)² 196 个 patch。每个 patch 展平后是 768 维向量这正好是 ViT-Base 的 hidden size。实现这一步的常用方式是使用一个卷积核大小和步长都等于 patch_size 的 Conv2d输入通道为 3输出通道为 embed_dim。这个操作等价于对每个 patch 做一次参数共享的线性变换卷积的权重就是那个线性映射矩阵。处理后的特征图经过 flatten 和 transpose变成 [B, N, C] 的 token 序列完成从像素空间到特征空间的转换。这里有个小技巧也是我一开始没注意到的如果直接手写一个 reshape 然后接 Linear效果和 Conv2d 是等价的但 Conv2d 实现更简洁还能利用底层优化的卷积算子速度更快。所以几乎所有开源 ViT 实现都是这么做的。3.2 CLS token 与 Transformer Encoder在 patch embedding 之后ViT 会额外拼接一个可学习的 cls token放在序列最前面。这个 token 的作用是充当全局信息汇聚器训练时最终序列输出里第一个 token 的向量被拿去接分类头其他 patch token 则只负责中间表示的学习。Transformer Encoder 本身由多层相同的 Block 组成。每个 Block 的核心结构是LayerNorm → Multi-Head Self-Attention → 残差连接 → LayerNorm → MLP → 残差连接。MLP 通常包含两层全连接中间用 GELU 激活隐藏维度一般为 embed_dim 的 4 倍ViT-Base 就是 768→3072→768。需要特别强调的是 LayerNorm 的位置。原版 Transformer 用的是 post-norm也就是 attention 之后才接 LayerNormViT 沿用了这个设计。但今天很多新模型包括 GPT 系列都改成了 pre-norm也就是在进入 attention 前先做 LayerNorm残差里没有额外正则。pre-norm 的好处是训练更稳定梯度流更干净。ViT 论文里其实用的是 post-norm我在手写实现时两种都试过肉眼可见 pre-norm 在小数据集上收敛得更快。如果你从零开始写自己的 ViT建议用 pre-norm 做默认选项。3.3 训练策略与数据规模ViT 的成功离不开大规模预训练。如果你只在 CIFAR-10 这种小数据集上从头训练效果通常不如 ResNet因为 ViT 没有 CNN 那种局部性和平移等变的先验知识必须靠大量数据让网络自己去学这些规律。训练时的常用策略包括随机裁剪、水平翻转、mixup、cutmix、随机擦除等数据增强配合 cosine 学习率衰减和 warmupAdamW 优化器weight decay 通常会设到 0.05 甚至更高。我在小数据集上实验时还发现一个有意思的现象适当增大 patch_size 反而能提升模型在小图上的稳定性。原因很直接patch 变大意味着序列长度变短计算量下降而且每个 token 包含的局部信息更丰富对小数据训练压力更小。当然patch 过大也会丢失细节需要根据任务平衡。4. 实操手写一个简化版 ViT4.1 准备数据与预处理今天下午我用了 CIFAR-10 当实验对象原因很简单小、快、大家都熟。CIFAR-10 每张图只有 32×32直接用 patch_size16 会让序列太短所以我把图 resize 到 224×224或者改用小 patch_size 比如 4。我这里演示的是 224 的版本方便对照原论文结构。数据预处理参考通用实践Resize 到 224×224随机水平翻转归一化到标准 ImageNet 统计值。训练集和测试集用同样的归一化参数但测试集不做随机增强。4.2 核心实现代码以下是完整可运行的简化版 ViT核心组件拆成了 PatchEmbedding、MultiHeadSelfAttention、TransformerBlock、ViT 四个类。我刻意把 attention 内部的前向过程写出来了方便对照公式理解。import torch import torch.nn as nn class PatchEmbedding(nn.Module): def __init__(self, in_channels3, patch_size16, embed_dim768, img_size224): super().__init__() self.patch_size patch_size self.n_patches (img_size // patch_size) ** 2 # 用 stridepatch_size 的卷积完成无重叠切块和线性映射 self.proj nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): x self.proj(x) # [B, embed_dim, H/p, W/p] x x.flatten(2) # [B, embed_dim, n_patches] x x.transpose(1, 2) # [B, n_patches, embed_dim] return x class MultiHeadSelfAttention(nn.Module): def __init__(self, embed_dim768, num_heads12, dropout0.0): super().__init__() assert embed_dim % num_heads 0 self.num_heads num_heads self.head_dim embed_dim // num_heads self.qkv nn.Linear(embed_dim, embed_dim * 3) self.proj nn.Linear(embed_dim, embed_dim) self.dropout nn.Dropout(dropout) def forward(self, x): 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) # 3, B, heads, N, head_dim q, k, v qkv[0], qkv[1], qkv[2] attn (q k.transpose(-2, -1)) * (self.head_dim ** -0.5) attn attn.softmax(dim-1) attn self.dropout(attn) x (attn v).transpose(1, 2).reshape(B, N, C) x self.proj(x) return x class TransformerBlock(nn.Module): def __init__(self, embed_dim768, num_heads12, 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.Linear(hidden_dim, embed_dim), nn.Dropout(dropout) ) def forward(self, x): x x self.attn(self.norm1(x)) # pre-norm 残差 x x self.mlp(self.norm2(x)) return x class ViT(nn.Module): def __init__(self, img_size224, patch_size16, in_channels3, num_classes10, embed_dim768, depth12, num_heads12, mlp_ratio4.0, dropout0.1): super().__init__() self.patch_embed PatchEmbedding(in_channels, patch_size, embed_dim, img_size) self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, 1 self.patch_embed.n_patches, embed_dim)) self.pos_drop nn.Dropout(dropout) self.blocks nn.ModuleList([ TransformerBlock(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) self._init_weights() def _init_weights(self): nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) for m in self.modules(): if isinstance(m, nn.Linear): nn.init.xavier_uniform_(m.weight) if m.bias is not None: nn.init.zeros_(m.bias) def forward(self, x): x self.patch_embed(x) # [B, N, C] cls_tokens self.cls_token.expand(x.shape[0], -1, -1) x torch.cat([cls_tokens, x], dim1) # [B, N1, C] x x self.pos_embed x self.pos_drop(x) for block in self.blocks: x block(x) x self.norm(x) cls x[:, 0] return self.head(cls)这段代码的可读性优先没有刻意压缩核心就是搞清楚 QKV 的 reshape、Attention 的计算、token 序列的拼接顺序。如果你打算拿它做实验可以把 depth、embed_dim、num_heads 都调小比如 depth6、embed_dim192、num_heads6在小数据集上跑起来会快很多。4.3 训练参数与调试经验训练 ViT 时我踩过的第一个坑是学习率。CNN 里常用的 0.1 起步学习率在 Transformer 上直接爆掉loss 变成 NaN。后来切成 AdamW 学习率 1e-4 warmup 5 个 epoch cosine 衰减训练才稳定下来。小 batch 和大学习率的组合在 ViT 上尤其灵敏建议 batch size 从 64 起步学习率按 batch 同比缩放。另一个值得亲自体会的点是 warmup 的作用。Transformer 在训练初期对学习率特别敏感因为没有足够的迭代来稳定 LayerNorm 和 attention 的统计量一上来就给大学习率很容易让梯度爆炸。warmup 相当于给网络一个“热身期”让参数慢慢适应优化方向。我试过省掉 warmup结果同样的数据下准确率掉了将近两个点这个差距非常明显。CIFAR-10 上用刚才那个小型 ViT 配置跑 100 个 epoch大概能到 78% 左右。这个数字不算高因为 ViT 在小数据上和 CNN 相比完全没有优势但这不影响你观察它的收敛曲线和训练行为。想更快迭代的话可以只跑 20 个 epoch 看趋势重点观察训练 loss 是否稳定下降、attention 的 logits 是否分布合理。5. 站在 Day 34 往四周看近期热门变体与应用5.1 从 ViT 到 Swin、DeiT、RestormerViT 是图像 Transformer 的起点但它并不是终点。Swin Transformer 是我比较推荐了解的下一代架构它引入了层次化设计和窗口注意力。所谓窗口注意力是限制每个 token 只和附近窗口内的 token 做 attention窗口之间通过 shift 操作交替连接。这样既保住了局部归纳偏置又通过窗口挪移实现了跨窗口信息流动。Swin 在目标检测、分割这类密集预测任务上表现出色因为它能输出多尺度特征图这是 ViT 做不到的。DeiT 则研究了如何在小数据上训 ViT核心技巧是用一个 teacher CNN 做知识蒸馏并额外增加一个 distillation token 参与训练。这么做的好处是让它用 ImageNet-1k 就能达到接近 CNN 的水平不用非得像原始 ViT 那样依赖巨大的 JFT 数据集。Restormer 是另一条路线它面向图像复原任务把注意力施加在通道维度上同时用多头转置注意力保持效率在去雨、去噪、超分等底层视觉任务上表现很强。5.2 Attention 机制的跨界应用Attention 的变形远不止图像分类。跨注意力Cross Attention常用于多模态任务比如文本描述和图片特征的对齐。和 Self-Attention 不同Cross Attention 的 Q 来自一个模态K 和 V 来自另一个模态让两种信息互相“查询”。多模态行人检测、RGB-T 融合这类任务就会用到 Deformable Cross-Attention 这类改进版目的是让注意力的采样位置是可学习的从而在不对齐的跨模态特征间找到对应关系。Coordinate Attention 和 Double Attention 则在 CNN 的语境里引入了一些注意力思想。前者在通道注意力里加入坐标信息让模型知道特征在空间上的位置后者同时聚合全局和局部信息用两组注意力矩阵组合出新特征。Point Transformer 把 Transformer 用在点云上对每个点做 k 近邻采样后计算注意力这让我意识到 Transformer 并不在意输入是像素、词还是点云它只关心集合中元素之间的关系。只要你能把数据组织成一组 tokenTransformer 就能试一下。6. 学习过程中踩过的坑快点记下来6.1 位置编码选择不当导致训练震荡我在复现时曾经把 ViT 的位置编码临时改成三角函数式的结果训练 loss 一直震荡。后来排查发现三角函数编码是按 1D 顺序频率生成的而图像 patch 是 2D 排列的直接套用等于丢掉了垂直方向上的空间结构。ViT 之所以用可学习位置编码是因为它不需要预设任何先验让模型从数据里自己学出“哪个 patch 在哪个位置”。所以如果你要在图像任务里改位置编码要么用可学习的 1D 编码要么考虑 Swin 那种 2D 相对位置偏置别直接用 NLP 里那套三角函数硬搬。6.2 训练不收敛要从这几个方向排查第一检查注意力分数有没有变成 NaN常见原因是 QK^T 数值过大被 softmax 放大确保除以根号 d_k。第二检查 LayerNorm 的 eps默认 1e-5 在某些情况下不够稳定可以调到 1e-6。第三检查残差连接有没有拼错维度。我记得有一次在 TransformerBlock 里把残差加到了 norm 之后而非 norm 之前训练 loss 始终下不去花了很久才发现。第四优化器权重衰减和梯度裁剪要配合Transformer 对梯度范数很敏感clip_grad_norm_ 设个 1.0 能避免意外爆掉。6.3 显存爆炸是常态学会轻量化ViT 的显存消耗主要来自注意力矩阵形状是 [B, num_heads, N, N]。N 是 token 总数它和 patch_size 的平方成反比patch 越小序列越长显存涨得越快。如果显存不够优先调大 patch_size 或者减小 depth。Flash Attention 这类 IO 感知注意力把注意力计算分块到 SRAM 上是训练长序列时非常实用的优化手段很多主流框架都已经内置支持我也把它加入了下一次实验的计划中。另一个经验是尽量在推理时把 dropout 关掉测试集上能稳定涨一点精度。最后分享一个我自己的体会今天整个学下来最大的收获不是记住了 ViT 的代码而是看懂了“Attention 是机制Transformer 是框架ViT 是应用”这三层关系。Attention 负责给出一个通用的关系建模方式Transformer 把它安排成可堆叠、可并行、可扩展的网络结构ViT 则表明这种结构可以跳出文本领域成为视觉任务的新底座。如果你正在学这块我的建议是别只读论文敲一遍代码比看十遍公式都有用。先从手写 Self-Attention 开始再叠一个 Block然后跑通 ViT 的前向和后向最后在小数据集上从头训练直到能够收敛。这个过程一定会踩坑但每一个坑都在帮你建立对模型真正直观的理解。后面我打算继续往下去写 Flash Attention 的原理和 Swin 的实现细节如果你们也有想深入的方向欢迎在评论区告诉我。
返回列表