ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

Transformer论文精读与PyTorch从零复现:注意力机制完全指南

2026/9/20 7:38:07 拓冰建站 浏览量
Transformer论文精读与PyTorch从零复现:注意力机制完全指南 1. 为什么值得花时间啃这篇论文如果你正在看大模型相关的东西不管是做应用、做微调还是单纯想搞明白ChatGPT这类产品底下到底在跑什么那《Attention Is All You Need》这篇论文是绕不过去的。2017年谷歌团队发在NeurIPS上的这篇文章直接干掉了RNN和CNN在序列建模里的统治地位提出了完全基于注意力机制的Transformer架构。今天你听到的BERT、GPT、LLaMA、Qwen、DeepSeek底层骨架全是它。但问题在于这篇论文满打满算11页公式密集、符号多很多人第一遍读下来感觉“每个字都认识连起来不知道在说什么”。更麻烦的是网上解读文章虽然多但要么只翻译不讲解要么只讲概念不给代码要么代码复现和论文对不上号。我前前后后把这篇论文读了不下十遍用PyTorch手撸过三版Transformer踩过的坑包括但不限于位置编码写错导致模型完全不收敛、mask矩阵维度搞反、多头注意力拆头之后忘记reshape回来、学习率warmup策略没对齐论文导致训练震荡。这篇文章就是把我这些年读论文、复现代码的经验一次性倒出来。我会先带你把论文全文过一遍逐段精读关键部分然后从零用PyTorch写一个完整的Transformer最后跑一个简单的序列到序列任务验证代码正确性。不管你是刚入门深度学习的小白还是已经用过大模型但没深究原理的工程师跟着走一遍Transformer这个东西在你脑子里就不再是黑盒了。提示读这篇论文之前建议先搞清楚什么是词向量、什么是softmax、什么是残差连接。如果这些概念还模糊先去补一下基础不然读起来会很痛苦。2. 论文全文翻译与逐段精读2.1 摘要与引言作者到底想解决什么问题论文摘要原文翻译过来是这样的主流的序列转换模型基于复杂的循环或卷积神经网络包含编码器和解码器。性能最好的模型还通过注意力机制连接编码器和解码器。我们提出一种新的简单网络架构——Transformer完全基于注意力机制彻底摒弃了循环和卷积。在两项机器翻译任务上的实验表明这些模型在质量上更优同时更可并行化训练时间显著减少。我们的模型在WMT 2014英德翻译任务上达到28.4 BLEU超过现有最佳结果包括集成模型2 BLEU以上。在WMT 2014英法翻译任务上我们的模型在8个GPU上训练3.5天后达到41.8 BLEU创下新的单模型最佳成绩训练成本仅为文献中最佳模型的一小部分。这段摘要信息量很大。作者明确说了三个核心卖点第一不用RNN不用CNN纯注意力第二训练快因为可以并行第三效果好BLEU分数高。你要知道在2017年序列建模的主流是LSTM、GRU这些循环网络它们有个致命问题必须按时间步顺序计算第t步依赖第t-1步的隐状态根本没法并行。这导致两个后果训练慢而且长序列容易梯度消失。Transformer直接把循环结构扔了所有位置同时计算这是它最根本的突破。引言部分作者先回顾了RNN和CNN在序列建模中的使用。RNN系列模型包括LSTM、GRU是当时的标准方案但它们的顺序计算特性限制了并行化在长序列上内存占用也大。CNN类模型比如ByteNet、ConvS2S虽然可以并行但要让两个远距离位置的信息交互需要的卷积层数随距离线性或对数增长。然后作者抛出核心观点注意力机制可以直接建模任意两个位置之间的依赖关系不管距离多远操作数都是常数。这就是Transformer的理论基础。我读这一段的时候最大的感受是作者不是在做一个增量改进而是在换赛道。他们看准了RNN的顺序计算是瓶颈直接把这个瓶颈拿掉用注意力替代。这种“重新定义问题”的思路比单纯调参刷榜有价值得多。2.2 模型架构总览编码器-解码器结构论文第3节给出了Transformer的整体架构。和之前的序列转换模型一样Transformer也是编码器-解码器结构。编码器把输入序列x1, ..., xn映射成连续表示序列z1, ..., zn。解码器拿到z之后自回归地生成输出序列y1, ..., ym每生成一个词就把它加到输入里再预测下一个。编码器由N6个相同的层堆叠而成。每一层有两个子层第一个是多头自注意力机制第二个是简单的位置全连接前馈网络。每个子层都加了残差连接和层归一化。具体来说每个子层的输出是LayerNorm(x Sublayer(x))。为了支持残差连接所有子层和嵌入层的输出维度都设为d_model512。解码器同样由N6个相同层堆叠。但每层有三个子层第一个是带掩码的多头自注意力第二个是对编码器输出的多头交叉注意力第三个是前馈网络。和编码器一样每个子层都有残差连接和层归一化。解码器的自注意力加了掩码确保预测位置i时只能看到小于i的位置保持自回归性质。这里有几个关键设计决策值得展开说。为什么是6层论文做了消融实验6层是效果和计算成本的平衡点。为什么d_model512这是base模型的配置big模型用的是1024。为什么用LayerNorm而不是BatchNorm因为序列长度可变BatchNorm在变长序列上统计量不稳定LayerNorm对每个样本独立归一化更适合NLP任务。残差连接的作用是缓解深层网络的梯度消失问题让梯度能直接流过跳跃连接。层归一化则稳定训练过程加速收敛。这两个技术组合在今天已经是标配了但在2017年把它们用在纯注意力架构上是需要勇气的。2.3 注意力机制Scaled Dot-Product Attention这是整篇论文最核心的部分。注意力函数的本质可以描述为把一个查询query和一组键值对key-value pair映射成输出。输出是值的加权和权重由查询和对应键的相似度决定。论文里用的注意力叫“缩放点积注意力”。输入包括维度为dk的查询和键维度为dv的值。计算步骤是查询和所有键做点积除以根号dk然后过softmax得到权重最后用权重对值加权求和。公式写出来就是Attention(Q, K, V) softmax(QK^T / sqrt(dk)) V为什么要除以根号dk论文给的解释是当dk很大时点积结果会变得很大导致softmax进入梯度极小的区域训练困难。除以根号dk可以把点积结果拉回合理范围保持梯度稳定。我实测过如果不做这个缩放在dk64时训练loss会剧烈震荡收敛很慢。用矩阵形式理解更直观Q是(n, dk)矩阵K是(m, dk)矩阵QK^T得到(n, m)的分数矩阵每一行表示一个查询对所有键的注意力分数。softmax按行归一化得到注意力权重。最后乘以V(m, dv)得到(n, dv)的输出。整个过程就是几次矩阵乘法GPU上跑起来非常快。2.4 多头注意力为什么需要多个头论文没有满足于单一注意力而是提出了多头注意力。核心思想是把查询、键、值分别用不同的线性变换投影到低维空间做h次注意力计算然后把结果拼接起来再投影一次。公式是MultiHead(Q, K, V) Concat(head_1, ..., head_h) W^O 其中 head_i Attention(Q W_i^Q, K W_i^K, V W_i^V)论文里h8dkdvd_model/h64。因为每个头的维度降低了总计算量和单头全维度注意力差不多。为什么要多头我的理解是不同的头可以关注不同的模式。比如在翻译任务里一个头可能关注语法依赖另一个头关注语义相似度还有一个头关注位置邻近关系。单头注意力只能学到一种加权方式多头让模型有多个“视角”。这就像你看一个东西从不同角度拍几张照片比只拍一张信息更丰富。实际代码里多头注意力的实现有两种方式。一种是老老实实循环h次每次算一个头另一种是用一个大矩阵一次性算完再reshape。后者效率高得多我后面代码复现部分会详细讲。2.5 位置编码没有循环怎么知道顺序Transformer完全抛弃了循环和卷积这意味着它本身对序列顺序没有感知。你把输入序列打乱注意力算出来的结果是一样的因为注意力是对称的。但语言是有顺序的“猫追老鼠”和“老鼠追猫”意思完全相反。所以必须把位置信息注入进去。论文的解决方案是位置编码。在编码器和解码器的输入嵌入上加上一个位置编码向量。位置编码的维度也是d_model这样可以直接相加。论文用了正弦和余弦函数PE(pos, 2i) sin(pos / 10000^(2i/d_model)) PE(pos, 2i1) cos(pos / 10000^(2i/d_model))其中pos是位置i是维度索引。偶数维度用sin奇数维度用cos。为什么选这个函数论文解释说对于任意固定偏移kPE(posk)可以表示为PE(pos)的线性函数。这让模型容易学到相对位置关系。而且这个编码可以扩展到比训练时更长的序列因为它是解析式而不是学出来的。我实际用下来正弦位置编码在大多数任务上够用但后来很多模型改用了可学习的位置嵌入比如BERT或旋转位置编码RoPE比如LLaMA。不过作为理解Transformer的起点正弦编码是最经典的。2.6 前馈网络与残差连接每个编码器和解码器层里都有一个前馈网络对每个位置独立应用。它包含两个线性变换中间夹一个ReLU激活FFN(x) max(0, xW1 b1) W2 b2输入输出维度是d_model512中间层维度d_ff2048。这个设计先升维再降维增加了模型的非线性表达能力。虽然叫“前馈网络”但它对每个位置是独立处理的不涉及位置间的交互位置间的交互全靠注意力层完成。残差连接和层归一化在每个子层后面。具体顺序论文写的是LayerNorm(x Sublayer(x))也就是先做子层计算加上残差再归一化。后来有些实现改成Pre-LN先归一化再进子层训练更稳定但论文原版是Post-LN。我建议复现时先按论文来跑通了再尝试变体。2.7 训练策略与正则化论文第5节讲了训练细节。优化器用Adambeta10.9beta20.98epsilon1e-9。学习率不是固定的而是用了一个warmup策略lrate d_model^(-0.5) * min(step_num^(-0.5), step_num * warmup_steps^(-1.5))warmup_steps设为4000。意思是前4000步学习率线性增加之后按步数的平方根倒数衰减。这个策略对Transformer训练很关键我试过不用warmup直接上大学习率模型直接发散。正则化用了两种残差dropout和标签平滑。Dropout加在子层输出上、嵌入和位置编码相加之后rate0.1。标签平滑值epsilon0.1让模型不要过度自信提升泛化。另外还有三个细节一是用了字节对编码BPE共享源语言和目标语言的词表二是beam searchbeam size4长度惩罚alpha0.6三是训练用了8个GPUbase模型每步约0.4秒。3. 从零用PyTorch复现Transformer3.1 环境准备与项目结构先把环境搭好。我用的配置是Python 3.10 PyTorch 2.1 CUDA 11.8。如果你没有GPUCPU也能跑就是慢一点。安装命令pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 pip install numpy matplotlib项目结构我习惯这样组织transformer_from_scratch/ ├── model.py # Transformer模型定义 ├── data.py # 数据加载和预处理 ├── train.py # 训练脚本 ├── config.py # 超参数配置 └── checkpoints/ # 保存模型权重这样拆分的理由是模型定义和数据逻辑分离方便单独调试。比如你想换数据集只改data.py就行不用动模型代码。3.2 缩放点积注意力的代码实现先写最核心的注意力函数。我直接上代码然后逐行解释import torch import torch.nn as nn import torch.nn.functional as F import math def scaled_dot_product_attention(query, key, value, maskNone): query: (batch, heads, seq_len, d_k) key: (batch, heads, seq_len, d_k) value: (batch, heads, seq_len, d_v) mask: (batch, 1, 1, seq_len) 或 (batch, 1, seq_len, seq_len) d_k query.size(-1) scores torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attention_weights F.softmax(scores, dim-1) output torch.matmul(attention_weights, value) return output, attention_weights这段代码有几个关键点。第一key.transpose(-2, -1)是把最后两个维度交换让key变成(d_k, seq_len)这样和query做矩阵乘法得到(seq_len, seq_len)的分数矩阵。第二除以math.sqrt(d_k)是论文里的缩放操作d_k是每个头的维度。第三mask的处理把mask为0的位置填成负无穷softmax之后这些位置的权重就变成0了。第四softmax的dim-1表示对最后一维也就是key的维度做归一化。注意mask填负无穷用float(-inf)不要用很大的负数比如-1e9因为在某些精度下softmax可能会溢出。PyTorch的masked_fill配合float(-inf)是安全的。3.3 多头注意力的高效实现多头注意力如果用循环写代码直观但慢。我用矩阵 reshape 的方式实现一次算完class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout0.1): super().__init__() assert d_model % num_heads 0, d_model必须能被num_heads整除 self.d_model d_model self.num_heads num_heads self.d_k d_model // num_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, query, key, value, maskNone): batch_size query.size(0) # 线性投影并拆分成多头 # (batch, seq_len, d_model) - (batch, seq_len, num_heads, d_k) - (batch, num_heads, seq_len, d_k) Q self.W_q(query).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) K self.W_k(key).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) V self.W_v(value).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) # 缩放点积注意力 attn_output, attn_weights scaled_dot_product_attention(Q, K, V, mask) # 拼接多头 # (batch, num_heads, seq_len, d_k) - (batch, seq_len, num_heads, d_k) - (batch, seq_len, d_model) attn_output attn_output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model) # 最终线性变换 output self.W_o(attn_output) return output, attn_weights这里最容易出错的是reshape和transpose的顺序。view(batch_size, -1, num_heads, d_k)把d_model维度拆成(num_heads, d_k)然后transpose(1, 2)把num_heads换到前面。算完注意力之后先transpose(1, 2)换回来再用contiguous().view()合并num_heads和d_k。如果你忘了contiguous()view会报错因为transpose之后内存不连续。我踩过的坑一开始把transpose(1, 2)写成了transpose(0, 1)结果batch维度被换掉了训练完全不收敛。调试了半天才发现。所以维度操作一定要写注释标清楚每一步的shape。3.4 位置编码的实现细节位置编码的代码看起来简单但有几个细节要注意class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000, dropout0.1): super().__init__() self.dropout nn.Dropout(dropout) # 创建位置编码矩阵 pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) 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) pe pe.unsqueeze(0) # (1, max_len, d_model) # 注册为buffer不参与梯度更新 self.register_buffer(pe, pe) def forward(self, x): # x: (batch, seq_len, d_model) x x self.pe[:, :x.size(1), :] return self.dropout(x)div_term的计算是10000^(-2i/d_model)用exp和log实现是为了数值稳定。pe[:, 0::2]表示所有行的偶数列pe[:, 1::2]是奇数列。register_buffer让pe成为模型的一部分但不更新梯度保存模型时也会一起存下来。提示max_len5000对大多数任务够用。如果你的序列更长调大这个值。但注意pe矩阵会占用内存max_len * d_model * 4字节50005124约10MB可以接受。3.5 编码器层与解码器层的组装有了多头注意力和位置编码就可以组装编码器和解码器层了class EncoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads, dropout) self.feed_forward nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Dropout(dropout), nn.Linear(d_ff, d_model) ) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) def forward(self, x, maskNone): # 自注意力子层 attn_output, _ self.self_attn(x, x, x, mask) x self.norm1(x self.dropout1(attn_output)) # 前馈子层 ff_output self.feed_forward(x) x self.norm2(x self.dropout2(ff_output)) return x解码器层多了一个交叉注意力子层class DecoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads, dropout) self.cross_attn MultiHeadAttention(d_model, num_heads, dropout) self.feed_forward nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Dropout(dropout), nn.Linear(d_ff, d_model) ) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.norm3 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) self.dropout3 nn.Dropout(dropout) def forward(self, x, encoder_output, src_maskNone, tgt_maskNone): # 掩码自注意力 attn_output, _ self.self_attn(x, x, x, tgt_mask) x self.norm1(x self.dropout1(attn_output)) # 交叉注意力query来自解码器key和value来自编码器 attn_output, _ self.cross_attn(x, encoder_output, encoder_output, src_mask) x self.norm2(x self.dropout2(attn_output)) # 前馈 ff_output self.feed_forward(x) x self.norm3(x self.dropout3(ff_output)) return x交叉注意力的query是解码器的中间表示key和value是编码器的输出。这样解码器在生成每个词时可以“看”到完整的输入序列。3.6 完整Transformer模型搭建把编码器和解码器堆叠起来加上嵌入层和输出层class Transformer(nn.Module): def __init__(self, src_vocab_size, tgt_vocab_size, d_model512, num_heads8, num_layers6, d_ff2048, max_len5000, dropout0.1): super().__init__() self.encoder_embedding nn.Embedding(src_vocab_size, d_model) self.decoder_embedding nn.Embedding(tgt_vocab_size, d_model) self.positional_encoding PositionalEncoding(d_model, max_len, dropout) self.encoder_layers nn.ModuleList([ EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.decoder_layers nn.ModuleList([ DecoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.fc_out nn.Linear(d_model, tgt_vocab_size) self.d_model d_model def forward(self, src, tgt, src_maskNone, tgt_maskNone): # 嵌入 位置编码 src_emb self.positional_encoding(self.encoder_embedding(src) * math.sqrt(self.d_model)) tgt_emb self.positional_encoding(self.decoder_embedding(tgt) * math.sqrt(self.d_model)) # 编码器 encoder_output src_emb for layer in self.encoder_layers: encoder_output layer(encoder_output, src_mask) # 解码器 decoder_output tgt_emb for layer in self.decoder_layers: decoder_output layer(decoder_output, encoder_output, src_mask, tgt_mask) # 输出投影 output self.fc_out(decoder_output) return output嵌入层输出乘以math.sqrt(d_model)是论文里的操作目的是让嵌入和位置编码的尺度匹配。因为位置编码的值在[-1, 1]之间而嵌入层初始化后值比较小乘以根号d_model可以放大嵌入的尺度。3.7 mask矩阵的生成与使用mask是Transformer里最容易搞错的部分。需要两种maskpadding mask和causal mask。def create_padding_mask(seq, pad_token0): # seq: (batch, seq_len) mask (seq ! pad_token).unsqueeze(1).unsqueeze(2) # (batch, 1, 1, seq_len) return mask def create_causal_mask(seq_len): # 下三角矩阵对角线及以上为1以下为0 mask torch.tril(torch.ones(seq_len, seq_len)).bool() # (1, 1, seq_len, seq_len) return mask.unsqueeze(0).unsqueeze(0)padding mask把填充位置屏蔽掉causal mask确保解码器只能看到当前位置及之前的位置。实际使用时解码器的自注意力需要同时应用两种mask通常是把它们做逻辑与操作。我踩过的坑causal mask的方向搞反了。torch.tril生成下三角矩阵位置(i, j)为1当且仅当j i表示位置i可以看到位置j。如果你用torch.triu就反了模型会看到未来的词训练loss会异常低但推理时完全不能用。4. 训练验证与常见问题排查4.1 用简单任务验证代码正确性代码写完了怎么知道对不对我建议先跑一个极简的序列复制任务输入一个序列让模型输出同样的序列。这个任务足够简单如果模型学不会说明代码有bug。# 生成数据 def generate_copy_data(batch_size, seq_len, vocab_size): src torch.randint(1, vocab_size, (batch_size, seq_len)) tgt src.clone() return src, tgt # 训练循环 model Transformer(src_vocab_size100, tgt_vocab_size100, d_model128, num_heads4, num_layers2, d_ff512) optimizer torch.optim.Adam(model.parameters(), lr0.0001, betas(0.9, 0.98), eps1e-9) criterion nn.CrossEntropyLoss(ignore_index0) for step in range(2000): src, tgt generate_copy_data(32, 10, 100) tgt_input tgt[:, :-1] tgt_output tgt[:, 1:] src_mask create_padding_mask(src) tgt_mask create_causal_mask(tgt_input.size(1)) output model(src, tgt_input, src_mask, tgt_mask) loss criterion(output.reshape(-1, 100), tgt_output.reshape(-1)) optimizer.zero_grad() loss.backward() optimizer.step() if step % 200 0: print(fStep {step}, Loss: {loss.item():.4f})如果代码正确loss应该从4.6左右ln(100)快速下降到0.1以下。我实测在2000步内能降到0.01以下说明模型完全学会了复制。4.2 常见问题速查表问题现象可能原因排查方法解决方案loss不下降学习率太大或太小打印梯度范数调整学习率加warmuploss变成nan除零或log(0)检查mask和softmax加epsilon检查mask训练loss低但推理差过拟合或mask错误检查causal mask方向修正mask加dropout显存溢出batch太大或序列太长打印各层shape减小batch梯度累积收敛极慢位置编码错误可视化位置编码检查sin/cos公式注意力权重全一样缩放因子缺失检查是否除以sqrt(d_k)补上缩放4.3 实操心得与避坑技巧第一个心得先在小数据上过拟合。拿到一个新模型不要一上来就跑完整数据集。先拿100条数据把模型调到能过拟合训练loss接近0这说明模型容量和代码都没问题。然后再上全量数据调正则化。第二个心得梯度裁剪是保命符。Transformer训练时梯度偶尔会爆炸加一行torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)能救命。我遇到过好几次loss突然飙到几千加了梯度裁剪之后就稳了。第三个心得学习率warmup不能省。论文的warmup公式我直接抄def get_lr(step, d_model, warmup_steps4000): return d_model ** (-0.5) * min(step ** (-0.5), step * warmup_steps ** (-1.5))前4000步线性增加之后衰减。我试过跳过warmup直接上1e-4的学习率前100步loss就炸了。第四个心得保存注意力权重可视化。调试的时候把注意力矩阵画出来看看模型到底关注哪里。正常的注意力应该是对角线附近有较高权重关注邻近词同时某些头会关注远距离依赖。如果所有头都长得一样说明多头没起作用检查reshape逻辑。第五个心得标签平滑提升泛化。论文用了0.1的标签平滑我在小数据集上试过不用标签平滑验证loss会高0.2左右。PyTorch里用nn.CrossEntropyLoss(label_smoothing0.1)就行。4.4 从复现到深入下一步可以做什么代码跑通之后你可以做几件事来加深理解。第一把base模型换成big模型d_model1024, num_heads16, d_ff4096, dropout0.3看看参数量和显存占用怎么变。第二把正弦位置编码换成可学习的位置嵌入对比效果。第三把Post-LN改成Pre-LN观察训练稳定性。第四用训练好的模型做注意力可视化分析不同头学到了什么模式。我个人在实际操作中的体会是读论文和复现代码是两件事。论文给你的是设计思路代码让你看到每个细节怎么落地。很多时候论文里一句话带过的东西代码里要写十几行而且处处是坑。但正是这些坑让你真正理解模型为什么这么设计。比如那个除以根号dk的缩放论文只说了“防止softmax梯度消失”你自己写代码时忘了加看到loss震荡才会真正记住这个操作的必要性。最后分享一个小技巧如果你不想从零写PyTorch官方有nn.Transformer模块可以直接用。但我强烈建议至少手写一遍因为官方模块封装了很多细节你调参的时候不知道底层发生了什么。手写一遍之后再用官方模块你会知道每个参数对应什么出问题也知道去哪里找。