Transformer 这个架构,现在几乎成了人工智能大模型的基础组件。从文本生成、机器翻译到图像分类、视频理解,凡是你能看到的大模型,底层基本都绕不开它。这篇文章就负责把 Transformer 从头到尾拆开,配合图解思路讲清楚每个模块为什么存在、怎么计算、维度怎么变、训练时容易在哪里翻车。
我默认读者已经掌握了基本的神经网络和深度学习概念,但不要求你手写过 Attention。这篇内容适合三类人:正在准备面试的算法岗同学、需要做 NLP 或视觉项目的工程师、以及想把模型架构搞明白再动手写代码的初学者。整篇按“结构认知 + 逐块图解 + 最小实现 + 排错”的顺序来写,不用按章节顺序读,卡在哪一章就回到哪一章看。
1. 先搞清楚:Transformer 到底解决了什么问题
1.1 为什么“最后是 Transformer”:从序列建模的困境说起
在 Transformer 出现之前,处理序列数据主要靠 RNN/LSTM 和 CNN。RNN 是顺序扫描的,第 t 个词要等前 t-1 个词算完,并行性很差。长距离依赖问题虽然有 LSTM 的门控机制缓解,但模型仍然是按时间步展开,序列一长,梯度传播和信息丢失依然很麻烦。CNN 可以并行处理窗口内的词,但需要靠堆层数扩大感受野,才能建立跨度很大的依赖关系。注意力机制虽然很早就被用在翻译模型里,但当时通常是配合 RNN 使用,不是单独作为主架构。
Transformer 的做法是彻底去掉循环,直接对整句话计算两两之间的注意力得分。这样有两个直接好处。第一,并行度高,一句话内部的所有 token 可以同时参与计算。第二,任意两个位置之间的信息传递只需要一步,不需要沿着时间步或层逐层传递。
所以“为什么最后是 Transformer”这个问题,答案并不复杂:它把序列建模里“长距离依赖”和“并行计算”两个核心难题同时解决掉了。当然代价也很明显,自注意力的计算复杂度是 O(n²),输入序列越长,计算量和内存占用涨得越快。这也是后来各种稀疏注意力、线性注意力、窗口注意力变体出现的原因。
1.2 一张图记住整体结构:编码器、解码器、注意力、前馈
原始 Transformer 是编码器-解码器结构。左边编码器负责把输入序列编码成一组上下文表示,右边解码器负责根据这些表示和已经生成的内容逐步输出目标序列。
先记四个关键部件:
- 自注意力(Self-Attention):处理序列内部的关系。
- 交叉注意力(Cross-Attention):编码器和解码器之间的桥梁。
- 前馈网络(FFN):对每个位置做非线性变换。
- 残差连接和层归一化:稳定训练并缓解深层网络退化。
实际项目中,现在很多人只使用编码器(如 BERT)来做理解任务,或者只使用解码器(如 GPT 系列)来做生成任务。但无论哪种变体,内部的模块基本都还是 Transformer Block,所以先把原始结构看懂,后面看任何变体都不会太慌。
图解思路是:输入句子经过 Embedding 变成向量,加上位置编码,进入编码器;编码器输出 K 和 V,给到解码器的交叉注意力;解码器输入已经生成的部分序列,经过带掩码的自注意力后,再和编码器输出做交叉注意力;最后通过线性层和 Softmax 输出下一个词的概率。
2. 图解核心组件:Self-Attention 到底是怎么算的
2.1 Q、K、V 是什么:不要死记公式
很多初学者被 Q、K、V 这三个字母劝退。其实可以把它想象成检索过程。
- Query(查询):你想找什么。
- Key(索引):每个候选位置的标签。
- Value(内容):每个候选位置真正携带的信息。
比如翻译“I love you”的时候,要决定“love”和后面哪个词关系最强,就把“love”的 Query 拿去和所有词的 Key 做点积,得到一组分数,然后按分数加权所有词的 Value。分数高,说明这两个位置的相关性强,输出里就会更多地保留那个位置的 Value 信息。
具体到计算上,输入 X 的形状是 [batch, seq_len, d_model],经过三个可学习的权重矩阵 W_Q、W_K、W_V,分别得到 Q、K、V。这三个矩阵通常形状都是 [d_model, d_model],所以 Q、K、V 的形状和 X 一致。这里“可学习”是关键,模型在训练时不断调整 W 矩阵,让注意力关注到真正有用的位置。
2.2 缩放点积注意力的计算过程和维度变化
计算分四步:
- 算注意力分数:Q 乘以 K 的转置,形状变成 [batch, seq_len, seq_len]。
- 缩放:除以 sqrt(d_k),防止点积过大导致 Softmax 梯度消失。
- Softmax:对每一行做归一化,得到注意力权重。
- 加权求和:注意力权重乘以 V,输出形状仍然是 [batch, seq_len, d_model]。
这里 d_k 是每个注意力头的维度,不是 d_model。为什么要除以 sqrt(d_k)?因为当维度变高时,点积的方差会变大,Softmax 之后容易进入饱和区,梯度会非常小。除以 sqrt(d_k) 是让方差回到一个可控范围,这样训练更稳定。
下面给一个精简版 PyTorch 实现,方便对照维度:
import torch import torch.nn as nn class ScaledDotProductAttention(nn.Module): def __init__(self, d_k): super().__init__() self.d_k = d_k def forward(self, q, k, v, mask=None): scores = torch.matmul(q, k.transpose(-2, -1)) / (self.d_k ** 0.5) if mask is not None: scores = scores.masked_fill(mask == 0, float('-inf')) weights = torch.softmax(scores, dim=-1) return torch.matmul(weights, v), weights这里的 mask 是可选参数,解码器或者 padding 位置通常会用到。用 masked_fill 把非法位置置为负无穷,Softmax 后对应位置概率变成 0,就相当于屏蔽掉了不合法位置的注意力。
2.3 多头注意力:为什么要“多头”
多头注意力不是算一次,而是把 d_model 拆成 h 份,每份维度 d_k = d_model / h,分别做注意力计算,最后拼接起来再过一次线性层。
拆成多头有什么用?可以从两个角度理解。
- 表达能力:不同的头可以关注不同类型的模式。有的头关注句法关系,有的头关注指代关系,有的头关注局部相邻词。
- 计算稳定性:每个头的维度变小,矩阵运算规模降低,整体计算量在可控范围。
但多头也不是越多越好。头数增加后,每个头的维度变小,能表达的信息量会下降。同时训练数据不够时,多头之间的重复度会变高。实际使用常见的头数是 8、12、16、32 等,具体要根据模型总维度和任务来决定。如果调参,不要只调头数,要把“头数 × 每个头的维度”和总维度一起看。
3. 位置编码:序列顺序信息怎么进来
3.1 为什么 Transformer 需要位置编码
Self-Attention 本身是集合操作,它不知道 token 在句子里的顺序。把“我打你”和“你打我”的 token 向量输入进去,如果不加位置信息,模型会认为两个句子完全一样。这个问题在 Transformer 里是靠显式注入位置向量解决的。
主流做法有两种:
- 固定位置编码:正余弦函数生成,不需要学习。原版 Transformer 用的就是这个。
- 可学习位置编码:把位置向量当作参数训练,BERT 等模型常用。
固定位置编码的好处是外推到更长序列时相对稳定。可学习编码在训练数据足够时效果通常也不错。到底哪个更好,要看任务和数据量。实际项目里,很多人直接用可学习编码,因为实现简单,而且在训练长度附近表现够用。
3.2 正余弦位置编码为什么这么设计
正余弦位置编码的公式是:
PE(pos, 2i) = sin(pos / 10000^(2i / d_model)) PE(pos, 2i+1) = cos(pos / 10000^(2i / d_model))这里 pos 是 token 的位置,i 是维度下标。设计思路是:不同维度使用不同的频率,低维度频率高,高维度频率低。这样一来,位置向量既能区分不同位置,又会表现出一定的平滑性。对模型来说,它可以学习利用这些位置的相对关系来推断“谁在谁的附近”。
这个公式还有一个好处:不需要额外参数,位置编码是确定的,任何输入长度都能计算,不占用训练参数。
如果使用可学习位置编码,要注意最大位置长度。如果训练时只覆盖 512 个位置,下游突然输入 1024 个 token,超出部分没有对应的位置向量,处理起来会很麻烦。常见做法是预先设置一个足够大的 max_len,或者在推理时用插值方式扩展。
3.3 位置编码的现代演进:绝对位置和相对位置
正余弦编码属于绝对位置编码,它把位置信息直接加到 token 上。BERT 的可学习位置编码也属于这一类。
但很多任务更需要的是“相对位置信息”。比如判断“猫”和“椅子”的关系时,真正重要的是它们之间距离多远,而不是它们在整句话里的绝对位置在第几个。于是出现了很多相对位置编码方案,比如在注意力分数计算时直接加上一个位置偏置项。近年来大模型里常用的 RoPE(旋转位置编码)也是一种相对位置方案,它让 Q 和 K 先旋转一定角度再计算点积,从而把相对位置信息编码进注意力分数里。
如果你在阅读 Transformer 变体论文时遇到位置编码相关的问题,只要记住一条主线:从绝对位置到相对位置,再到把位置信息融合进注意力计算。这个演进方向都是为了解决同一个问题,让模型知道“顺序”和“距离”。
4. 完整拆解一个 Transformer Block:从输入到输出
4.1 编码器层:Norm、注意力、残差、前馈网络的顺序
一个标准的编码器层包含两个子层:
- 第一子层:多头自注意力 + 残差连接 + Layer Norm。
- 第二子层:前馈网络 + 残差连接 + Layer Norm。
很多初学者会纠结 LayerNorm 放在哪里。原版采用“先残差相加,再 LayerNorm”,也就是:
x = LayerNorm(x + Attention(x)) x = LayerNorm(x + FFN(x))后来很多模型使用 Pre-LN,也就是先把 x 归一化再进入子层:
x = x + Attention(LayerNorm(x)) x = x + FFN(LayerNorm(x))两种方式在训练稳定性上有区别。Pre-LN 在深层网络里更容易训练,不一定需要特别的 warmup 策略,因此很多大模型实际采用 Pre-LN。学习的时候,不要把某一派的顺序当成唯一真理,重点理解每个组件的作用。
前馈网络通常是两个线性层加一个激活函数,中间层维度一般是 d_model 的 4 倍。以 d_model=512 为例,FFN 中间层通常是 2048。这个模块对每个词位置独立计算,不做跨位置交互。真正负责“跨位置交互”的就是注意力。
4.2 解码器层:掩码注意力和交叉注意力
解码器层比编码器多了一个子层:与编码器输出做交叉注意力。解码器里有三块:
- 带掩码的自注意力:生成第 t 个词时,只能看到当前位置之前的内容,不能看未来。这个掩码是一个上三角矩阵,未来位置被置为负无穷。
- 交叉注意力:Query 来自解码器,Key 和 Value 来自编码器输出。解码器通过这一步去“读取”输入句子的信息。
- 前馈网络,和编码器一致。
为什么必须加掩码?因为在训练阶段,我们同时把整个目标句子交给模型,如果不加掩码,模型就能“偷看”到后面要预测的词,训练会崩。推理阶段虽然不是整体输入,但掩码保证了训练和推理的一致性。
交叉注意力的出现也解释了一个常见问题:为什么纯解码器模型(如 GPT)没有交叉注意力?因为纯解码器只做文本生成,没有单独的编码器输入来源,它的 K、V 都来自自身历史上下文。而机器翻译这类任务有“源语言”和“目标语言”两个序列,所以才需要交叉注意力做桥梁。
4.3 输出层:如何把向量变成词概率
解码器最后一层输出的向量,经过线性层映射到词表大小,然后经过 Softmax,得到每个候选词的概率分布。训练时用交叉熵损失计算预测和真实下一个词的差距;推理时每次选一个词,把它拼到已有序列中,反复执行直到遇到结束符或达到最大长度。
这部分看起来简单,但推理时有个容易被忽略的细节:解码是逐步进行的,每一步都要重新计算当前所有位置的注意力。如果句子很长,这个逐步计算的过程会越来越慢。工程上常用 KV Cache,把已经算好的 Key 和 Value 缓存起来,避免每一步都重复计算旧位置。这也是很多大模型推理性能优化的核心手段之一。
5. 图解变体:ViT 和 Swin Transformer 怎么把注意力用到图像上
5.1 ViT 的核心思路:把图像切成 Patch
Vision Transformer(ViT)把图像当作“一系列单词”来处理。假设输入是 224x224 的 RGB 图像,先切成 16x16 的 Patch,共 14x14=196 个 Patch。每个 Patch 展平成一个向量,经过一个线性投影变成 d_model 维,这就类似文本里的 token embedding。再加上位置编码和可选的 [CLS] token,最后扔进标准 Transformer 编码器。
ViT 的最大优势是:结构统一,图像和文本可以用同一套架构来建模,这对多模态模型非常友好。但它有一个明显的边界:如果训练数据不够,它不像卷积神经网络那样自带局部性先验,可能不容易收敛。所以 ViT 早期主要是在大规模数据上效果突出,在小数据集上直接训练,效果往往不如 CNN。
5.2 Swin Transformer:窗口注意力解决什么问题
ViT 把整个图像当作一个全局序列,Token 数量一多,自注意力的 O(n²) 复杂度就很难受。比如高分辨率图像,Patch 数量上万后,直接算全局注意力内存会爆炸。
Swin Transformer 的核心改动:只在窗口内部做注意力。窗口大小固定,比如 7x7 个 Patch,这样注意力复杂度只和窗口大小有关,而不是整个图像大小。为了让不同窗口之间能交换信息,Swin 还引入了“移动窗口”机制,在相邻层之间把窗口偏移一些,让原本边界处的 Patch 在新的窗口里能互相看到。
Swin Transformer 还做了类似于 CNN 金字塔的层级结构:前面层的分辨率高、尺寸大,后面层通过 Patch Merging 逐步降低分辨率、增大通道数。这个设计让它做检测、分割等密集型任务时,可以直接套用目标检测框架里的多尺度思路。
5.3 图像任务里为什么需要层级和归纳偏置
很多人问:Transformer 做图像,为什么不能像文本一样直接全局注意力就好?原因主要有两个。
- 图像分辨率高,Token 数远多于文本,全局注意力内存和计算代价太大。
- 图像本身有强局部性,自然图像的多尺度结构决定了模型最好能分层提取特征。
窗口注意力本质上是在没有用卷积的情况下,给模型注入了一种“先看局部、再跨窗口看全局”的偏置。这也是为什么 Swin 在检测和分割上往往比 ViT 更容易落地。
| 对比角度 | ViT | Swin Transformer |
|---|---|---|
| 核心思路 | 全局注意力 | 窗口注意力 + 移动窗口 |
| 计算复杂度 | O(n²) | O(window_size²) |
| 层级结构 | 基本没有多尺度 | 有 Patch Merging 多尺度 |
| 适合任务 | 中等分辨率图像分类等 | 检测、分割等大分辨率场景 |
如果你想学图像 Transformer,我建议先跑 ViT 理解“图像怎么变成 token”,再跑 Swin 理解“窗口和移动窗口”,最后把两者对比起来看。直接啃论文图容易晕,先用现成库跑一个分类任务会更靠谱。
6. 动手理解的最好方式:跑一个最小 Attention 实现和可视化
6.1 环境准备和依赖
要看懂图解,最好的验证方式是亲手跑一次。如果只是想理解 Attention 计算流程,普通电脑就够了,不需要大显卡。常见依赖:
- Python 3.8 及以上
- PyTorch 1.13 或 2.x 均可
- NumPy
- Matplotlib(可视化用)
如果做的是 ViT 或者 Swin 的完整训练,再考虑 GPU 显存问题。ViT-Base 在 224x224 分辨率下,如果用完整训练,显存需求会比较高;但如果只是前向推理一个 Batch,8G 显存通常也可以做实验。低配置机器能跑通 Demo,不代表适合批量训练,这一点要提前有预期。
6.2 用 PyTorch 手写一个精简版 Attention
先跑一个随机输入的最小样例,确认形状没问题:
import torch import torch.nn.functional as F def attention(q, k, v, mask=None): d_k = q.size(-1) scores = torch.matmul(q, k.transpose(-2, -1)) / (d_k ** 0.5) if mask is not None: scores = scores.masked_fill(mask == 0, float('-inf')) weights = F.softmax(scores, dim=-1) return torch.matmul(weights, v), weights batch, seq_len, d_model, heads = 2, 8, 64, 4 d_k = d_model // heads x = torch.randn(batch, seq_len, d_model) w_q = torch.randn(d_model, d_model) w_k = torch.randn(d_model, d_model) w_v = torch.randn(d_model, d_model) q = x @ w_q k = x @ w_k v = x @ w_v # 拆成多头 q = q.view(batch, seq_len, heads, d_k).transpose(1, 2) k = k.view(batch, seq_len, heads, d_k).transpose(1, 2) v = v.view(batch, seq_len, heads, d_k).transpose(1, 2) output, attn_weights = attention(q, k, v) print(output.shape) # [batch, heads, seq_len, d_k] print(attn_weights.shape) # [batch, heads, seq_len, seq_len]这个例子只是为了验证维度,不是完整实现。建议跑的时候把输出形状打印出来,对照上面的说明看一遍,比单纯看文章有效很多。
6.3 检查注意力权重、输出形状和梯度
如果自己实现,容易出现三类问题:
- 形状错乱:多头拼接后没有恢复到 d_model,最后输出维度对不上。
- Softmax 维度选错:应该对最后一维 seq_len 方向做归一化,不是对 batch 方向。
- 掩码位置处理错误:负无穷没有准确覆盖,导致未来位置仍被关注。
判断标准很简单:
- 任何中间输出的最后一个维度,最后都应该能对应回 d_model。
- 注意力权重每一行的和应该等于 1,在有效位置内。
- 训练时梯度不为 NaN,loss 能稳定下降。
如果发现注意力权重分布非常平均,先不要急着改结构。这通常意味着训练不充分、残差和 Norm 没有正确叠加,或者输入数据本身没有区分度。
6.4 可视化注意力的思路
可视化可以用 Matplotlib 把 attn_weights 画成热力图。横轴和纵轴都是序列位置,颜色越深代表注意力权重越高。这样可以直观看到某个 token 主要关注了哪些位置。
我一般会这样检查:
- 如果是机器翻译样例,看解码时是否关注了源句中的对应词。
- 如果有一段长文本,看模型是否关注了关键实体而忽略无关词。
- 多头模式下,把不同头分开画,通常会看到有的头关注局部、有的头关注远距离指代。
训练充分的模型注意力图往往比随机初始化清晰很多。这也是“图解 Transformer”最有意思的部分,把一个公式变成看得见的矩阵。
7. 从图解到实战:常见坑和排查顺序
7.1 维度对不上:先查 batch、seq_len、d_model
Transformer 代码里最常见错误就是维度。遇到报错,先打印 Q、K、V 的形状,再确认拆分多头之后的形状。遵循一个原则:形状错误时先改维度,再看逻辑,不要一上来就调学习率或换模型结构。
| 报错现象 | 优先排查 |
|---|---|
| 维度不匹配 | Q/K/V 形状、多头拆分和拼接方式 |
| loss 为 NaN | 学习率、输入数据、掩码 |
| 注意力权重几乎一样 | 训练不充分、随机初始化、Norm 顺序 |
| 内存溢出 | 序列长度、batch size、头数、d_model |
7.2 训练不收敛:先看掩码、学习率和数据格式
如果 loss 不下降,或者掉到 NaN,排查顺序是:
- 输入数据是否包含 NaN 或异常值。先检查数据,再改模型。
- 损失函数和标签是否对齐。文本任务里经常是标签错位。
- 学习率是否过大。Transformer 对学习率比较敏感,常用 warmup + 动态衰减。
- 是否加入了正确掩码。padding 位置不应该参与注意力,否则模型会把无意义位置也纳入计算。
- 初始化策略。不同框架的默认初始化不同,不要随意叠加自定义缩放。
如果只是学习而不是生产,建议先跑小数据集、小模型、短序列。把一条样例跑通,再慢慢扩大范围。
7.3 要不要从头手写:什么时候用现成框架
“手撕 Transformer”确实能加深理解,尤其是准备面试时,自己动手写一遍和只看文章完全不一样。但从工程角度,生产项目不需要自己实现 Attention,直接用 PyTorch、Hugging Face 等现成库更稳。
原因有几个:
- 现成库考虑了大量边界条件、数值稳定性和计算优化。
- 手写实现可能在短序列上能跑,但长序列、批处理、量化场景下性能和稳定性都差很多。
- 后续维护成本更高。
我的建议是:把“手写”当作学习手段,把“现成库”当作落地手段。学习阶段一定要写,生产阶段不一定要写。
7.4 学习 Transformer 的一个合理路径
如果你刚入门,可以按这个顺序走:
- 先理解整体结构:输入、编码器、解码器、输出。
- 再理解自注意力:Q/K/V、缩放点积、多头。
- 自己跑一个最小注意力样例,确认输出形状。
- 复现一个单层编码器 Block,加入残差和 LayerNorm。
- 看 ViT 或 Swin,理解怎么把图像转成 token。
- 最后再看源码、看论文,结合公式对照代码。
如果时间紧张,前四步就足以应付绝大多数面试中的架构问题。后面两步是为了视觉和多模态场景准备的。
最后说一个个人体会。Transformer 图解的价值,不在于背出某个公式,而在于你能在写代码时预判“这一步输出形状是什么”“这个掩码会影响哪些位置”“这个参数拉大后资源占用会怎么涨”。把这三个问题想清楚,比多背十篇论文都有用。