Transformer 这个架构,我从 2019 年开始断断续续接触,最初看论文的时候也是一头雾水,什么 Self-Attention、Multi-Head、Positional Encoding,每个词都认识,连在一起就不知道在说什么。后来逼着自己手写了一遍代码,又拿它做了几个时序预测的项目,才算真正把里面的结构吃透了。这篇文章我打算把 Transformer 拆开揉碎,从整体架构到每一个子模块,把 Encoder、Decoder、Multi-Head Attention、Feed Forward 这些核心组件讲清楚,同时把位置编码、残差连接、层归一化这些容易被忽略但极其关键的细节也一并说透。不管你是刚入门的新手,还是已经跑过几个模型但对内部机制还不太确定的朋友,看完应该都能有一个清晰的认识。
1. Transformer 整体架构拆解与设计思路
1.1 为什么 Transformer 要设计成 Encoder-Decoder 结构
Transformer 最初是为机器翻译任务设计的,输入一种语言,输出另一种语言。这个场景天然就适合 Encoder-Decoder 架构:Encoder 负责理解输入序列,把它压缩成一组包含语义信息的表示;Decoder 负责根据这组表示,一步步生成目标序列。
你可以把它想象成一个翻译官的工作流程。Encoder 就像翻译官先把整段外文读完,在脑子里形成完整的理解;Decoder 就像翻译官开始用中文写译文,每写一个字都要参考原文的理解,同时还要看自己前面已经写了什么。
但这里有个关键点:Encoder 和 Decoder 内部的结构并不完全一样。Encoder 是双向的,每个位置都能看到整个输入序列;Decoder 是单向的,每个位置只能看到当前位置及之前的位置。这个差异直接决定了它们内部 Attention 的计算方式不同,后面我会详细展开。
1.2 Encoder 和 Decoder 的堆叠数量怎么选
原始论文《Attention Is All You Need》里,Encoder 和 Decoder 各堆叠了 6 层。这个数字不是随便定的,也不是必须的。6 层是一个在效果和计算成本之间比较平衡的选择。
实际项目中,这个层数是可以调整的。比如 BERT-base 用了 12 层 Encoder,BERT-large 用了 24 层。GPT 系列也是类似,层数从 12 到 96 不等。层数越多,模型的表达能力越强,但计算量和显存占用也会线性增长。
我个人的经验是:如果你在做小规模的任务,比如文本分类或者简单的时序预测,4 到 6 层通常就够了。如果数据量很大、任务很复杂,再考虑加到 12 层甚至更多。但要注意,层数增加带来的收益是递减的,而且训练难度会显著上升。
1.3 残差连接和层归一化的位置选择
Transformer 每个子层(Self-Attention、Feed Forward)外面都包了一层残差连接和层归一化。原始论文用的是 Post-LN,也就是先做残差加法,再做 Layer Normalization:
output = LayerNorm(x + Sublayer(x))但后来的研究发现,Pre-LN 更稳定,也就是先做 Layer Normalization,再做子层计算,最后残差加法:
output = x + Sublayer(LayerNorm(x))这两种方式的区别在于梯度传播的路径。Post-LN 在深层网络中容易出现梯度消失或爆炸,训练时需要非常小心地调学习率和 warmup 策略。Pre-LN 则稳定得多,很多现代 Transformer 变体都默认用 Pre-LN。
我踩过的坑:早期用 Post-LN 训练一个 12 层的模型,不加 warmup 直接训,loss 直接飞了。后来换成 Pre-LN,同样的学习率就能稳定训练。所以如果你在复现论文结果时遇到训练不稳定的问题,可以先检查一下 LN 的位置。
2. Multi-Head Attention 的核心机制与实操细节
2.1 Self-Attention 到底在算什么
Self-Attention 的核心思想是:序列中每个位置的表征,都应该由整个序列中所有位置的表征加权求和得到。权重的大小取决于当前位置和其他位置的关联程度。
具体计算过程分三步:
- 把输入向量分别乘以三个权重矩阵 W_Q、W_K、W_V,得到 Query、Key、Value 三个矩阵。
- 用 Q 和 K 做点积,除以 sqrt(d_k) 做缩放,再经过 Softmax 得到注意力权重。
- 用注意力权重对 V 加权求和,得到输出。
用公式表示就是:
Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V这里的 sqrt(d_k) 缩放非常关键。如果不做缩放,当 d_k 很大时,Q 和 K 的点积会变得很大,Softmax 的输出会趋近于 one-hot,梯度会变得极小,训练就会停滞。除以 sqrt(d_k) 可以把点积的方差控制在 1 左右,保证 Softmax 的输出在一个合理的范围内。
2.2 为什么要用 Multi-Head 而不是 Single-Head
单个 Attention 头只能学到一种注意力模式。但语言中的关系是多种多样的:有的位置关注语法结构,有的位置关注语义关联,有的位置关注位置邻近关系。Multi-Head Attention 就是让模型同时学习多种注意力模式。
具体做法是把 Q、K、V 分别投影到 h 个低维子空间,在每个子空间里独立做 Attention,最后把 h 个头的输出拼接起来,再经过一个线性变换。
原始论文里 h=8,每个头的维度是 d_model/h=64。这样总的计算量和单个 d_model 维度的 Attention 差不多,但表达能力更强。
我实测下来的感受是:头数不是越多越好。8 个头在大多数任务上表现都不错。如果头数太多,每个头的维度太小,反而学不到有意义的关系。如果头数太少,又退化成 Single-Head 了。一般建议每个头的维度不要低于 32。
2.3 Masked Multi-Head Attention 的实现要点
Decoder 里的 Self-Attention 需要加 Mask,保证每个位置只能看到自己和之前的位置。这个 Mask 是一个上三角矩阵,对角线及以下为 0,对角线以上为负无穷。
实现的时候,通常是在 Softmax 之前把 Mask 加到注意力分数上:
attn_scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k) attn_scores = attn_scores.masked_fill(mask == 0, float('-inf')) attn_weights = torch.softmax(attn_scores, dim=-1)这里有个细节:masked_fill 用的值是负无穷,而不是 0。因为 Softmax 之后,负无穷对应的权重会变成 0,而如果直接填 0,Softmax 之后仍然会有非零权重,那就起不到 Mask 的作用了。
还有一个容易出错的地方:Mask 的形状要和注意力分数的形状匹配。注意力分数的形状是 (batch_size, num_heads, seq_len, seq_len),Mask 需要广播到这个形状。我见过不少人在这一步因为维度不匹配而报错。
2.4 Cross-Attention 在 Decoder 中的作用
Decoder 中间还有一个 Cross-Attention 层,它的 Q 来自 Decoder 上一层的输出,K 和 V 来自 Encoder 的输出。这个设计让 Decoder 在生成每个词的时候,都能参考 Encoder 对输入序列的完整理解。
你可以把它理解为:Decoder 在写译文的时候,每写一个词都会回头看一眼原文,看看哪些部分和当前要写的词最相关。这就是 Cross-Attention 的作用。
需要注意的是,Cross-Attention 的 K 和 V 是共享的,也就是所有 Decoder 层用的都是同一个 Encoder 输出。但 Q 是每层独立的,来自上一层 Decoder 的输出。
3. Feed Forward Network 与位置编码的深度解析
3.1 Feed Forward 层为什么是两层而不是一层
Transformer 里的 Feed Forward 层其实就是一个两层的全连接网络,中间加了 ReLU 激活:
FFN(x) = max(0, xW_1 + b_1)W_2 + b_2第一层把维度从 d_model 扩展到 d_ff,第二层再投影回 d_model。原始论文里 d_ff=2048,d_model=512,扩展倍数是 4 倍。
为什么要先扩展再压缩?我的理解是:Attention 层主要在做信息的加权聚合,而 Feed Forward 层在做非线性变换,给模型提供更强的表达能力。扩展到更高维度,可以让模型在这个高维空间里学到更复杂的特征组合,然后再压缩回原维度,保持整个网络的维度一致。
这个 4 倍的扩展比例在大多数情况下都够用。如果任务特别复杂,可以适当增大,但要注意参数量会平方级增长。因为 FFN 的参数量是 2 * d_model * d_ff,当 d_ff 增大时,参数量线性增长,但计算量也是线性增长。
3.2 位置编码的计算方式和选择
Transformer 本身没有循环结构,也没有卷积结构,如果不加位置编码,它就无法区分序列中不同位置的词。所以位置编码是必须的。
原始论文用的是正弦位置编码:
PE(pos, 2i) = sin(pos / 10000^(2i/d_model)) PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))这个设计的好处是:对于任意固定的偏移量 k,PE(pos+k) 可以表示为 PE(pos) 的线性函数。这意味着模型可以很容易地学到相对位置关系。
但后来很多模型改用可学习的位置编码,比如 BERT。可学习的位置编码就是把每个位置的编码当成一个可训练的参数,让模型自己学。这种方式更灵活,但泛化到训练时没见过的长度时表现会差一些。
我个人的经验是:如果序列长度比较固定,可学习的位置编码效果通常更好;如果序列长度变化很大,或者需要泛化到更长的序列,正弦编码更稳妥。另外还有相对位置编码、旋转位置编码等变体,这些在长序列场景下表现更好,但实现复杂度也更高。
3.3 词嵌入矩阵的初始化与共享
Transformer 的词嵌入矩阵通常是随机初始化的,然后随着训练更新。但有一个技巧:Encoder 和 Decoder 的词嵌入矩阵可以共享,而且词嵌入矩阵和最后的输出投影矩阵也可以共享。
共享的好处是减少参数量,而且可以让词嵌入和输出投影学到一致的表征。这个技巧在机器翻译任务里很常用,效果也不错。
初始化方面,一般用正态分布,均值 0,标准差 0.02 或者 1/sqrt(d_model)。标准差太大会导致训练初期梯度爆炸,太小又会导致梯度消失。我试过用 Xavier 初始化,效果也还可以,但不如正态分布稳定。
4. 实操过程中的关键步骤与参数计算
4.1 从零手写一个 Transformer 的完整流程
手写 Transformer 是理解它最好的方式。我建议按以下顺序来:
- 先实现 Scaled Dot-Product Attention,这是最核心的部分。
- 再实现 Multi-Head Attention,把多个 Attention 头拼起来。
- 然后实现 Positional Encoding,加到词嵌入上。
- 接着实现 Encoder Layer 和 Decoder Layer。
- 最后把多层堆叠起来,加上输出层。
每一步都要写单元测试,确保输出形状和数值都正确。比如 Scaled Dot-Product Attention 的输出形状应该是 (batch_size, seq_len, d_v),Multi-Head Attention 的输出形状应该是 (batch_size, seq_len, d_model)。
我当初手写的时候,在 Multi-Head Attention 的维度变换上卡了很久。Q、K、V 的形状是 (batch_size, seq_len, d_model),需要先 reshape 成 (batch_size, seq_len, num_heads, d_k),再 transpose 成 (batch_size, num_heads, seq_len, d_k)。这一步的维度顺序很容易搞错,建议画个图辅助理解。
4.2 模型参数量怎么估算
Transformer 的参数量主要来自以下几个部分:
| 组件 | 参数量公式 | 说明 |
|---|---|---|
| 词嵌入 | vocab_size * d_model | 词表大小乘以模型维度 |
| 位置编码 | max_len * d_model | 如果可学习的话 |
| Multi-Head Attention | 4 * d_model * d_model | Q、K、V、输出投影各一个矩阵 |
| Feed Forward | 2 * d_model * d_ff | 两个全连接层 |
| Layer Norm | 2 * d_model | 缩放和平移参数 |
以一个 d_model=512、d_ff=2048、num_layers=6、vocab_size=30000 的模型为例:
- 词嵌入:30000 * 512 = 15,360,000
- 每层 Attention:4 * 512 * 512 = 1,048,576
- 每层 FFN:2 * 512 * 2048 = 2,097,152
- 每层 LN:2 * 512 * 2 = 2,048
- 每层总计:约 3,147,776
- 6 层 Encoder:约 18,886,656
- 6 层 Decoder:约 18,886,656(加上 Cross-Attention 的 1,048,576 * 6)
- 总计:约 55M 参数
这个估算方法在选模型规模的时候很有用。如果你只有 8GB 显存,大概能训 100M 参数左右的模型,再大就要考虑梯度累积或者模型并行了。
4.3 训练时的学习率调度策略
Transformer 的训练对学习率非常敏感。原始论文用了一个 warmup 策略:学习率先线性增加,然后再按步数的平方根倒数衰减。
lr = d_model^(-0.5) * min(step^(-0.5), step * warmup_steps^(-1.5))warmup_steps 一般设为 4000 或者总步数的 10%。这个策略的目的是在训练初期用较小的学习率让模型稳定下来,然后再逐步增大学习率加速收敛,最后再衰减以保证收敛到好的局部最优。
我实测下来,如果不加 warmup,直接用固定学习率,模型很容易在训练初期就发散。加了 warmup 之后,训练稳定很多。另外,Adam 优化器的 beta2 建议设为 0.98 而不是默认的 0.999,这样对 Transformer 更友好。
5. 常见问题排查与避坑经验
5.1 训练 loss 不下降或者震荡怎么办
这是最常见的问题,可能的原因和排查思路如下:
| 现象 | 可能原因 | 排查方法 |
|---|---|---|
| loss 完全不降 | 学习率太小 | 增大学习率,检查 warmup 是否生效 |
| loss 震荡 | 学习率太大 | 减小学习率,增大 batch size |
| loss 先降后升 | 过拟合 | 加 dropout,减小模型规模 |
| loss 变成 NaN | 梯度爆炸 | 加梯度裁剪,检查 LN 位置 |
| loss 降得很慢 | 初始化不好 | 检查参数初始化,尝试 Pre-LN |
我遇到最多的是梯度爆炸导致的 NaN。解决方法是在反向传播后、优化器更新前加梯度裁剪:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)max_norm 一般设 1.0 或者 5.0。这个操作几乎不会影响正常训练,但能有效防止梯度爆炸。
5.2 注意力权重全是均匀分布怎么办
如果注意力权重接近均匀分布,说明模型没有学到有意义的注意力模式。可能的原因有:
- 学习率太大,模型还没稳定下来
- 训练数据太少,模型没有足够的信息来学习
- 位置编码有问题,模型无法区分不同位置
- 初始化不好,Q 和 K 的点积太小
排查的时候,可以先可视化注意力权重,看看是不是真的均匀。如果确实是均匀的,先检查位置编码是否正确加到了输入上。然后检查 Q、K 的初始化,确保它们的方差在合理范围内。
5.3 显存不够用的优化技巧
Transformer 的显存占用主要来自注意力矩阵,它的形状是 (batch_size, num_heads, seq_len, seq_len)。当 seq_len 很大时,这个矩阵会非常占显存。
优化方法有几种:
- 减小 batch size,这是最直接的方法
- 用梯度累积,模拟大 batch 的效果
- 用混合精度训练,把 float32 换成 float16,显存占用减半
- 用 Flash Attention 或者 Memory-Efficient Attention,这些实现会优化注意力矩阵的存储和计算
我实测下来,混合精度训练是最划算的,几乎不损失精度,显存直接减半。Flash Attention 效果也很好,但需要特定的 GPU 架构支持。
5.4 Decoder 生成时重复输出同一个词怎么办
这是生成任务里的常见问题,通常是因为模型陷入了局部最优。解决方法有:
- 用 beam search 代替 greedy decoding
- 加 repetition penalty,对已经生成过的词降低概率
- 加 temperature 参数,让分布更平滑
- 加 top-k 或 top-p 采样,避免总是选概率最高的词
我一般会先用 beam search 试试,如果还是重复,再加 repetition penalty。temperature 和 top-p 要根据具体任务调,没有万能的值。
6. Transformer 在不同场景下的变体与扩展
6.1 Vision Transformer 是怎么把图像变成序列的
Vision Transformer 的思路很直接:把图像切成固定大小的 patch,每个 patch 展平成一个向量,再加上位置编码,就变成了一个序列,然后直接扔给标准的 Transformer Encoder。
比如一张 224x224 的图片,切成 16x16 的 patch,就得到 196 个 patch。每个 patch 展平后是 16163=768 维,正好和 d_model 一致。然后加上一个可学习的分类 token,放在序列最前面,最后用这个 token 的输出做分类。
这个设计的美妙之处在于:它几乎不需要修改 Transformer 的结构,就能直接用在视觉任务上。但缺点是计算量比 CNN 大很多,因为注意力是平方复杂度的。
6.2 Transformer 做时序预测的关键调整
用 Transformer 做时序预测,和做 NLP 有几个关键区别:
- 位置编码要改成适合时序的形式,比如可学习的位置编码或者时间戳编码
- 输出层通常是一个回归头,而不是分类头
- 损失函数用 MSE 或者 MAE,而不是交叉熵
- 可能需要处理多变量输入,也就是每个时间步有多个特征
我做过一个正弦函数预测的实验,用 Transformer 预测未来 10 个时间步的值。关键是要把输入序列和目标序列错开,输入是 t 到 t+n,目标是 t+1 到 t+n+1。训练的时候用 teacher forcing,推理的时候用自回归生成。
实测下来,Transformer 在时序预测上表现不错,尤其是当序列有长距离依赖的时候。但如果序列很短,或者主要是局部模式,CNN 或者 RNN 可能更合适。
6.3 新手跑 Transformer 模型的建议路线
如果你是第一次跑 Transformer,我建议按这个路线来:
- 先用 HuggingFace 的 transformers 库跑一个预训练模型,感受一下输入输出。
- 然后找一个简单的任务,比如文本分类,微调一下模型。
- 接着试着手写一个最小的 Transformer,比如 2 层 Encoder,做个小规模的翻译或者复制任务。
- 最后再尝试从头训练一个完整的模型,处理真实数据。
这个路线的好处是循序渐进,每一步都有正反馈,不会一上来就被复杂的细节劝退。我当初就是直接从零手写,结果卡在维度变换上好几天,差点放弃。后来退回去先用现成的库跑通,再回头手写,就顺畅多了。
6.4 Transformer 和 CNN、RNN 的核心区别
| 特性 | Transformer | CNN | RNN |
|---|---|---|---|
| 并行能力 | 完全并行 | 完全并行 | 无法并行 |
| 长距离依赖 | 直接建模 | 需要堆叠多层 | 容易梯度消失 |
| 计算复杂度 | O(n^2) | O(n) | O(n) |
| 位置感知 | 需要位置编码 | 天然有 | 天然有 |
| 参数量 | 较大 | 较小 | 中等 |
Transformer 最大的优势是并行能力和长距离依赖建模。但代价是计算复杂度是平方级的,序列很长时计算量会爆炸。CNN 的复杂度是线性的,但感受野有限,需要堆叠很多层才能覆盖长距离。RNN 天然适合序列,但无法并行,训练速度慢。
实际选型的时候,如果序列长度在几百以内,Transformer 通常是首选。如果序列很长,比如几千甚至几万,就要考虑用稀疏注意力或者线性注意力的变体。如果计算资源有限,CNN 或者 RNN 可能更实际。
7. 我个人的实操心得与建议
7.1 调试 Transformer 的几个实用技巧
第一个技巧是打印中间张量的形状。Transformer 的维度变换很多,很容易搞错。我习惯在每个关键步骤后打印形状,确保和预期一致。比如 Multi-Head Attention 里,Q、K、V 的形状变换就有好几步,每一步都打印出来,出问题的时候一眼就能定位。
第二个技巧是用小规模数据先跑通。不要一上来就用全量数据训练,先用几百条数据跑几个 epoch,确保模型能过拟合。如果能过拟合,说明模型结构没问题,再上全量数据。如果不能过拟合,说明结构或者训练逻辑有问题,先修好再扩大规模。
第三个技巧是可视化注意力权重。把注意力权重画成热力图,能直观地看到模型在关注哪些位置。如果注意力权重看起来有规律,比如对角线附近权重高,说明模型学到了位置关系。如果看起来杂乱无章,可能模型还没训练好。
7.2 关于学习路线的一点建议
Transformer 涉及的知识点很多,不要试图一次全部搞懂。我的建议是先抓住主干:Attention 机制、Encoder-Decoder 结构、位置编码。这三个搞懂了,其他的细节可以慢慢补。
看论文的时候,第一遍不要纠结公式推导,先看图和文字描述,理解整体流程。第二遍再仔细看公式,自己推导一遍。第三遍看代码实现,对照论文理解每一行代码在做什么。
手写代码是必须的,但不要一开始就追求完美。先写一个能跑通的版本,哪怕效率低一点、代码丑一点都没关系。跑通之后再优化,比如加上 Multi-Head、加上 Mask、加上位置编码。
7.3 后续可以深入的方向
如果你已经把标准 Transformer 搞懂了,可以往这几个方向深入:
- 高效注意力机制:Linformer、Performer、Flash Attention,解决平方复杂度问题
- 长序列建模:Longformer、BigBird,处理超长序列
- 多模态 Transformer:CLIP、DALL-E,同时处理文本和图像
- 稀疏化与剪枝:减少参数量,提升推理速度
- 位置编码的变体:旋转位置编码、相对位置编码,提升长度泛化能力
每个方向都有大量的论文和开源实现,选一个你感兴趣的方向深入进去,比泛泛地看要有效得多。
我在实际项目里用得最多的还是标准的 Transformer Encoder,配合预训练模型做微调。这套组合在大多数任务上都能拿到不错的效果,而且有大量的开源工具支持,踩坑的成本比较低。如果你刚开始接触,建议也从这个路线入手,等熟悉了再尝试更复杂的变体。