跳到主要内容

给生成用的注意力 — 因果掩码、dropout 与多头

这一章做三次改造: 把上一章的注意力改成「只许看过去」的因果版; 给它加一层训练时才开的随机扰动;再把它复制成多份并行。 三次改造完,得到的就是 GPT 模型里真正运行的那个模块。 这是「阶段 1 造机器」的发动机部分,下一章把它装进整车。

1. 这一章讲什么

上一章结尾有一个没解决的问题:那个注意力版本里,每个词打分的时候能看到全句——包括它后面的词。 用来理解句子没问题;但我们的机器是干「猜下一个词」的,训练时如果让它看到答案,等于考试抄书。 这一章的三节分别解决:怎么遮(因果掩码)、怎么防背题(dropout)、怎么让它同时看好几个方面(多头)。

2. 顶层全景

上一章的注意力矩阵(6 词互看,每行和为 1):

Your journey starts with one step
Your [0.19 0.16 0.16 0.15 0.17 0.15 ] ← 问题:Your 能看到 step,
journey[0.20 0.16 0.16 0.14 0.16 0.14 ] 而生成时「未来」还不存在


三次改造:
① 因果掩码:上三角全部抹掉,每行只留「自己及之前」再重新归一
② dropout:训练时随机把一部分权重清零(剩下的加倍补偿),防背题
③ 多头:同一份输入复制几路、各自投影各自算,最后拼回来

图说:① 改的是「能看谁」,② 改的是「训练时怎么扰动」,③ 改的是「并行几份」。
三者互不冲突,叠加起来就是 MultiHeadAttention 类。

3. 核心原理

3.1 因果掩码:遮住未来,且一步遮干净

「猜下一个词」这个任务有个铁规矩:算第 i 个位置的上下文向量时,只许用第 1 到第 i 个词。 满足这条规矩的注意力叫因果注意力(causal attention,也叫掩码注意力)1

书里先演示一个「笨但直」的三步法,用来把道理摆清楚:

① 照旧算 softmax,得 6×6 权重矩阵
② 乘上一个下三角掩码(对角线以上全是 0),未来位置清零
③ 每行重新归一(除以剩下的和),让每行重新和为 1

真实数字长这样:journey 那行原来是 [0.2041, 0.1659, 0.1662, 0.1496, 0.1665, 0.1477], 遮完后四位置清零,重新归一化成 [0.5517, 0.4483, 0, 0, 0, 0]—— journey 只能看 Your 和自己2

看到这里应该起疑:第 ① 步的 softmax 已经把未来位置算进分母了,第 ③ 步再归一, 未来的信息是不是已经漏进来了?书里专门设框回答:不漏。 掩码清零再归一, 数学上等价于「一开始就只在未遮蔽的位置上算 softmax」——被遮的位置对结果没有任何贡献3

虽然三步法已经正确,实战用的是更省的一步法:在进 softmax 之前, 先把分数矩阵的上三角填成负无穷(-inf)。softmax 里有一项 e 的幂,而 e 的负无穷次方是 0, 所以这些位置出来自动是 0,行和自动为 1,不需要再归一4。 同一个 -inf 技巧后面还会用到:第 07 章的 top-k 采样(按概率从候选词里挑一个的那一步)用它砍掉长尾词。

代码层面的两个细节顺带记住,因为第 05 章组装整机时会撞上:

  • 掩码用 register_buffer 注册进模块——这样它随模型自动搬到 GPU, 不用手工保证「掩码和参数在同一块卡上」5;
  • 类里支持成批输入(两份文本 × 6 词 × 3 维,形状 [2, 6, 3]),一次算完6

3.2 dropout:训练时故意制造残缺

过拟合是第 02 章埋下的词:模型把训练数据背下来,而不是学会通用规律。 dropout 是对付它的经典手段:训练时随机把一部分神经单元的输出清零, 逼模型不能依赖任何特定的某几个单元;评估和使用时关掉7

在 GPT 这类模型里,dropout 有两个可加的位置:注意力权重算完之后,或权重乘上值向量之后。 书里选了前者,理由是实践中更常见8

演示用 50% 的丢弃率(真训 GPT 只用 0.1 或 0.2,50% 是为了肉眼可见): 一张全 1 的 6×6 矩阵过完 dropout,约一半变 0,剩下的从 1 变成 29。 为什么要翻倍?清掉一半元素后,整体「用力」会衰减;乘以 1/0.5 = 2 补偿, 保证训练和推理两个阶段注意力的平均影响力一致10

3.3 多头:一份注意力的钱,买几份不同的看法

多头注意力(multi-head attention)就是把上面的因果注意力复制几份, 每份有自己独立的一套 Q/K/V 矩阵,各看各的,最后把结果拼起来11。 直觉:一份注意力只能学一种「什么跟什么相关」,几份就能同时学好几种—— 比如一份学语法邻近,一份学指代关系。

书里给了两种实现,这是一个漂亮的「先讲明白再讲快」:

写法一:叠罗汉。 造两个 CausalAttention 实例并排跑,把两份输出沿最后一维拼接。 2 个头 × 每个头输出 2 维 = 拼出 4 维上下文向量,书里印出了 [2, 6, 4] 的真实输出12。 好懂,但慢:每个头各做各的矩阵乘法,而矩阵乘是整套计算里最贵的步骤。

写法二:切开看。 只造一份大的 W_query/W_key/W_value,一次矩阵乘算出所有头的 查询/键/值,然后把最后一维切成 头数 × 每头维度(head_dim = d_out ÷ num_heads)。

再转置成四维张量(排成多维格子的一堆数:向量是一维、矩阵是二维、它是四维)。

然后让 PyTorch 做批量(所有头塞进同一次调用、同时算完)矩阵乘——相当于把「每个头分别算」折进一次调用13

书里用一个小张量验证了两种算法(同一道题的两种解法步骤——叠罗汉与切开看)逐位相同14

写法一: 输入 ─→ [头1: 自己的 Wq/Wk/Wv 各乘一遍] ─┐
─→ [头2: 自己的 Wq/Wk/Wv 各乘一遍] ─┴→ 拼接

写法二: 输入 ─→ 一份大的 Wq/Wk/Wv 各乘一遍 ─→ 切成 num_heads 份
─→ 批量矩阵乘(所有头同时) ─→ 拼回 ─→ 再过一层线性(out_proj)

图说:数学上是同一件事;写法二把「最贵的矩阵乘」从每头一次降到总共一次。

写法二末尾还多一层 out_proj(输出投影):拼接后再过一个线性层。 不是必须,但很多 LLM 都这么干,书里为完整性加上了15

规模的参照物:书里演示用 2 头、2 维,而最小的 GPT-2(1.17 亿参数)是 12 头、768 维、 上下文长度 1,024;最大的 GPT-2(15 亿参数)是 25 头、1,600 维。 GPT 里输入维度和输出维度相同(d_in = d_out)16

4. 作者的判断与证据

有证据的: 掩码前后的权重矩阵([0.5517, 0.4483, …])、dropout 的 1→2 补偿、 两种多头实现的输出一致——全部是书里印出的可复现输出2914

作者的判断:

  • dropout 加在注意力权重之后是「实践中更常见的变体」——经验陈述8
  • out_proj 是「非必须但常见」——书里明说,附录 B 给了文献15
  • 50% dropout 只是演示;真训 GPT 时用 0.1~0.2——经验值17

判断(我们的,不是书里的): 「先叠罗汉再切开看」这一节是全书教学法的缩影: 先用最直白的写法把「是什么」固定下来,再证明高效写法与它逐位等价。 学任何高性能库(FLAGS 满天飞的那些)都该这么读:先找它的「叠罗汉版」。 如果错,会错在: 如果某天硬件变化让「一次大乘法再切」不再比「多次小乘法」快, 这个效率结论会翻转——但数学等价性不受影响。

5. 边界与局限

  • 因果掩码保证「训练时不看未来」,但它不解决生成的根本慢法:每生成一个词仍要重算一遍 (第 05 章的生成循环会看到这有多直接)。
  • dropout 是防过拟合的一种手段,不是全部;数据太小时它救不了场——第 06 章会亲眼看到 开着 dropout 照样把小说背下来。
  • 多头「不同头学不同方面」是广泛流传的说法,书里也是这么讲的; 但每个头具体学了什么,书里没给证据,今天可解释性研究也仍在争论。引用这句时留意分寸。
  • 这本书的多头是标准多头。2024 年后为省显存,主流模型改用分组查询注意力(GQA)等变体—— 超出本书范围,同族的 Grigorov 书(我们拆过)有专章。

6. 可带走的

  1. 生成式(逐个词往外写、每步都拿自己刚写的当输入的那一类)模型训练的铁规矩:算每个位置时只许看它及它之前;违反就是考试抄书。
  2. 遮未来的标准做法:分数矩阵上三角填 -inf 再 softmax——一步完成,无泄漏,免二次归一。
  3. 「先 softmax 再遮再归一」在数学上与一步法等价——这个等价性本身就是「无泄漏」的证明。
  4. dropout:训练时随机清零一部分单元,剩下的按比例放大补偿;使用时必须关(.eval())。
  5. 多头 = 几份独立投影各看各的再拼回;「叠罗汉」与「切大矩阵」数学等价,后者省矩阵乘。
  6. head_dim = d_out ÷ num_heads;GPT-2 small 是 12 头 × 64 维 = 768 维。
  7. register_buffer 注册掩码这类「不是参数但要随模型搬设备」的张量。
  8. -inf 这个技巧值得记住:它在这一章遮未来,在第 07 章还会用来砍采样长尾。

7. 原文地图

主题原书章原文位置
因果注意力的规矩3 Coding attention mechanismstext/11-ch03-3-coding-attention-mechanisms.txt:1052(搜「Hiding future words」)
三步法与真实矩阵同上text/11-ch03-3-coding-attention-mechanisms.txt:1156(搜「0.1921, 0.0000」) · text/11-ch03-3-coding-attention-mechanisms.txt:1174(搜「1.0000, 0.0000」)
信息泄漏框同上text/11-ch03-3-coding-attention-mechanisms.txt:1182(搜「Information leakage」)
-inf 一步法同上text/11-ch03-3-coding-attention-mechanisms.txt:1213(搜「negative infinity」)
dropout 与补偿同上text/11-ch03-3-coding-attention-mechanisms.txt:1253(搜「Masking additional attention weights」) · text/11-ch03-3-coding-attention-mechanisms.txt:1332(搜「factor of 1/0.5 = 2」)
register_buffer 与批输入同上text/11-ch03-3-coding-attention-mechanisms.txt:1390(搜「register_buffer」)
叠罗汉多头同上text/11-ch03-3-coding-attention-mechanisms.txt:1520(搜「MultiHeadAttentionWrapper」)
切大矩阵多头与批量矩阵乘同上text/11-ch03-3-coding-attention-mechanisms.txt:1632(搜「An efficient multi-head attention class」) · text/11-ch03-3-coding-attention-mechanisms.txt:1784(搜「exactly the same results」)
out_proj 与 GPT-2 规模同上text/11-ch03-3-coding-attention-mechanisms.txt:1800(搜「output projection layer」) · text/11-ch03-3-coding-attention-mechanisms.txt:1845(搜「117 million parameters」)

Footnotes

  1. 出处:「3 Coding attention mechanisms」第 1055 段(text/11-ch03-3-coding-attention-mechanisms.txt:1055,搜「masked attention」)。 原文:「It restricts a model to only consider previous and current inputs in a sequence」。

  2. 出处:「3 Coding attention mechanisms」第 1156 段(text/11-ch03-3-coding-attention-mechanisms.txt:1156,搜「0.1921, 0.0000」)与第 1175 段(text/11-ch03-3-coding-attention-mechanisms.txt:1175,搜「0.5517, 0.4483」)。 2

  3. 出处:「3 Coding attention mechanisms」第 1182 段(text/11-ch03-3-coding-attention-mechanisms.txt:1182,搜「Information leakage」)。 原文:「after masking and renormalizing, the effect of the masked positions is nullified」。

  4. 出处:「3 Coding attention mechanisms」第 1213 段(text/11-ch03-3-coding-attention-mechanisms.txt:1213,搜「negative infinity」)。 原文:「When negative infinity values (-∞) are present in a row, the softmax function treats them as zero probability.」

  5. 出处:「3 Coding attention mechanisms」第 1390 段(text/11-ch03-3-coding-attention-mechanisms.txt:1390,搜「register_buffer」)。

  6. 出处:「3 Coding attention mechanisms」第 1368 段(text/11-ch03-3-coding-attention-mechanisms.txt:1368,搜「torch.stack」)。

  7. 出处:「3 Coding attention mechanisms」第 1254 段(text/11-ch03-3-coding-attention-mechanisms.txt:1254,搜「randomly selected hidden layer units」)。

  8. 出处:「3 Coding attention mechanisms」第 1260 段(text/11-ch03-3-coding-attention-mechanisms.txt:1260,搜「two specific times」)。 2

  9. 出处:「3 Coding attention mechanisms」第 1322 段(text/11-ch03-3-coding-attention-mechanisms.txt:1322,搜「2., 2., 0., 2.」)。 2

  10. 出处:「3 Coding attention mechanisms」第 1332 段(text/11-ch03-3-coding-attention-mechanisms.txt:1332,搜「factor of 1/0.5 = 2」)。

  11. 出处:「3 Coding attention mechanisms」第 14 段(text/11-ch03-3-coding-attention-mechanisms.txt:14,搜「multi-head attention」)。

  12. 出处:「3 Coding attention mechanisms」第 1539 段(text/11-ch03-3-coding-attention-mechanisms.txt:1539,搜「d_out*num_heads=4」)与第 1583 段(text/11-ch03-3-coding-attention-mechanisms.txt:1583,搜「-0.4519」)。

  13. 出处:「3 Coding attention mechanisms」第 1632 段(text/11-ch03-3-coding-attention-mechanisms.txt:1632,搜「An efficient multi-head attention class」)与第 1709 段(text/11-ch03-3-coding-attention-mechanisms.txt:1709,搜「head_dim = d_out / num_heads」)。

  14. 出处:「3 Coding attention mechanisms」第 1784 段(text/11-ch03-3-coding-attention-mechanisms.txt:1784,搜「exactly the same results」)。 2

  15. 出处:「3 Coding attention mechanisms」第 1800 段(text/11-ch03-3-coding-attention-mechanisms.txt:1800,搜「output projection layer」)。 原文:「This output projection layer is not strictly necessary…but it is commonly used in many LLM architectures」。 2

  16. 出处:「3 Coding attention mechanisms」第 1845 段(text/11-ch03-3-coding-attention-mechanisms.txt:1845,搜「the smallest GPT-2 model (117 million parameters)」)。

  17. 出处:「3 Coding attention mechanisms」第 1264 段(text/11-ch03-3-coding-attention-mechanisms.txt:1264,搜「dropout rate of 50%」)。 原文:「When we train the GPT model in later chapters, we will use a lower dropout rate, such as 0.1 or 0.2.」