Transformer原理深度拆解:从Self-Attention到ViT与工程实践
发布时间:2026/9/9 9:34:54 锦皓数字建站

1. 为什么Transformer能取代RNN一个被反复低估的本质原因做了几年深度学习看过的模型架构不算少。但每次被问到Transformer的原理到底是什么这个问题我都觉得很难用一两句话讲清楚。因为它不是一个单一的创新而是一整套设计决策的组合。我最早接触Transformer不是从论文原文《Attention Is All You Need》开始的而是从翻译任务里发现LSTM怎么调都不太对劲——长句子效果差、训练慢、还经常梯度消失——然后才认真去啃Transformer的实现原理。先说一个核心结论Transformer最本质的贡献是把序列建模从逐步递推变成了一步到位。RNN/LSTM处理序列时隐状态像一个接力棒从第一个词传到第二个词再传到第三个词。这个设计天然有个致命问题距离越远信息损耗越大。虽然LSTM用门控机制缓解了这个问题但本质上它还是在传话传话总有失真。而Transformer里的Self-Attention机制让序列中任意两个位置直接建立联系无论它们隔得多远信息传递路径都是一条直线没有中间人。这个路径长度的差异比很多人想象的重要得多。你可以把RNN想成一条只能逐个站点停靠的公交线路从起点到终点必须经过所有中间站而Transformer是直达航班任意两个城市之间都能直飞。当序列长度拉长到几百上千的时候这种差异就变成了能不能用的问题而不仅仅是好不好的问题。另一个被忽视的点是并行计算。RNN必须按时间步串行计算第t步依赖第t-1步的隐状态所以GPU再强也快不起来。Transformer没有这种依赖关系一个序列里所有位置可以同时计算这也是为什么它能在大规模数据上训练起来。说句题外话当年我为了提升RNN的训练速度试过各种trick包括调整batch size、换优化器、混合精度效果都很有限。换到Transformer之后训练速度反而上去了而且模型效果还更好这个对比是相当直观的。所以如果你要理解Transformer原理第一个要建立的认知就是Self-Attention是主干位置编码是辅助残差和归一化是保障。它不是像LSTM那样在Rnn上加注意力机制而是把注意力机制本身做成了整个网络。2. Self-Attention的计算拆解Q/K/V背后到底在做什么2.1 从词向量到Q、K、V三次线性投影的意义Self-Attention的第一步是把输入向量变成三个不同的向量分别叫Query查询、Key键、Value值。很多人第一次接触这三个概念会懵其实用一个生活化的场景就很好理解。想象你在图书馆找一本书Query就是你想找的问题比如深度学习入门有哪些经典教材Key是每本书的标题和标签是图书馆用来索引的Value就是书本身的内容。Attention的过程就是拿你的Query去和每一本书的Key做匹配算出相似度分数然后用这个分数对Value做加权求和——相似度高的书内容占的权重大相似度低的权重小。在代码实现上这个过程就是三次矩阵乘法import torch import torch.nn as nn class SelfAttention(nn.Module): def __init__(self, d_model, d_k, d_v): super().__init__() self.W_Q nn.Linear(d_model, d_k) self.W_K nn.Linear(d_model, d_k) self.W_V nn.Linear(d_model, d_v) def forward(self, x): # x: [batch_size, seq_len, d_model] Q self.W_Q(x) # [batch_size, seq_len, d_k] K self.W_K(x) # [batch_size, seq_len, d_k] V self.W_V(x) # [batch_size, seq_len, d_v] return Q, K, V为什么要做三次不同的投影而不是直接用原始向量因为如果直接用原始向量所有位置之间的关系就固定死了表达力不够。通过训练学习到的W_Q、W_K、W_V矩阵模型可以灵活决定从哪个角度去衡量两个位置的相关性。比如在翻译任务里代词it可能需要在性别、单复数等不同维度上去找它的先行词单一空间很难同时满足这些需求。2.2 缩放点积为什么要除以根号d_k得到Q、K、V之后下一步是计算注意力分数。公式是[ \text{Attention}(Q,K,V) \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V ]两句话就能说清公式Q乘以K的转置得到形状为[batch_size, seq_len, seq_len]的相似度矩阵除以(\sqrt{d_k})做缩放然后softmax归一化成权重最后和V相乘得到输出。关于为什么要除以(\sqrt{d_k})这是很多教程一笔带过但很重要的问题。假设Q和K的每个元素都是均值为0、方差为1的随机变量那Q和K的点积结果方差会是(d_k)。当(d_k)比较大时点积结果的数值会很大导致softmax进入饱和区——也就是大部分概率都集中在一个位置上其他位置几乎为0梯度会变得非常小训练不动。除以(\sqrt{d_k})后方差重新回到1softmax的输入分布更平缓梯度更健康。有次我写代码时把缩放系数漏了模型表现很奇怪——训练loss下降得特别慢而且一旦学习率调大一点就爆炸。排查了很久才意识到是注意力分数太大导致softmax饱和。这个细节如果你自己手写Transformer几乎一定会踩一次。2.3 Mask解码器里不能偷看未来的约束在实际训练中解码器的Self-Attention有一个特殊要求预测第t个词时只能看到位置1到t-1的信息不能看到未来的词。实现方式就是在softmax之前对未来的位置加上一个非常大的负数通常用-inf或者-1e9这样softmax之后这些位置的权重就趋近于0。def causal_mask(seq_len): mask torch.triu(torch.ones(seq_len, seq_len), diagonal1) mask mask.masked_fill(mask 1, float(-inf)) return mask这个mask矩阵是个上三角矩阵对角线以上的位置全是-inf。除此之外在实际应用中还经常会遇到padding mask也就是把序列中填充的无效位置也mask掉。两种mask可以结合使用先padding mask再causal mask或者反过来具体取决于框架习惯。2.4 复杂度问题O(n²)从哪里来怎么办Self-Attention的时间复杂度是(O(n^2 \cdot d))n是序列长度d是特征维度。这个n²就是QK^T这个矩阵乘法带来的——它把一个长度为n的序列两两配对产生了n×n的注意力矩阵。当n是128、512这种规模时没有问题但如果n到几千、几万比如处理整篇长文档、高分辨率图像、长视频n²就会爆炸。这里需要记住一个概念token数量决定复杂度特征维度只影响常数因子。所以后续的Sparse Attention、滑动窗口注意力、Swin Transformer的分窗注意力本质上都是在想办法把n²这个大矩阵变稀疏只计算有效位置的注意力。3. 多头注意力与位置编码顺序感与信息分工的隐藏逻辑3.1 多头机制到底在做什么多头注意力Multi-Head Attention就是把刚才的Self-Attention过程重复h次每次用不同的W_Q、W_K、W_V参数然后把h个结果拼在一起再经过一个线性层输出。为什么要多头因为语言里的相关性是多维的。还是拿it举例子在一个句子里it可能既和cat有关指代对象又和fed有关动作对象还和kitchen有关地点环境。单头注意力只能算一组相关性分数无法同时捕捉这么多维度的关系。多头相当于让模型从h个不同的角度去理解序列每个头专注一种关系类型。实际训练中有些头会学到相邻词的局部依赖有些头会学到远距离的核心指代关系有些头则倾向于关注句法结构。这就是为什么几乎所有的Transformer变体都保留了多头结构4头、8头、16头、32头都有具体数量根据模型大小和数据规模调整。我在实际项目里的一般经验是模型越大、数据量越大头数可以适当增加但头数乘以每个头的维度要等于d_model这是一个约束。多头注意力的完整调用在PyTorch里可以直接用现有的类不用手动实现但理解内部的维度变换很重要import torch.nn as nn # d_model 512, n_head 8, 每头维度 64 mha nn.MultiheadAttention(embed_dim512, num_heads8, batch_firstTrue)输入和输出都是[batch_size, seq_len, d_model]内部先把d_model拆成n_head个d_k64的子空间计算完之后再合并回d_model。3.2 位置编码没有它Transformer就是一个词袋模型这个是我觉得很多人理解Transformer原理时最容易忽略的点。Self-Attention本身是置换不变的——它对输入顺序不敏感。把句子里的词顺序打乱Self-Attention计算出来的结果是相同的因为注意力权重只取决于Q、K、V之间的点积关系而不取决于位置。如果直接去掉位置编码Transformer就变成了一个高级词袋模型完全丧失了顺序信息。原始论文用的是正弦余弦位置编码公式如下[ PE_{(pos, 2i)} \sin\left(\frac{pos}{10000^{2i/d_{model}}}\right) ][ PE_{(pos, 2i1)} \cos\left(\frac{pos}{10000^{2i/d_{model}}}\right) ]pos是位置下标i是维度下标。这个公式的巧妙之处在于不同维度使用不同频率的正弦/余弦波位置编码就能同时编码绝对位置和相对位置信息。为什么这么说因为三角函数有一个恒等式(\sin(\alpha\beta)\sin\alpha\cos\beta\cos\alpha\sin\beta)意味着位置posk的编码可以用位置pos的编码经过线性变换得到这让模型在理论上更容易学到相对位置关系。不过这已经是2017年的设计。现代模型越来越多地使用AliBi或RoPE旋转位置编码这类相对位置编码。RoPE在LLaMA等模型中被广泛采用它的核心思路是在Q和K上注入旋转矩阵让注意力分数自然包含相对位置信息。3.3 一个容易踩的坑位置编码和词嵌入相加时的维度对齐在代码实现里位置编码是直接加到词嵌入上的不是拼接。这意味着位置编码的维度必须和词嵌入的维度一致都是d_model。如果你把词嵌入维度设成768那位置编码也得是768否则就会报维度不匹配。还有一点需要特别注意不同的实现中位置编码可学习参数的初始化方式会影响训练效果。Transformer原始论文用的是固定不可学习的sin/cos编码但BERT等模型使用的是可学习的位置嵌入。在数据量不够大的情况下固定编码更稳妥不容易过拟合数据量足够大时可学习编码上限更高。4. 归一化与残差训练稳定性的工程化设计4.1 LayerNorm为什么比BatchNorm更适合TransformerTransformer中的每个子层Self-Attention和Feed-Forward后面都接一个LayerNorm层而且每个子层还外加残差连接。之所以用LayerNorm而不是BatchNorm原因很简单BatchNorm是在batch维度上做归一化对batch size敏感而且当序列长度不固定时很难统一处理。LayerNorm是在特征维度上做归一化不依赖batch大小而且对每个token单独做完美适配变长序列。你可以这样理解BatchNorm是给整个班级的学生统一调平均分每个学生都减去班级平均分再除以班级标准差这要求班级足够大才有统计意义LayerNorm是给每个学生单独做标准化不管其他同学成绩如何只把自己成绩分布调整为均值0方差1。Transformer的训练通常batch size不大尤其是大模型受限于显存所以LayerNorm是唯一正确的选择。4.2 Pre-LN还是Post-LN一个影响训练成败的细节原始Transformer论文用的是Post-LN也就是把LayerNorm放在残差相加之后公式是[ \text{Output} \text{LayerNorm}(x \text{SubLayer}(x)) ]但这个设计在实践中有一个众所周知的坑训练不稳定尤其深层模型很容易发散。后来的研究发现如果改成Pre-LN把归一化放在子层之前[ \text{Output} x \text{SubLayer}(\text{LayerNorm}(x)) ]训练稳定性会大幅提升收敛也更快。现在的主流开源框架包括GPT系列、LLaMA、ChatGLM等基本都采用Pre-LN或者它的变体。我自己的经验是如果你从零开始训练一个Transformer优先选Pre-LN如果你是在已有预训练模型上做微调可以保持原模型的归一化位置不变。有次我微调一个Post-LN的模型刚开始训练loss就出现NaN排查了一圈数据、学习率、优化器都没问题最后把归一化方式改成Pre-LN就正常了。这个坑值得记住但更值得记住的是不要轻易改动预训练模型的架构设计。4.3 残差连接为什么没有它深层网络根本训练不动残差连接就是直接把输入x加到子层输出上[ \text{Output} x \text{SubLayer}(x) ]它的作用是让梯度和信息可以直接从最后一层流向第一层避免梯度在多层传播中消失。没有残差连接深层的Transformer很难训练基本是学界共识。当你手写代码时一定要保证残差连接的路径不经过任何破坏性操作比如dropout之外不要加别的随机性模块。实际工程中Dropout通常加在残差支路子层输出上而不是主路径上。这个顺序也影响训练效果如果在主路径上加Dropout残差的信息传递会被干扰模型会更难收敛。5. 从NLP到CVViT、Swin、Deformable和Restormer都在改什么5.1 ViT把图像当成句子Vision TransformerViT的技术路线很简单把一张图像切成固定大小的patch比如16×16像素一块然后每块拉平成一个向量经过一个线性投影变成patch embedding最后拼接上位置编码送进标准的Transformer encoder。这个操作的动机是图像本身是一种序列数据只是序列元素从词变成了图像块。ViT在JFT-300M这种超大规模数据集上训练后效果超过了当时的CNN。但我必须提醒一下ViT对数据量的要求很高。如果你只有几万张图像的训练集直接上ViT很可能会比ResNet差。原因也好理解开始切patch的时候模型并不知道像素之间的局部空间关系——比如相邻像素往往颜色相近、纹理连续。这层归纳偏置是CNN通过卷积核自带的ViT里没有只能靠数据喂出来。所以ViT的实践操作里通常会加一个在CNN预训练后用ViT微调或者先用小patch增大局部建模能力的技巧。5.2 Swin Transformer用层级和窗口换效率Swin Transformer解决的是ViT的两个痛点一是计算复杂度高全局Self-Attention的n²复杂度在图像上尤其疼二是缺乏多尺度特征图像里的物体有大有小单一尺寸的patch很难适配。Swin的做法是把注意力限制在每个局部窗口内窗口之间通过移位窗口操作来交换信息。这么做有两个直接收益计算复杂度从O(n²)降到了O(n)同时通过窗口划分-合并形成了类似CNN特征金字塔的层级结构。在做目标检测、语义分割这类需要多尺度特征的任务时Swin的优势非常明显。如果你在思考为什么Swin和ViT长得很不一样核心答案就一句话因为图像和文本的结构先验不同。文本天然是线性的全局注意力合理图像天然是二维的局部相关性更重要所以要在注意力的视野范围上做文章。5.3 Deformable Attention和Restormer稀疏化的进一步探索Deformable Attention的思路是不去计算所有位置的注意力而是学习预测每个query应该关注哪些key。就像人看一幅图时眼睛不会平均扫描每个像素而是重点关注几个显著的物体轮廓或运动区域。这种方法在弱对齐的RGB-T双模态行人检测等任务上有显著收益因为模态间的对齐关系本身不固定需要模型自己学会聚焦。Restormer则是在低层视觉任务去雨、去噪、超分上的Transformrer改造。低层视觉和分类最大的不同是每个像素都需要精细的输出所以不能用Swin那样粗暴的窗口缩小感受野。Restormer提出的多头转置注意力在通道维度上计算注意力核心公式是[ \text{Attention}(Q,K,V) V \cdot \text{Softmax}(K^T \cdot Q / \alpha) ]注意这里Q和K乘法的顺序变了矩阵形状从[N×N]变成了[C×C]C是通道数。对于高分辨率图像输入N像素数通常几十万上百万而C可能只有64或128所以转置注意力把复杂度从像素维度转移到通道维度算力消耗大幅降低。如果你在做图像复原类任务想用Transformer又担心算力不够Restormer这种在通道维度做注意力的思路非常值得参考。这个方案我第一次看到的时候也觉得有点反直觉但仔细推演一下复杂度就明白了。5.4 一个值得思考的趋势Transformer正在成为通用接口从NLP到CV再到多模态的CLIP、多模态UAV感知中的Geometry-aware Alignment TransformerToken化Self-Attention已经变成了一套相当统一的方法论。不管输入是文本、图像、点云还是光谱数据第一步都是把原始输入切成token并嵌入到同一特征空间第二步是用Self-Attention建模token之间的关系第三步把输出的token映射回目标任务空间。这种归一化的建模方式的最大好处是不同模态的数据可以在同一个框架里融合。比如你有一组无人机多模态感知数据可见光红外激光雷达点云用三个不同的编码器提取特征再全部转成token后丢给同一个Transformer交互层每个模态的信息就能在特征层面直接匹配和融合。这在过去CNN时代是很难想象的因为卷积核天然绑定单模态的图像grid结构。6. 手写Transformer时最容易踩的五个坑来自实战的排除经验6.1 维度不匹配大多数bug的源头我见过太多新手在实现Transformer时卡在维度问题上。这里有一个自查方法把每个张量的形状在关键步骤打印出来一步步核对。输入x [batch_size, seq_len, d_model]Q、K、V投影后各自为[batch_size, seq_len, d_k]或[batch_size, seq_len, d_v]Q K^T结果为[batch_size, seq_len, seq_len]这里如果shape不对先查投影维度softmax后再乘V变为[batch_size, seq_len, d_v]多头合并拼接后是[batch_size, seq_len, head*d_v]必须等于d_model如果你用了nn.MultiheadAttention且batch_firstFalsePyTorch默认值是False输入要变成[seq_len, batch_size, hidden_dim]这个order搞反也会得到莫名其妙的报错或更糟糕的错误结果。用batch_firstTrue可以省心很多但并不是所有预训练代码都这么写。6.2 Mask顺序先padding mask还是先causal mask在做训练时encoder的Self-Attention只需要padding maskdecoder的Self-Attention要同时用padding mask和causal mask。顺序上是先把两个mask对齐都是[seq_len, seq_len]形状然后取或或者按位相加确保如果一个位置被padding mask标记为不可见那它也不参与causal mask的计算。一个常见的bug是只做了causal mask但忘了padding mask导致模型在预测时还会看到padding位置的权重训练loss明显偏高但又不至于不收敛属于那种很难察觉的隐性错误。排查方法很简单把attention weights打印出来看padding位置的权重是否接近0。6.3 位置编码加在输入上还是加在每层上标准做法是只加在输入端。但也有一部分模型在每一层的输入都注入位置信息比如ALiBi是直接在注意力分数上加一个与距离相关的偏置项。如果你在魔改架构必须清楚自己的位置编码方案是绝对位置还是相对位置两者不能混着用否则位置信息会被重复编码模型反而困惑。6.4 训练中的Loss Spike和NaN问题用fp16训练Transformer时loss spike是一个经典问题。最常见的原因是梯度溢出尤其是Attention里softmax前的数值过大。解决手段很成熟梯度裁剪gradient clipping一般max_norm设置在0.5到1.0混合精度训练时的loss scaling策略调整在某些层用bf16代替fp16特别是在使用LLaMA类模型时另一个我自己踩过的坑是使用Adam优化器时如果beta2设得太大比如0.999而训练初期梯度噪声又大会导致自适应学习率分母过大模型直接原地踏步。把warmup steps设长一点比如总steps的5%-10%能有效缓解。6.5 推理时的KV Cache最该优化的性能点训练完成后部署推理时Transformer的自回归生成有一个天然的性能瓶颈每生成一个token都要重新计算之前所有token的K和V。KV Cache的思路是把已经计算过的K和V缓存下来下一步就不重复计算。具体实现是在循环中维护两个列表past_key_values None # 初始为空 for step in range(max_new_tokens): outputs model(input_ids, past_key_valuespast_key_values, use_cacheTrue) past_key_values outputs.past_key_values next_token_logits outputs.logits[:, -1, :] # 采样、贪心或其他策略生成下一个token使用KV Cache后每一步的计算量只取决于当前token而不是整个序列长度推理速度可以成倍提升。这个优化在输入序列很长的场景下尤其重要——比如聊天机器人的多轮对话历史越长cache收益越明显。我在实际部署Transformer模型做生成任务时这句话已经成了我排查性能问题的第一步先确认cache有没有生效再考虑模型结构优化。因为很多推理框架默认会开启这个功能但如果你自己纯手写推理脚本非常容易漏掉。7. 关于学Transformer原理这件事的一点额外心得可能你已经发现了Transformer的原理并没有想象中那么高不可攀。它核心的东西就三个Self-Attention、位置编码、归一化和残差。但要做到真正理解光看文章是不够的——你一定要自己动手从零写一遍或者至少把一个开源实现逐行读懂。我自己当年是把一份最小的Transformer实现大概300行PyTorch反复看了好几遍手动画出每个张量的shape变化图才算真正理解了。更关键的体会是理解一种架构的原理不是为了背诵它的结构而是为了获得设计判断力。当你会从这组数据里元素之间的关系是什么这个角度去思考模型选型时自然就知道该用RNN还是Transformer该上全局注意力还是局部窗口注意力该加什么位置编码。这种判断力只能在实际项目中积累没有捷径。如果你正处在代码能跑但说不清原理或者论文读了不少但一动手就错的阶段我的建议是挑一个你最关注的下游任务比如文本分类、图像复原或目标检测找到对应的Transformer实现把它从数据输入到loss输出完整走一遍。这个过程会比你看一百篇论文都有用。
锦
锦皓数字建站
深耕本土企业品牌数字化升级,专注原创端正雅致商务官网,从视觉设计到稳定运维全程保驾护航。