ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Transformer 论文精读:自注意力、多头注意力与实现

Transformer 论文精读:自注意力、多头注意力与实现 1. 从标题出发这篇论文到底在讲什么《Attention Is All You Need》这个标题起得非常嚣张翻译过来就是“你需要的只是注意力”。第一次看到的人往往有两种反应一种觉得这是标题党另一种看完摘要之后直接愣住——它真的把循环和卷积全都扔掉了。这篇 2017 年的论文提出的 Transformer 架构后来成了几乎所有主流大模型的底座从机器翻译一路蔓延到语言建模、视觉识别、语音处理、蛋白质结构预测。可以说今天的 AI 从业者不管做哪个方向读这篇论文都是绕不过去的一课。我在不同阶段读过它好几次。第一次读是为了复现囫囵吞枣地把代码抄了一遍结果训练不收敛第二次读是为了给别人讲被迫逐段翻译才发现自己上次漏掉的全是关键细节——比如缩放因子为什么是 $\sqrt{d_k}$、位置编码为什么用正弦余弦、warmup 到底在救什么。所以这篇解读我打算换个方式写先讲论文为什么这么设计再讲每个模块的实现要点最后把训练和调参的坑摊开来说。适合的读者范围很宽只要你懂一点线性代数、看过几张神经网络的结构图就能跟下来如果你已经能写 PyTorch那第四、第六节大概是你最想跳过去看的部分。1.1 论文出现之前的困局在 Transformer 之前序列建模基本是循环神经网络的地盘。RNN、LSTM、GRU 这一脉的思路是“一个词一个词地走过去”第 $t$ 步的隐状态 $h_t$ 依赖 $h_{t-1}$ 和当前输入。这个设计符合直觉但有个致命问题它是串行的。你没法在算出 $h_5$ 之前先算 $h_{50}$整个时间维度的计算被锁死了。在 GPU 这种靠大规模并行吃饭的硬件上这就意味着算力利用率极低。序列越长这个瓶颈越明显。第二个问题是长距离依赖。虽然 LSTM 的门控机制缓解了梯度消失但信息要在 $n$ 个时间步之间传递路径长度是 $O(n)$。路径越长梯度回传时被稀释得越厉害远距离的关联就越难学到。论文里专门用一张表对比了不同层类型的“最大路径长度”循环层是 $O(n)$卷积层是 $O(\log_k n)$而自注意力层是 $O(1)$——任意两个位置之间直接建立连接一步到位。这张表是理解论文动机的关键很多人读的时候直接翻过去了其实它是整篇论文的立论基础。第三个问题是卷积方案的局限。ConvS2S、ByteNet 这类模型确实能并行但卷积核的感受野是局部的要覆盖长距离关系就得堆很多层或者用膨胀卷积。层数一多信息在层间传递的路径又变长了。所以论文的核心诉求是找到一种既能并行计算、又能让任意两个位置直接交互的算子。答案就是注意力。1.2 摘要逐句翻译与要点提取原文摘要我按句拆开翻译顺便把每句的信息含量标出来The dominant sequence transduction models are based on complex recurrent or convolutional neural networks that include an encoder and a decoder. The best performing models also connect the encoder and decoder through an attention mechanism.主流序列转换模型基于复杂的循环或卷积神经网络包含编码器和解码器两部分性能最好的模型还会通过注意力机制把编码器和解码器连接起来。这句话交代了背景也埋了一个伏笔注意力机制早就存在只是过去它被当作循环网络的“辅助配件”而不是主角。We propose a new simple network architecture, the Transformer, based solely on attention mechanisms, dispensing with recurrence and convolutions entirely.我们提出一种新的简单网络架构——Transformer它完全基于注意力机制彻底摒弃了循环和卷积。这是全篇最核心的一句。“solely”和“entirely”这两个词用得很重作者在明确宣告这不是改良是替换。Experiments on two machine translation tasks show these models to be superior in quality while being more parallelizable and requiring significantly less time to train.在两项机器翻译任务上的实验显示这些模型质量更优同时更易于并行化训练时间也显著缩短。三个卖点质量、并行性、训练时间。Our model achieves 28.4 BLEU on the WMT 2014 English-to-German translation task, improving over the existing best results, including ensembles, by over 2 BLEU.我们的模型在 WMT 2014 英德翻译任务上取得 28.4 BLEU比此前最佳结果包括集成模型高出 2 BLEU 以上。单模型打赢别人的集成模型这是当年最有冲击力的一条。On the WMT 2014 English-to-French translation task, our model establishes a new single-model state-of-the-art BLEU score of 41.8 after training for 3.5 days on eight GPUs, a small fraction of the training costs of the best models from the literature.在 WMT 2014 英法翻译任务上模型在 8 块 GPU 上训练 3.5 天后取得 41.8 BLEU 的单模型最优成绩训练成本仅为文献中最佳模型的一小部分。注意“a small fraction”这个措辞作者在强调性价比而不是单纯刷分。We show that the Transformer generalizes well to other tasks by applying it successfully to English constituency parsing both with large and limited training data.我们还展示了 Transformer 的良好泛化能力无论训练数据充足还是有限它在英语成分句法分析任务上都取得了成功。最后一句是防守型论证回应“这只是为翻译定制的架构”这种质疑。1.3 论文结构地图与阅读顺序建议论文正文分八节背景、模型架构、为什么用自注意力、训练、结果、结论。其中第 3 节“Why Self-Attention”经常被跳过但它其实是整篇的论证核心解释了作者为什么敢把循环结构整个拿掉。我建议的阅读顺序不是从头到尾而是顺序章节读它的目的1Model Architecture第 3 节先建立整体结构图知道数据怎么流动2Scaled Dot-Product Attention3.2.1抓住唯一的核心算子3Why Self-Attention第 4 节理解设计动机回答“凭什么”4Training第 5 节拿到可复现的超参数5Results Ablation第 6 节看哪些设计真的有用6Background第 2 节补历史脉络可选按这个顺序读你会在最有动力的时候先拿到结构再用动机去验证最后用消融实验来确认理解得对不对。反过来读的话第 2 节一堆前人工作很容易让人失去耐心。2. 核心机制逐层拆解自注意力到底在算什么Transformer 的骨架其实只有三样东西注意力、前馈网络、残差加层归一化。位置编码算是第四样但它更像一个补丁。把这三样吃透整个模型就没有黑箱了。2.1 缩放点积注意力那个 $\sqrt{d_k}$ 从哪来论文给的公式只有一行$$\text{Attention}(Q,K,V) \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$$Q$、$K$、$V$ 分别叫查询、键、值。用检索系统打比方最好懂你手上有一个查询我想找什么数据库里有若干条记录的键每条记录打什么标签以及对应的值记录的实际内容。注意力的做法是拿查询去和每一条键算相似度把相似度归一化成权重再对所有的值做加权求和。相似度高的记录它的值就占更大比重。“点积”指的就是用向量内积来度量相似度。内积越大说明方向越一致。问题在于内积的结果会随维度增长。假设 $q$ 和 $k$ 的每一维都是独立同分布、均值 0、方差 1 的随机变量那么内积 $q \cdot k \sum_{i1}^{d_k} q_i k_i$ 的均值是 0方差是 $d_k$。维度 $d_k 64$ 的时候标准差就是 8数值分布相当宽。这会带来什么后果softmax 在输入数值差异很大的时候会变得极其尖锐最大值那一项的输出接近 1其余接近 0梯度几乎全部消失。除以 $\sqrt{d_k}$ 之后内积的方差被压回 1softmax 的输入落在一个比较温和的区间梯度才能正常回传。这就是缩放的全部理由没有任何神秘之处。import math import torch import torch.nn.functional as F def scaled_dot_product_attention(q, k, v, maskNone): # q: (batch, heads, q_len, d_k) # k: (batch, heads, k_len, d_k) # v: (batch, heads, k_len, d_v) d_k q.size(-1) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn F.softmax(scores, dim-1) return torch.matmul(attn, v), attn这几行代码里有两个容易写错的点。第一transpose的两个维度必须是最后两维也就是序列长度那一维和 $d_k$ 那一维写错了会变成对 batch 做矩阵乘形状对不上但报错信息很难懂。第二mask 要在 softmax 之前加加完之后被遮住的位置变成负无穷softmax 之后自然就是 0。有人喜欢在 softmax 之后乘 mask那样归一化分母会算错权重加起来不等于 1训练会不稳定。2.2 多头注意力为什么不是一个大头如果只做一次注意力模型只能学到一种“关注模式”。但语言里的关系是多层次的有的位置需要关注语法主语有的需要关注相邻词有的需要关注标点边界。单个注意力头很难同时兼顾。多头注意力的做法是把 $d_{model}$ 维的向量切成 $h$ 份每份独立做一次注意力最后拼接再线性变换。论文里 $d_{model} 512$$h 8$所以每个头的维度 $d_k d_v d_{model} / h 64$。注意总计算量和单头差不多因为维度被摊薄了并没有变成 8 倍开销。这一点常被误解很多人以为是“八个头等于八倍计算”。import torch.nn as nn class MultiHeadAttention(nn.Module): def __init__(self, d_model512, n_heads8, dropout0.1): super().__init__() assert d_model % n_heads 0 self.d_model d_model self.h 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 _split_heads(self, x, batch): seq_len x.size(1) # 把 d_model 拆成 (h, d_k)再把 heads 提到 seq 前面 x x.view(batch, seq_len, self.h, self.d_k) return x.transpose(1, 2) def forward(self, query, key, value, maskNone): batch query.size(0) q self._split_heads(self.w_q(query), batch) k self._split_heads(self.w_k(key), batch) v self._split_heads(self.w_v(value), batch) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn self.dropout(F.softmax(scores, dim-1)) out torch.matmul(attn, v) # (batch, h, seq, d_k) out out.transpose(1, 2).contiguous().view(batch, -1, self.d_model) return self.w_o(out), attn_split_heads里的view加transpose是这套实现最容易翻车的地方。view要求张量内存连续如果前面的操作导致了不连续得先.contiguous()。另外view(batch, seq, h, d_k)的切分顺序决定了哪几维归到哪个头必须和后面拼接的顺序严格对应否则信息会被打乱模型照样能训练但学到的表示是错位的表现会明显变差。2.3 位置编码扔掉循环之后顺序信息从哪来自注意力有个天然的缺陷它是置换等变的。把输入序列里的词顺序打乱输出只会跟着一起打乱注意力的计算结果本身对顺序不敏感。这对语言来说显然不行“狗咬人”和“人咬狗”完全不是一回事。所以必须显式地把位置信息注入进去。论文选择了正弦余弦函数$$PE_{(pos, 2i)} \sin\left(\frac{pos}{10000^{2i/d_{model}}}\right), \quad PE_{(pos, 2i1)} \cos\left(\frac{pos}{10000^{2i/d_{model}}}\right)$$$pos$ 是位置$i$ 是维度索引。偶数维用 sin奇数维用 cos。这么设计的巧妙之处在于对任意固定的偏移量 $k$$PE_{posk}$ 都可以表示成 $PE_{pos}$ 的线性函数用三角函数的和角公式就能推出来。这意味着模型有可能通过线性变换学会“相对位置”的概念而不只是死记绝对位置。另外这个函数是确定性的不需要训练参数测试时可以外推到比训练集更长的序列上。作者的实验还对比了可学习的位置嵌入两者效果几乎一样。选正弦版本主要是为了外推能力和参数效率。class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000, dropout0.1): super().__init__() pe torch.zeros(max_len, d_model) pos torch.arange(0, max_len).unsqueeze(1).float() div torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] torch.sin(pos * div) pe[:, 1::2] torch.cos(pos * div) self.register_buffer(pe, pe.unsqueeze(0)) # (1, max_len, d_model) self.dropout nn.Dropout(dropout) def forward(self, x): x x self.pe[:, :x.size(1)] return self.dropout(x)用register_buffer而不是普通属性是为了让这部分张量跟着模型一起搬到 GPU但不会被当成需要训练的参数。这是个细节但很多人第一次写的时候用self.pe pe结果报“张量在不同设备上”的错。2.4 那些不起眼但缺一不可的配角残差连接。每个子层注意力和前馈网络外面都套了一层x Sublayer(x)。因为加法要求维度一致所以论文里把所有子层和嵌入层的输出维度统一成 $d_{model} 512$。残差的作用是给梯度提供一条高速公路让 6 层甚至更深的网络能稳定训练。层归一化。论文用的是 Post-LN也就是LayerNorm(x Sublayer(x))归一化放在残差相加之后。要注意后来的很多实现改成了 Pre-LN先归一化再进子层训练更稳定不需要 warmup 也能收敛但论文原版是 Post-LN复现的时候别搞混否则超参对不上。前馈网络。结构是两层线性加 ReLU中间维度 $d_{ff} 2048$也就是先升到 4 倍再降回来。论文的解释是它相当于两个 $1\times1$ 卷积提供逐位置的非线性变换。我个人的理解是注意力负责“在位置之间搬运信息”前馈层负责“对每个位置的信息做加工”两者分工明确。Dropout。论文在三个地方用了 dropout位置编码加完之后的输出、每个子层的输出残差相加之前、注意力权重上。基础模型用 0.1大模型用 0.3。注意力权重上做 dropout 是个容易被忽略的点它在 softmax 之后、和 $V$ 相乘之前作用是防止某些头过度依赖固定的位置关系。3. 整体架构与训练配置精读3.1 编码器与解码器的堆叠方式编码器由 $N 6$ 个相同的层堆叠每层两个子层多头自注意力、前馈网络。解码器同样是 6 层但每层有三个子层带掩码的多头自注意力、对编码器输出的交叉注意力、前馈网络。三个子层都套了残差加层归一化。输入侧的流程是词元索引 → 嵌入层乘上 $\sqrt{d_{model}}$→ 加位置编码 → 进编码器。乘 $\sqrt{d_{model}}$ 这一步是为了让嵌入的数值量级和位置编码匹配否则位置编码会盖过词嵌入的信号。这个细节在论文正文里只用了一个脚注说明但实测确实有影响。输出侧有个 shift right 的操作也就是把目标序列整体右移一位在开头补上起始符。这样在第 $t$ 步预测第 $t$ 个词的时候模型只能看到前 $t-1$ 个词符合自回归生成的因果约束。这个约束靠掩码实现。3.2 三种注意力的掩码差异这是实际写代码时最容易搞错的地方我把它单独拎出来讲。三种注意力用的掩码完全不同位置查询来源键值来源掩码类型目的编码器自注意力编码器输入编码器输入padding mask忽略补齐位解码器自注意力解码器输入解码器输入padding causal忽略补齐位且不能看未来解码器交叉注意力解码器编码器输出padding mask忽略编码器侧的补齐位causal mask 是一个下三角矩阵位置 $(i,j)$ 在 $j i$ 时为 0其余为 1。生成方式很简单def make_padding_mask(seq, pad_id0): # seq: (batch, seq_len) return (seq ! pad_id).unsqueeze(1).unsqueeze(2) # (batch, 1, 1, seq_len) def make_causal_mask(seq_len, device): mask torch.tril(torch.ones(seq_len, seq_len, devicedevice)) return mask.bool().unsqueeze(0).unsqueeze(1) # (1, 1, seq_len, seq_len)解码器自注意力需要把两个掩码做逻辑与因为形状分别是(batch, 1, 1, k_len)和(1, 1, q_len, k_len)广播之后正好是(batch, 1, q_len, k_len)。我踩过的一个坑是掩码用 0/1 还是 True/False 混用masked_fill的条件判断在布尔和整型上行为不一样统一成布尔最省事。还有一个隐蔽的坑如果整行都被 mask 掉比如某个样本全是补齐位softmax 会遇到全-inf的输入输出是 NaN。稳妥做法是把-inf换成一个很大的负数比如-1e9或者用torch.finfo(dtype).min。3.3 训练超参数的完整拆解论文的训练配置写得很实基本可以直接抄。我把关键项列成表项目基础模型大模型$d_{model}$5121024前馈中间维度20484096注意力头数816层数66Dropout0.10.3参数量65M213M训练步数100K300K总耗时12 小时8 卡3.5 天8 卡优化器用 Adam$\beta_1 0.9$$\beta_2 0.98$$\epsilon 10^{-9}$。注意 $\beta_2$ 是 0.98 而不是默认的 0.999这是专门为这个任务调过的二阶动量估计的衰减更快对稀疏梯度的响应更灵敏。学习率调度是这篇论文的另一个亮点公式是$$lrate d_{model}^{-0.5} \cdot \min(step^{-0.5},\ step \cdot warmup_steps^{-1.5})$$warmup_steps取 4000。这个调度分两段前 4000 步学习率线性增长之后按步数的平方根倒数衰减。warmup 的作用是让模型在参数还随机的时候先小步走等 Adam 的二阶动量估计稳定下来再加速。如果一上来就用大学习率注意力层的参数很容易被推到极端值直接训练崩溃。class NoamScheduler: def __init__(self, optimizer, d_model, warmup_steps4000, factor1.0): self.optimizer optimizer self.d_model d_model self.warmup_steps warmup_steps self.factor factor self.step_num 0 def step(self): self.step_num 1 lr self.factor * (self.d_model ** -0.5) * \ min(self.step_num ** -0.5, self.step_num * self.warmup_steps ** -1.5) for group in self.optimizer.param_groups: group[lr] lr self.optimizer.step()其他训练技巧还有三个标签平滑$\epsilon_{ls} 0.1$把目标分布从 one-hot 变成软分布缓解过拟合、提升 BLEU虽然困惑度会变差但作者明确说了“困惑度变差不要紧BLEU 才是目标”检查点平均把最后若干个检查点的参数平均起来几乎零成本拿到一点提升批量大小按词元数算每个批次约 25000 个源词元和 25000 个目标词元。3.4 复杂度与路径长度对比论文第 4 节那张表我完整翻译一下它是理解“为什么自注意力更好”的核心论据层类型每层复杂度最小顺序操作数最大路径长度自注意力$O(n^2 \cdot d)$$O(1)$$O(1)$循环层$O(n \cdot d^2)$$O(n)$$O(n)$卷积层$O(k \cdot n \cdot d^2)$$O(1)$$O(\log_k n)$受限自注意力$O(r \cdot n \cdot d)$$O(1)$$O(n/r)$$n$ 是序列长度$d$ 是表示维度$k$ 是卷积核大小$r$ 是邻域大小。三个指标各有含义复杂度决定算力开销顺序操作数决定并行度路径长度决定学习长距离依赖的难度。结论很清晰序列长度 $n$ 小于表示维度 $d$ 时自注意力的复杂度更低顺序操作数是常数并行性最好路径长度是常数远距离关联最容易学。当 $n$ 特别大的时候$n^2$ 项会成为瓶颈所以论文也提到了受限自注意力作为改进方向——这直接启发了后来一大票稀疏注意力和线性注意力的工作。4. 手写实现从零搭一个能跑通的 Transformer4.1 张量形状约定与调试准备写 Transformer 的代码九成的 bug 都出在形状上。我建议在动手之前先把形状规则写在纸上张量形状含义输入索引(batch, seq_len)词元 id词嵌入(batch, seq_len, d_model)稠密向量加位置编码后(batch, seq_len, d_model)保持不变拆头之后(batch, h, seq_len, d_k)头维度提前注意力分数(batch, h, q_len, k_len)每个头一份注意力输出(batch, h, q_len, d_k)加权求和结果合并头之后(batch, q_len, d_model)回到统一维度调形状问题有个很土但很有效的办法随便造一个小批量比如batch2, seq_len5, d_model8, h2然后逐层打印形状看哪一步和预期不符。大模型上跑不通的代码在小形状上往往一眼就能看出问题。4.2 编码器层与解码器层的实现class PositionwiseFeedForward(nn.Module): def __init__(self, d_model512, d_ff2048, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(dropout) def forward(self, x): return self.linear2(self.dropout(F.relu(self.linear1(x)))) class EncoderLayer(nn.Module): def __init__(self, d_model512, n_heads8, d_ff2048, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, n_heads, dropout) self.ffn PositionwiseFeedForward(d_model, d_ff, dropout) 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, src_mask): attn_out, _ self.self_attn(x, x, x, src_mask) x self.norm1(x self.dropout1(attn_out)) # Post-LN ffn_out self.ffn(x) x self.norm2(x self.dropout2(ffn_out)) return x class DecoderLayer(nn.Module): def __init__(self, d_model512, n_heads8, d_ff2048, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, n_heads, dropout) self.cross_attn MultiHeadAttention(d_model, n_heads, dropout) self.ffn PositionwiseFeedForward(d_model, d_ff, dropout) 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, memory, src_mask, tgt_mask): # 自注意力查询、键、值都来自解码器 a1, _ self.self_attn(x, x, x, tgt_mask) x self.norm1(x self.dropout1(a1)) # 交叉注意力查询来自解码器键值来自编码器输出 a2, _ self.cross_attn(x, memory, memory, src_mask) x self.norm2(x self.dropout2(a2)) ffn_out self.ffn(x) x self.norm3(x self.dropout3(ffn_out)) return xcross_attn那一行是很多人写错的地方。查询必须是解码器当前的表示键和值必须是编码器的输出。三个参数传反了代码照样能跑但模型学不到东西——因为查询和键值来自同一个分布的话交叉注意力和自注意力就没区别了。4.3 拼装完整模型与形状追踪class Transformer(nn.Module): def __init__(self, src_vocab, tgt_vocab, d_model512, n_heads8, n_layers6, d_ff2048, dropout0.1, max_len5000): super().__init__() self.src_embed nn.Embedding(src_vocab, d_model) self.tgt_embed nn.Embedding(tgt_vocab, d_model) self.pos_enc PositionalEncoding(d_model, max_len, dropout) self.encoder_layers nn.ModuleList( [EncoderLayer(d_model, n_heads, d_ff, dropout) for _ in range(n_layers)]) self.decoder_layers nn.ModuleList( [DecoderLayer(d_model, n_heads, d_ff, dropout) for _ in range(n_layers)]) self.fc_out nn.Linear(d_model, tgt_vocab) self.d_model d_model self.scale math.sqrt(d_model) def encode(self, src, src_mask): x self.pos_enc(self.src_embed(src) * self.scale) for layer in self.encoder_layers: x layer(x, src_mask) return x def decode(self, tgt, memory, src_mask, tgt_mask): x self.pos_enc(self.tgt_embed(tgt) * self.scale) for layer in self.decoder_layers: x layer(x, memory, src_mask, tgt_mask) return x def forward(self, src, tgt, src_mask, tgt_mask): memory self.encode(src, src_mask) out self.decode(tgt, memory, src_mask, tgt_mask) return self.fc_out(out) # (batch, tgt_len, tgt_vocab)拼装好之后一定要做一次形状自检。我习惯用 torchinfo 那种库打印一遍每层的输入输出形状或者手写一个 for 循环逐层打印。这一步花两分钟能省掉后面几个小时的调试。4.4 小规模训练验证与观察在正式跑翻译任务之前建议先做一个“过拟合单批次”的验证。做法是构造一个只有 2 到 4 个样本的小数据集反复训同一个批次几百步看损失能不能降到接近 0。如果降不下去说明模型或数据管道有问题跟数据量、超参没关系先修代码。model Transformer(src_vocab1000, tgt_vocab1000, d_model128, n_heads4, n_layers2, d_ff512, dropout0.1) optimizer torch.optim.Adam(model.parameters(), lr0.0, betas(0.9, 0.98), eps1e-9) scheduler NoamScheduler(optimizer, d_model128, warmup_steps400) criterion nn.CrossEntropyLoss(ignore_index0, label_smoothing0.1) for step in range(2000): logits model(src, tgt_in, src_mask, tgt_mask) loss criterion(logits.reshape(-1, logits.size(-1)), tgt_out.reshape(-1)) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scheduler.step() if step % 200 0: print(fstep {step:5d} | loss {loss.item():.4f} | lr {optimizer.param_groups[0][lr]:.6f})梯度裁剪这一步论文没强调但在实际实现里几乎是必备的。Transformer 的梯度偶尔会突然变大尤其是训练早期裁剪到范数 1.0 能显著降低崩溃概率。我试过不裁剪直接跑大概每三到五次训练就有一次在两千步左右炸掉加上裁剪之后基本没再遇到过。5. 实验结果与消融实验的读法5.1 翻译成绩怎么看模型英德 BLEU英法 BLEU训练开销FLOPs基础模型单模型27.338.1$3.3\times10^{18}$大模型单模型28.441.8$2.3\times10^{19}$此前最佳集成26.441.3约 $1.2\times10^{21}$把它读明白要注意两点。第一基础模型就已经打赢了之前所有模型包括别人的集成第二大模型的训练开销比此前最佳方案低将近两个数量级而效果还更好。作者想说的是这不只是刷分而是把性价比整个拉高了一个台阶。英德和英法两个方向的表现差异也值得一提。英法任务上提升更大从 41.3 涨到 41.8英德任务上从 26.4 涨到 28.4提升幅度更明显。一般认为英德的语言结构差异更大更依赖长距离依赖建模而这正是自注意力的强项。5.2 消融实验改动哪一项最伤模型论文的表 3 做了消融我把它整理成更容易理解的版本改动英德 BLEU 变化结论单头注意力$h1$维度不变下降约 0.9多头是真实增益不是装饰头数过多$h32$$d_k16$下降约 0.4头太多维度太小表达能力受损减小键维度 $d_k$略有下降键维度影响相似度度量的精度去掉位置编码下降明显顺序信息必须显式注入用可学习位置嵌入基本持平正弦版本不是关键可外推才是去掉 dropout下降明显dropout 是主要正则手段我最想强调的是第一行。很多人觉得多头是个花哨设计实测数据说明单头会掉将近 1 个 BLEU这在翻译任务里是相当大的差距。多头真正的价值在于提供了多个“观察视角”不同的头可以分工关注不同的语言现象。论文附录里画了注意力可视化的图能看到有的头专门盯紧相邻词有的头会在句法结构上形成明显的模式。5.3 注意力可视化能看出什么论文附录展示了几个注意力头的权重视图。比较有意思的现象是某些头在编码器里呈现出类似句法依存的结构比如动词会集中关注它的主语和宾语解码器的某些头会稳定地关注下一个位置的词像是在做一种隐式的对齐。这些模式不是人工设计的是训练自己长出来的。不过这里我要泼一点冷水可视化看着漂亮不等于模型真的“理解”了句法。后来的研究做了很多探针实验发现注意力权重和模型的实际行为之间关系复杂有时候改动权重分布并不影响输出。所以把注意力可视化当成一种诊断辅助工具就好别过度解读。6. 常见问题排查与踩坑实录6.1 训练侧的问题损失不下降一直卡在一个值附近。我遇到过几次原因各不相同。最常见的是学习率调度写错了比如把min写成了max或者step从 0 开始导致了除零。其次是标签平滑和目标序列的错位没对上模型在预测第 $t$ 个词但你喂给它的标签是第 $t-1$ 个那它就永远学不会。排查方法是把学习率固定成一个小常数比如 1e-4跑几百步如果损失能动说明是调度问题如果还是不动就是数据或标签的问题。损失突然变成 NaN。按概率从高到低排查全 mask 行导致的 softmax NaN、学习率过大导致的梯度爆炸、除零$\sqrt{d_k}$ 相关的实现错误、数值溢出。加梯度裁剪、把-inf换成-1e9、检查学习率曲线这三招基本能解决八成的情况。训练早期损失下降很快然后突然卡住。这通常是 warmup 步数设得太短的信号。论文用 4000 步是基于 25000 词元的大批量如果你用自己的小批量训练warmup 步数要按比例放大。经验公式是让 warmup 覆盖前 5% 到 10% 的总训练步数。6.2 实现侧的经典 bug展示维度搞错。softmax 必须在最后一维键的维度上做因为要对所有键做归一化。如果写成dim1就变成在头维度上归一化了权重完全没有意义。这个 bug 特别隐蔽因为形状没错损失也能慢慢下降只是效果差很多。掩码方向搞反。causal mask 的设计是“允许看自己和自己左边不允许看右边”。有人把上三角和下三角弄反结果模型只能看未来不能看过去训练时损失下降得特别慢。自检方法很简单打印掩码矩阵看第一行是不是只有第一个元素是 1。位置编码的除零或维度不匹配。torch.arange(0, d_model, 2)生成的是 $d_{model}/2$ 个数同时赋给pe[:, 0::2]和pe[:, 1::2]才刚好填满。如果 $d_{model}$ 是奇数两边长度不一致会直接报错所以 $d_{model}$ 必须能被头数整除、也最好是偶数。形状广播静默出错。PyTorch 的广播很方便但也很危险。比如掩码形状是(batch, 1, 1, k_len)分数形状是(batch, h, q_len, k_len)广播没问题但如果掩码不小心写成了(batch, k_len)广播规则会把维度对齐到错误的位置结果看起来能跑实际掩错了东西。养成打印形状的习惯。6.3 常见问题速查表现象可能原因排查动作损失卡住不动学习率调度错误、标签错位换固定小学习率试跑损失变 NaN全 mask 行、梯度爆炸、除零加裁剪、替换-inf、查形状效果远差于论文交叉注意力参数传反、softmax 维度错检查cross_attn三个入参训练极慢没并行、批量太小、多余同步检查 DataLoader 的num_workers推理重复输出同一词束搜索实现问题、长度惩罚缺失检查束搜索的归一化方式长序列效果崩位置编码外推不足换成相对位置编码或 RoPE显存爆掉批量过大、注意力矩阵 $O(n^2)$减批量或用梯度累积6.4 我踩过的几个具体坑第一个坑是嵌入层忘了乘 $\sqrt{d_{model}}$。这个缩放看着不起眼但因为位置编码的值域在 $[-1,1]$而随机初始化的嵌入值域大概在 $[-0.1, 0.1]$ 量级取决于初始化不加缩放的话位置信号会明显压过语义信号模型会先学位置再学语义收敛变慢。加上之后两者量级就匹配了。第二个坑是检查点保存了模型但没保存优化器状态。Transformer 的 Adam 优化器状态占的显存和参数差不多训练中断后如果只恢复模型参数、重置优化器二阶动量要重新累积前几百步的学习率曲线相当于浪费了效果会有肉眼可见的退化。第三个坑是过早下结论。我第一次复现的时候训到两万步看到 BLEU 只有十几个觉得论文有问题。后来才知道小规模配置下的 BLEU 曲线在前五万步都很平最后才突然抬起来。判断一个训练是否正常不要看绝对数值要看损失曲线的形状健康的曲线是前期快速下降、中期缓慢下降、后期在一个低水平上抖动。如果中期就完全平了那才是有问题。最后一个心得是关于验证频率。翻译任务上用 BLEU 做验证比用损失可靠但 BLEU 计算本身有开销每步都算不现实。我的做法是每 2000 步算一次 BLEU同时每步记录损失用损失曲线判断趋势、用 BLEU 判断质量。两个曲线偶尔会出现背离——损失在降但 BLEU 不动这通常意味着模型在优化那些不影响翻译质量的词上比如标点和功能词是正常现象。7. 从原论文延伸到后来的变体7.1 三大流派的分野《Attention Is All You Need》之后Transformer 的演化大致分成三条线。第一条是编码器系代表是 BERT 一脉只用编码器双向注意力靠掩码语言建模来预训练适合理解类任务比如分类、抽取、句子相似度。第二条是解码器系代表是 GPT 一脉只用解码器因果掩码靠自回归预测下一个词来预训练擅长生成。第三条是编码器-解码器系也就是论文的原版适合序列到序列的任务比如翻译、摘要、语音识别。理解这个分野有个实用价值当你接到一个新任务第一件事是判断它属于哪一类然后直接选对应的预训练模型而不是从头训一个原版 Transformer。这三条线在工程上的差异其实就是掩码方式、层数配置和预训练目标的不同核心算子还是那个缩放点积注意力。7.2 原始设计的现代改造论文里几个设计后来被普遍替换掉了。Post-LN 换成了 Pre-LN训练稳定性提升明显尤其对深层模型。正弦位置编码换成了可学习的绝对位置或 RoPERoPE 用旋转矩阵编码相对位置外推能力更好现在基本成了主流。层归一化在很多实现里换成了 RMSNorm去掉了均值中心化速度更快效果相当。激活函数从 ReLU 换成了 GELU 或 SwiGLU前馈层的结构也变成了门控形式。但注意这些改动都是在原设计基础上做的优化不是否定。原论文的那套配置在 2017 年的硬件和数据集条件下已经调得很到位了很多现代改动的收益在中小规模上并不明显。所以如果你是在做小规模实验我的建议是先用原版配置跑通再逐项替换一次只改一个变量用消融的方式确认每项改动确实带来了提升。7.3 视觉与其他领域的迁移Swin Transformer 这类工作把注意力搬到了视觉领域。核心的适配有两个一是把图像切成小块当作“词元”二是引入窗口化的注意力来降低 $n^2$ 的开销。因为图像的像素数远大于句子的词数直接用全局注意力算不动窗口注意力把计算范围限制在局部窗口内再用移位窗口来跨窗口通信。语音、蛋白质结构、时间序列这些领域的适配思路也类似先把原始信号切成离散单元再用某种方式编码位置或结构信息然后套用同一套注意力机制。这说明原论文的贡献不只是翻译效果好而是提出了一个足够通用的计算原语。如果你想顺着这篇论文往下读我的推荐顺序是先读视觉侧的 Swin Transformer理解窗口注意力的动机再选一个解码器系的语言模型论文看预训练目标怎么设计最后读关于位置编码外推的工作比如 RoPE 相关的论文。这四篇读完你基本能覆盖注意力机制从提出到成熟的主干路径。至于代码把这篇论文的实现从零手写一遍比读十篇解读都管用——我第一次真正搞懂多头注意力的维度变换就是在纸上把view和transpose的每一步形状都画出来之后。那种“原来是这样”的感觉是抄别人的代码永远得不到的。
返回列表