首页 / 资讯中心 / 文章详情

Transformer架构深度拆解:从Self-Attention到Encoder-Decoder的实战指南

Transformer架构深度拆解:从Self-Attention到Encoder-Decoder的实战指南 ★ FEATURED ARTICLE
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 层 EncoderBERT-large 用了 24 层。GPT 系列也是类似层数从 12 到 96 不等。层数越多模型的表达能力越强但计算量和显存占用也会线性增长。我个人的经验是如果你在做小规模的任务比如文本分类或者简单的时序预测4 到 6 层通常就够了。如果数据量很大、任务很复杂再考虑加到 12 层甚至更多。但要注意层数增加带来的收益是递减的而且训练难度会显著上升。1.3 残差连接和层归一化的位置选择Transformer 每个子层Self-Attention、Feed Forward外面都包了一层残差连接和层归一化。原始论文用的是 Post-LN也就是先做残差加法再做 Layer Normalizationoutput 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 个头的输出拼接起来再经过一个线性变换。原始论文里 h8每个头的维度是 d_model/h64。这样总的计算量和单个 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而如果直接填 0Softmax 之后仍然会有非零权重那就起不到 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_ff2048d_model512扩展倍数是 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, 2i1) cos(pos / 10000^(2i/d_model))这个设计的好处是对于任意固定的偏移量 kPE(posk) 可以表示为 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 Attention4 * d_model * d_modelQ、K、V、输出投影各一个矩阵Feed Forward2 * d_model * d_ff两个全连接层Layer Norm2 * d_model缩放和平移参数以一个 d_model512、d_ff2048、num_layers6、vocab_size30000 的模型为例词嵌入30000 * 512 15,360,000每层 Attention4 * 512 * 512 1,048,576每层 FFN2 * 512 * 2048 2,097,152每层 LN2 * 512 * 2 2,048每层总计约 3,147,7766 层 Encoder约 18,886,6566 层 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 sizeloss 先降后升过拟合加 dropout减小模型规模loss 变成 NaN梯度爆炸加梯度裁剪检查 LN 位置loss 降得很慢初始化不好检查参数初始化尝试 Pre-LN我遇到最多的是梯度爆炸导致的 NaN。解决方法是在反向传播后、优化器更新前加梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.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 展平后是 16163768 维正好和 d_model 一致。然后加上一个可学习的分类 token放在序列最前面最后用这个 token 的输出做分类。这个设计的美妙之处在于它几乎不需要修改 Transformer 的结构就能直接用在视觉任务上。但缺点是计算量比 CNN 大很多因为注意力是平方复杂度的。6.2 Transformer 做时序预测的关键调整用 Transformer 做时序预测和做 NLP 有几个关键区别位置编码要改成适合时序的形式比如可学习的位置编码或者时间戳编码输出层通常是一个回归头而不是分类头损失函数用 MSE 或者 MAE而不是交叉熵可能需要处理多变量输入也就是每个时间步有多个特征我做过一个正弦函数预测的实验用 Transformer 预测未来 10 个时间步的值。关键是要把输入序列和目标序列错开输入是 t 到 tn目标是 t1 到 tn1。训练的时候用 teacher forcing推理的时候用自回归生成。实测下来Transformer 在时序预测上表现不错尤其是当序列有长距离依赖的时候。但如果序列很短或者主要是局部模式CNN 或者 RNN 可能更合适。6.3 新手跑 Transformer 模型的建议路线如果你是第一次跑 Transformer我建议按这个路线来先用 HuggingFace 的 transformers 库跑一个预训练模型感受一下输入输出。然后找一个简单的任务比如文本分类微调一下模型。接着试着手写一个最小的 Transformer比如 2 层 Encoder做个小规模的翻译或者复制任务。最后再尝试从头训练一个完整的模型处理真实数据。这个路线的好处是循序渐进每一步都有正反馈不会一上来就被复杂的细节劝退。我当初就是直接从零手写结果卡在维度变换上好几天差点放弃。后来退回去先用现成的库跑通再回头手写就顺畅多了。6.4 Transformer 和 CNN、RNN 的核心区别特性TransformerCNNRNN并行能力完全并行完全并行无法并行长距离依赖直接建模需要堆叠多层容易梯度消失计算复杂度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处理超长序列多模态 TransformerCLIP、DALL-E同时处理文本和图像稀疏化与剪枝减少参数量提升推理速度位置编码的变体旋转位置编码、相对位置编码提升长度泛化能力每个方向都有大量的论文和开源实现选一个你感兴趣的方向深入进去比泛泛地看要有效得多。我在实际项目里用得最多的还是标准的 Transformer Encoder配合预训练模型做微调。这套组合在大多数任务上都能拿到不错的效果而且有大量的开源工具支持踩坑的成本比较低。如果你刚开始接触建议也从这个路线入手等熟悉了再尝试更复杂的变体。
阅读完成 · 觉得有帮助?
咨询建站