做文本建模或者大模型方向的朋友,早晚会绕不开一个东西——Transformer。我第一次对着那篇原始论文里的模型结构图看的时候,光是把 Query、Key、Value 三个矩阵在脑子里对齐就花了大半个下午,更别说后面多头拆分、位置编码、残差和归一化到底放在哪一层这种细节。后来自己从零手写了一遍,跑通了小规模机器翻译任务,才算真正把 Transformer 模型结构吃透。这篇就是我整理的一份超详细解读,从整体骨架到每一个子模块,再到能直接抄的 PyTorch 实现和踩过的坑,尽量说人话。不管你是刚入门想搞懂注意力机制的新手,还是已经能调库但说不清结构细节的熟手,应该都能从里面捞到点东西。
1. 从整体架构看Transformer的骨架设计
很多人一上来就钻进注意力公式里,其实先把整体骨架啃清楚,后面很多细节是顺理成章的。Transformer 的原始设计是一个编码器-解码器结构,专门为序列到序列任务准备的。但今天大家嘴上说的 Transformer,往往泛指这一整类基于自注意力的结构,包括只保留编码器的 BERT 系、只保留解码器的 GPT 系。
1.1 编码器与解码器的分工到底差在哪
编码器的任务是"读懂输入",它由 N 个完全相同的层堆叠而成,每一层里有两个子模块:一个是多头自注意力,一个是前馈网络。注意这里只有自注意力,没有掩码,因为编码器处理的是完整的输入句子,每个位置都能看到句子里所有其他位置,这叫双向可见。
解码器的任务则是"生成输出",它的每一层里有三个子模块:带掩码的多头自注意力、编码器-解码器注意力、前馈网络。第一个自注意力加了因果掩码,保证第 t 个位置只能看到前 t 个位置,不能偷看未来。第二个注意力比较特殊,它的 Query 来自解码器当前状态,Key 和 Value 来自编码器的输出,这一步是把"读懂的输入"和"正在生成的输出"对齐起来。
我个人的理解是,编码器像是一个把整段话压缩成向量表示的过程,解码器像是一个一边看压缩表示一边逐字往外吐的过程。搞清这个分工,后面看到 mask 为什么只加在解码器第一层、为什么编码器-解码器注意力不需要因果掩码,就不会迷糊了。
1.2 为什么选择"堆叠"而不是一味"加宽"
原始论文里编码器和解码器都堆了 6 层,模型维度 d_model 是 512。这里有个很关键的设计哲学:深度优先于宽度。堆叠层数带来的是抽象层级的提升,底层关注局部词形和短距离依赖,中层开始捕捉句法结构,高层才能表征语义和长距离关系。而单纯加宽只是增加单层容量,很难形成这种逐级抽象。
当然深度也不是越多越好。层数一多,梯度消失和训练不稳定的问题就来了,这也是后面 Post-LN 被 Pre-LN 逐渐取代的重要原因之一。我在实际做小任务时,4 层到 6 层通常够用;做大一点的任务,12 层是个常见起点。层数选择本质上是在表达能力和训练难度之间找平衡。
1.3 结构总览与全文的维度约定
为了后面推导不混乱,这里先把符号约定死,全文都按这套来:
- B:batch size,一次喂多少条样本
- L:序列长度,也就是 token 数量
- d_model:模型隐藏维度,比如 512
- h:注意力头数,比如 8
- d_k:每个头的维度,满足 d_k = d_model / h
- d_ff:前馈网络中间层维度,通常是 4 * d_model
举个例子,输入张量形状是 (B, L, d_model) = (2, 10, 512),h=8,那么 d_k = 64。整套结构从头到尾就是在这几个维度之间来回搬运、拆分、合并。把这张"维度地图"记在脑子里,看任何 Transformer 变体都会快很多。
2. 核心组件逐个拆解:从多头注意力到前馈网络
骨架理清了,接下来一个组件一个组件地拆。这一部分我会尽量把每一步的形状变化写出来,因为形状对不上是新手写代码时最高频的报错来源。
2.1 自注意力机制的计算全过程与维度推导
自注意力的本质是一句话:让序列里每个位置,去综合其他所有位置的信息,综合的权重由相关性决定。具体分三步。
第一步,把输入 X 分别乘上三个可学习矩阵 W_q、W_k、W_v,得到 Q、K、V。这三个矩阵形状都是 (d_model, d_model)。所以 Q = X @ W_q,形状仍是 (B, L, d_model)。
第二步,算注意力分数。用 Q 乘 K 的转置,得到形状 (B, L, L) 的分数矩阵,再除以根号 d_k。为什么要除以根号 d_k?因为当维度较大时,Q 和 K 的点积数值会随维度增长而变大,softmax 之后会变得非常尖锐,几乎只有一个位置接近 1,其余接近 0,梯度会变得极小。除以根号 d_k 相当于把方差拉回到 1 附近,这是数值稳定性的保障,不是可有可无的装饰。
第三步,对分数做 softmax 得到注意力权重,再乘 V,输出形状 (B, L, d_k)。整个过程用公式写就是:Attention(Q,K,V) = softmax(QK^T / sqrt(d_k)) V。这里有个容易忽略的点:softmax 是沿最后一维做的,也就是每个 query 位置对全体 key 位置归一化,方向反了结果就完全错了,我在调试时被这个坑过不止一次。
2.2 多头注意力的拆分与合并操作
单头注意力只能捕捉一种相关性模式,多头则是让模型同时从多个"视角"去看输入。实现上并不是真的跑了 h 次独立的注意力,而是把 d_model 维度切成 h 份,每份 d_k 维,h 个头并行计算,最后再拼回来。
具体流程是这样的:线性投影后得到 (B, L, d_model),用 view 重塑成 (B, L, h, d_k),再用 transpose 把 h 维换到前面,变成 (B, h, L, d_k)。这样每个头就是独立的一层二维注意力。算完以后输出 (B, h, L, d_k),transpose 回 (B, L, h, d_k),用 contiguous 保证内存连续,再 view 回 (B, L, d_model),最后乘输出矩阵 W_o。
这里有个技术细节值得强调:transpose 之后内存布局变了,直接 view 会报错,必须先 .contiguous()。这个报错信息往往很长很吓人,其实原因就一句话——不连续的内存没法直接拉平。还有一点,h 必须整除 d_model,否则切不匀,实现里通常加一句 assert 兜底。
提示:多头的意义不只是"多个视角",它还降低了单头的计算复杂度总开销,因为每个头只在 d_k 维度上算,总计算量和单头差不多。这也是它比"堆 h 个完整注意力再拼"更划算的原因。
2.3 位置编码:正弦编码与可学习编码的取舍
自注意力有个先天缺陷:它对输入顺序是无感的。你把句子里两个词调换位置,只要它们携带的内容不变,输出几乎不变。这对语言来说是灾难,因为"猫追狗"和"狗追猫"完全不同。解决办法就是位置编码,把位置信息注入进去。
原始论文用的是正弦位置编码,公式是:偶数维用 sin(pos / 10000^(2i/d_model)),奇数维用 cos(pos / 10000^(2i/d_model))。其中 pos 是位置,i 是维度索引。这套设计的巧妙之处在于,任意位置 pos+k 的编码可以表示成 pos 编码的线性变换,这让模型有能力外推到更长的序列。
另一种常见做法是可学习位置编码,就是给每个位置准备一个可训练向量,BERT 用的就是这个。它实现简单、效果稳定,但最大长度在训练时就固定死了,想扩展到更长序列得重新处理。我在短文本任务里通常用可学习编码,处理变长或需要外推的场景更倾向正弦编码或旋转位置编码。选哪个没有绝对答案,看任务。
2.4 残差连接与LayerNorm的位置之争
每个子模块外面都套了一层"残差连接 + 层归一化",写作 LayerNorm(x + Sublayer(x))。残差连接的作用是给梯度开一条高速公路,让深层网络也能训得动;LayerNorm 则是稳定每层的激活分布,加速收敛。
但这里有个被反复讨论的细节:归一化到底放在残差之前还是之后。原始论文是 Post-LN,也就是先做子层再相加再归一化。后来研究发现 Post-LN 在层数多时训练很不稳定,需要精细的学习率预热,于是 Pre-LN 流行起来,变成先归一化再做子层。Pre-LN 训练更稳、对学习率不那么敏感,代价是最终效果在同等层数下可能略逊,需要配合更深的堆叠来弥补。现在主流大模型几乎清一色 Pre-LN 或其变体,如果你自己训练遇到深层不收敛,先试试把归一化提到前面。
3. 前馈网络、激活函数与归一化细节
注意力负责在位置之间"通信",前馈网络负责在每个位置上"独立加工"。这两个部分交替出现,构成了 Transformer 层的基本节奏。这一部分聊几个容易被忽视但影响很大的细节。
3.1 前馈网络为什么是4倍扩张
前馈网络结构很朴素:两层线性变换中间夹一个激活函数,写作 FFN(x) = W_2 · activation(W_1 · x)。关键在于中间层的维度 d_ff 通常是 d_model 的 4 倍。原始论文里 d_model=512,d_ff=2048,正好是4倍。
为什么是4倍而不是2倍或8倍?我理解这是一个经验性的容量平衡点。注意力层已经负责了跨位置的信息聚合,前馈层需要足够的宽度来对每个位置做非线性变换,把特征投影到更高维再压回来,形成一个"瓶颈-扩张-瓶颈"的结构,类似自编码器的思路,能增强表达能力。倍数太小,非线性表达受限;倍数太大,参数量和计算量陡增,收益递减。4 倍是大量实验下来的常用值,T5 用过 8 倍,也有用 2.67 倍的变体,实际可以调。
顺带说一句,这部分参数量其实很可观。对 d_model=512、d_ff=2048 的单层来说,FFN 参数量约 2 × 512 × 2048 ≈ 200 万,而一层多头注意力四个投影矩阵加起来才约 4 × 512 × 512 ≈ 100 万。也就是说,前馈网络占了单层参数的大头,这一点很多人没意识到。
3.2 GELU与ReLU的选用逻辑
激活函数早期用 ReLU,简单高效。后来 GELU 逐渐成为 Transformer 的标配,BERT、GPT 系列都用它。GELU 的形式是对输入做高斯分布的累积概率加权,可以粗略理解为"平滑版的 ReLU",在零点附近是光滑过渡的,负值区域也不是直接截断为零,而是保留一小部分。
这个平滑性的好处是梯度更连续,训练更稳定,尤其深层网络里更明显。代价是计算比 ReLU 稍贵。如果你在做资源紧张的部署,ReLU 或它的变体 Swish、SiLU 都是可以的替代,实测差异在小任务上未必看得出来。我的经验是,先用 GELU 跑通,真到了要抠性能再换,别一开始就在激活函数上纠结。
3.3 归一化层的两种主流实现
LayerNorm 和 BatchNorm 的区别经常被问。BatchNorm 是在 batch 维度上统计均值和方差,依赖 batch 大小,序列任务里 batch 内的长度还不一致,用它很别扭。LayerNorm 则是针对每个样本、每个位置,在自己这一条特征向量上做归一化,和 batch 大小无关,非常适合变长序列。
LayerNorm 的计算是:(x - mean) / sqrt(var + eps) * gamma + beta。其中 mean 和 var 是沿最后一维(特征维)算的,gamma 和 beta 是可学习的缩放和平移参数,初始化为 1 和 0。eps 是个很小的数,防止除零,一般取 1e-5 或 1e-6。
再进阶一点,现在不少模型把 LayerNorm 换成了 RMSNorm,去掉了减均值和 beta 那一步,只保留缩放,计算更快,效果基本持平。这是工程优化里的常见取舍,属于知道就行、不必强求的细节。
4. 从零手写一个Transformer
光看懂结构不够,自己写一遍才能发现所有藏起来的坑。这一部分给一份可以跑的 PyTorch 实现,顺便把关键位置的形状标注清楚。我用的是自己教学时反复改过的版本,去掉了花哨的东西,突出主干。
4.1 张量维度约定与整体代码骨架
先定好超参:d_model=512,n_heads=8,d_ff=2048,层数 6,dropout 0.1。整体代码分成四块:多头注意力、位置编码、编码器层、解码器层,然后拼成完整模型。先看多头注意力,这是最核心、也最容易写错的部分。
import torch import torch.nn as nn import math class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads, dropout=0.1): super().__init__() assert d_model % n_heads == 0, "d_model 必须能被 n_heads 整除" self.d_model = d_model self.n_heads = n_heads self.d_k = d_model // n_heads self.w_q = nn.Linear(d_model, d_model) self.w_k = nn.Linear(d_model, d_model) self.w_v = nn.Linear(d_model, d_model) self.w_o = nn.Linear(d_model, d_model) self.dropout = nn.Dropout(dropout) def forward(self, q, k, v, mask=None): B = q.size(0) # 投影 + 拆多头: (B, L, d_model) -> (B, h, L, d_k) Q = self.w_q(q).view(B, -1, self.n_heads, self.d_k).transpose(1, 2) K = self.w_k(k).view(B, -1, self.n_heads, self.d_k).transpose(1, 2) V = self.w_v(v).view(B, -1, self.n_heads, self.d_k).transpose(1, 2) # 缩放点积注意力: (B, h, L, L) scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores = scores.masked_fill(mask == 0, -1e9) attn = torch.softmax(scores, dim=-1) attn = self.dropout(attn) # 加权求和后合并多头: (B, L, d_model) out = torch.matmul(attn, V) out = out.transpose(1, 2).contiguous().view(B, -1, self.d_model) return self.w_o(out)这段里有三个点必须盯住。第一,view 之前要能整除,assert 是保险。第二,transpose 换轴后必须 contiguous 再 view,否则报错。第三,masked_fill 用的值要足够小,-1e9 是常用做法,别用 0,因为 softmax 后 0 还会分到权重。
4.2 位置编码与前馈网络的实现
位置编码可以写成一个预计算的矩阵,训练时按序列长度取前 L 行加到词嵌入上。注意是加,不是拼接。
class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=5000): super().__init__() pe = torch.zeros(max_len, d_model) position = torch.arange(0, max_len).unsqueeze(1).float() div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) self.register_buffer('pe', pe.unsqueeze(0)) def forward(self, x): return x + self.pe[:, :x.size(1)]这里 div_term 用的就是 10000^(2i/d_model) 的倒数形式,写成 exp 能避免数值溢出。前馈网络更简单,就是两个线性层加激活:
class FeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout=0.1): super().__init__() self.net = nn.Sequential( nn.Linear(d_model, d_ff), nn.GELU(), nn.Dropout(dropout), nn.Linear(d_ff, d_model) ) def forward(self, x): return self.net(x)然后是编码器层,把注意力和前馈串起来,每个外面套 Pre-LN 残差:
class EncoderLayer(nn.Module): def __init__(self, d_model, n_heads, d_ff, dropout=0.1): super().__init__() self.attn = MultiHeadAttention(d_model, n_heads, dropout) self.ffn = FeedForward(d_model, d_ff, dropout) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def forward(self, x, mask=None): x = x + self.dropout(self.attn(self.norm1(x), self.norm1(x), self.norm1(x), mask)) x = x + self.dropout(self.ffn(self.norm2(x))) return x注意这里三次传入的都是 norm1(x),这是自注意力的标志——Q、K、V 同源。解码器层会多一个交叉注意力,其中 Q 来自解码器自身,K 和 V 来自编码器输出,这个区分是写对解码器的关键。
4.3 掩码生成与训练调试要点
因果掩码是个下三角矩阵,形状 (L, L),下三角(含对角线)为 1,其余为 0。表示第 t 个位置只能看到不超过 t 的位置。生成方式很简单:
def causal_mask(size): mask = torch.tril(torch.ones(size, size)).bool() return mask.unsqueeze(0).unsqueeze(0) # (1, 1, L, L)同时别忘了处理 padding 掩码,把补齐位置的注意力屏蔽掉,否则模型会去关注无意义的填充符。训练时的几个调试心得:学习率要配预热,前几千步线性升到峰值再衰减,这是原始论文的做法,没有预热深层很容易发散;dropout 在小数据集上别省,能显著缓解过拟合;如果 loss 一直不降,先检查 mask 方向有没有反,再看位置编码有没有真的加进去,这两处是最隐蔽的错。
5. 常见问题与排查技巧实录
这一部分是实打实用时间换来的经验。很多问题看起来是"训练技巧",根子其实在结构实现上。我整理成速查表,方便你对症下药。
5.1 形状不匹配类问题
最常见的一类报错就是形状对不上。我按照出现频率列一下,对应的原因和修法都写清楚。
| 报错现象 | 可能原因 | 修复办法 |
|---|---|---|
| view 处报 size 不匹配 | transpose 后内存不连续 | 加 contiguous() 再 view |
| reshape 维度错误 | d_model 不能被 n_heads 整除 | 调整头数或维度,加 assert |
| 矩阵乘维度冲突 | 注意力里 K 忘了转置 | 对 K 用 transpose(-2, -1) |
| 广播失败 | mask 形状与 scores 不一致 | mask 补成 (B,1,L,L) 或 (1,1,L,L) |
| 拼接后维度翻倍 | concat 时算错头维度 | 确认拼回的是 h*d_k = d_model |
这类问题九成是维度换算没记清。我的习惯是在每个模块入口和出口都打印一次 shape,跑一遍小数据,确认整条链路都对,再上大规模训练,能省下大量瞎猜的时间。
5.2 训练不收敛与效果异常排查
形状对了但训练不收敛,或者效果诡异,也有套路可循。第一个要怀疑的是掩码方向。如果因果掩码打反了,模型看到的就是未来信息,训练 loss 可能很低,但推理时完全崩,表现是"训练集无敌、生成一塌糊涂"。第二个是学习率,没有预热、峰值过大,深层模型会在前几百步就崩掉。
还有就是数值精度问题。用半精度训练时,注意力里的分数容易溢出,建议保留缩放那一步,并且对 softmax 输入做一下 clamp。位置编码加错了也不会立刻报错,只是模型学不到顺序,表现是对词序不敏感,你可以在一个小样本上故意打乱词序,看输出是否变化,变化不大就说明位置信息没生效。
注意:如果输出退化成所有位置生成同一个 token,先检查 softmax 维度是不是用错,再确认 logits 有没有被 mask 全屏蔽导致全是 -1e9。全屏蔽时 softmax 会输出均匀分布,表现为随机或重复,很迷惑人。
5.3 分模块定位问题的思路
面对一个不工作的 Transformer,别一上来就调超参,先做结构自检。我的顺序是:先用一条长度为 2 的假数据跑前向,确认不报错;再把模型过拟合一个极小的批次,比如 8 条样本,如果连这点数据都过拟合不了,说明结构或损失肯定有 bug;最后才加正则、调学习率。
这个"先能过拟合、再谈泛化"的思路非常实用。因为过拟合小数据是模型具备基本表达能力的必要条件,做不到就说明前向或反向有问题,跟超参无关。我见过太多人一开始就猛调学习率,结果 bug 在结构里,调多久都没用。
6. 结构变体与现代工程取舍
原版 Transformer 是 2017 年的设计,这些年结构上做了不少演进。理解变体不是为了追新,而是为了知道每个设计选择背后的权衡,面试和实际选型都用得上。
6.1 主流变体改了哪些结构点
按改动部位来梳理会更清楚。归一化方面,Post-LN 基本被 Pre-LN 取代,部分模型换成 RMSNorm。位置编码方面,可学习编码和正弦编码之外,旋转位置编码(RoPE)在长文本场景流行起来,它通过旋转矩阵把相对位置信息编码进注意力,外推性更好。注意力计算方面,为了省显存和算力,出现了分组查询注意力和多查询注意力,让多个 Query 头共享少量 Key、Value 头,大幅降低推理时的 KV 缓存开销。
前馈网络这块也有变化,出现了用门控机制的变体,比如把 FFN 拆成两路,一路做门控,能提升效果但增加参数。视觉领域的 Swin Transformer 则把注意力限制在局部窗口内,降低计算量的同时引入层级结构。这些都是同一个自注意力内核在不同约束下的工程取舍,理解原版结构是看懂它们的前提。
6.2 被问最多的几个结构问题
最后列几个我在交流里被问得最多的问题,答案其实都藏在前面。
第一个:为什么注意力要缩放?答案是高维点积方差过大导致 softmax 饱和、梯度消失。
第二个:多头和单头的计算量谁大?答案是差不多,因为每个头维度被切小了,总运算量基本持平,收益主要来自表达多样性。
第三个:为什么解码器第一层要加因果掩码,交叉注意力不用?因为解码器生成时必须保证不看未来,而交叉注意力看的是完整编码结果,本来就没有未来可言。
第四个:残差连接到底解决什么?最直接的是让梯度能顺畅回传,深层网络才训得动,同时让每层只需学习"增量",优化目标更简单。
把这些"为什么"能讲清楚,比背公式有用得多。我自己也是写了、调了、错了,才慢慢把这些点连成一片的。