07 课 · 多头与掩码注意力
LESSON 07 · 卷七 挡板与分身

遮住未来,多长几双眼睛

第 6 课那张注意力表已经能让每个字「环顾全班」了,可它还有两个大问题。 第一,一张表只能表达一种「看谁」的方式,而语言里同时存在好几种关系; 第二,我们的任务是「猜下一个字」——训练时如果第 3 个字能直接看到第 7 个字, 它根本不用学,抄答案就行了。 这一课补上这两块:多头因果掩码

STAGE 01
一张表不够
同时存在的几种关系
STAGE 02
多头
拆宽度、各算各的、再拼回来
STAGE 03
挡板
因果掩码,装在哪一步
STAGE 04
装起来
一层多头掩码注意力的账

先玩这块挡板。下面是一句话 6 个字排成的 6 × 6 表:行 = 谁在问,列 = 谁能被看到。 右上那半个三角被一块灰板盖着 —— 那就是「未来」。 点任意格子看「第 i 个字能不能看到第 j 个字」,再点「掀掉挡板看看」。

招牌动画因果掩码 —— 一块挡板盖住未来,点格子问「能不能看」
谁在问 ↓ / 能看谁 →??×××××××××××××××挡板:不许看未来这一行能看到 / 偷看」能看到 1 个;掀掉挡板会偷看 0%」能看到 2 个;掀掉挡板会偷看 0%」能看到 3 个;掀掉挡板会偷看 0%」能看到 4 个;掀掉挡板会偷看 0%」能看到 5 个;掀掉挡板会偷看 0%?」能看到 6 个;掀掉挡板会偷看 0%
能看:✓被挡:×(分数是 −∞,百分比是 0%)条的长度 = 掀掉挡板后「偷看未来」的注意力占比
5 个字「?」看第 2 个字「」:能看,它拿到的注意力是 14.7%(这一行没有被挡的格子,掀不掀挡板都是这个数)。
第 0 个字「」只能看到自己:它那一行是 100% 给自己 —— 输出就等于它自己的 V,一点上下文都没有。位置越靠后能看的越多, 到最后一个字「?」才第一次能看全 6 个字。
整张表一共 6 × 6 = 36 个格子,被挡掉 6 × 5 ÷ 2 = 15 个(正好是右上那半个三角)。
👆 点格子看「能不能看」;点「掀掉挡板看看」把灰板拿掉 —— 每个位置都会偷看掉一大截(最后一格是 0%,因为它本来就没有未来)。 真实代码里这块挡板叫 因果掩码(causal mask), 每一行要看的格子和这里一模一样,只是句子长得多。
三种视图:谁能看谁(✓ / ×)、分数(−∞ 还是具体数字)、百分比(0% 还是多少)。右边那排条是「掀掉挡板后会偷看多少」—— 位置越靠前偷看越多(第 0 个字偷看 82.5%), 最后一个位置是 0%(它本来就没有未来)。
图 7-1 · 第 i 个字只许看第 0 ~ i 个;被挡的格子分数是 −∞、百分比严格是 0%
挡板不是「少看一点」,
是把 100% 重新分给能看的那些。

三句话讲完这一课

  • 一张注意力表只能表达一种关系(第 6 课那张 6 × 6 的表)—— 把宽度拆成几份,每个头各算一张,最后拼回来,就有了几种视角(多头)。
  • 猜下一个字时不许看未来:把 j > i 那些格子的分数设成 −∞, softmax 之后严格是 0%(因果掩码)。
  • 挡板要装在 softmax 之前:不是「少看一点」, 而是把 100% 重新分给能看的那几个 —— 装错地方,注意力会被白白扔掉。
换个熟悉的地方想 · 考场隔板

每个考生只能看自己左边的卷子(前文),不能看右边(未来)。 这就是因果掩码。真实代码里它只有一行:

mask = 上三角全 True(j > i 的位置)
scores = Q · Kᵀ ÷ √d
scores = scores.masked_fill(mask, −∞) ← 挡板装在这里
w = softmax(scores) ← 被挡的格子自动变成 0%

注意为什么是 −∞ 而不是 0:softmax 第一步是取指数,e⁰ = 1 照样能分到一份百分比; 只有 e^(−∞) = 0 才是真正的「看不到」。这一条面试很爱问。

第 1 站

一张表装不下

第 6 课那张表长这样:第 i 行是「第 i 个字怎么分配它的 100%」。 它只能装下一种分配方案。

可同一句话里,关系从来不止一种。「三加五等于?」这 6 个字里, 至少同时存在着三件事:谁挨着谁(顺序)、数字该找哪个运算符(语法)、答案位该凑哪两个操作数(语义)。

核心实验三种关系 —— 同一句话,三种完全不同的「看谁」
0 个头:挨着看
100×××××4060××××4060×××4060××4060×4060
行 = 谁在看,列 = 看谁;「×」是因果掩码挡住的未来
1 个头:数字找运算符
50×××××100××××5050×××100××100×100
行 = 谁在看,列 = 看谁;「×」是因果掩码挡住的未来
2 个头:答案位找算式的两个操作数
100×××××100××××100×××100××100×4060
行 = 谁在看,列 = 看谁;「×」是因果掩码挡住的未来
现在选的是「第 2 个头:答案位找算式的两个操作数」——「?」同时盯住「加」和「五」—— 它要凑出答案,得先知道「加什么、加谁」。
三张表说的是同一句话,可它们关注的地方完全不同。第 6 课那张6 × 6 的表只能装下一种关系;而语言里这些关系是同时存在的。 多头就是让模型同时摆几张表:每个头一张,各自负责一种看法, 最后把结论拼起来。
⚠️ 这三张是手工画的示意图,用来讲清楚「多头想解决什么问题」。 真实模型里每个头长什么样,是训练自己长出来的 —— 没人告诉它「第 3 个头去盯运算符」。第 14 课训完模型之后,你会看到真头的样子。
点三张卡片切换:它们说的是同一句话,关注的地方却完全不同。一张 6 × 6 的表只能装下其中一种
图 7-2 · 三种关系示意图(手工画的):挨着看 / 数字找运算符 / 答案位找操作数
如果模型只有一个头,它会怎么处理这三种关系?先别往下翻,想 30 秒。
面试常问 · 多头不是「多花算力」
  • 宽度不变:d 维拆成 H 份,每头 8 ÷ H 维,加起来还是 8 维。
  • 参数不变:每头 Q/K/V 三张 [8, d/H],H 个头合起来正好 3d² —— 和 1 个头一模一样。
  • 计算量不变:每头算 T × T × (d/H) 次乘法,H 个头加起来还是 T × T × d。
  • 变的是「怎么看」:同一份算力,从「一个大视角」变成「几个小视角」。 工程上通常 8 ~ 32 个头,每个头 64 ~ 128 维。
第 2 站

多头:几副眼睛并行

做法就是「拆开算、拼回来」,一步都不新:

① 拆 d → H × d_h
每个头分到 d_h = d ÷ H 维。8 维、4 个头 → 每头 2 维;真实模型 28 个头 → 每头 128 维。
② 各算各的 自己的 Wq/Wk/Wv
第 h 个头用自己的一套矩阵([8, 2])算 Q/K/V,再走第 6 课那条老路:打分 → softmax → 取货。
③ 拼回来 concat → Wo
H 个头的输出按宽度接成 [T, d](宽度不变), 再乘一张输出矩阵 Wo 把所有头的结论混一次。
招牌动画多头 —— 换头数、点头选头、一个头一个头点亮
几个头
看哪个字
02
33.333.533.2×××
这一行加起来 100%
12
32.932.534.6×××
这一行加起来 100%
22
36.927.435.7×××
这一行加起来 100%
32
31.133.935×××
这一行加起来 100%
每张小表 = 一个头自己的「谁看谁」(6 × 6点一张表可以选中它,下面看它的细节「×」= 被因果掩码挡住(15 个格子)
拼起来:4 个头各出 2 个数,接成一整行 2 × 4 = 8 个数(宽度一个没多)
0.34-0.210.080.200.21-0.190.56-0.12→ 乘 Wo →-0.13-0.190.180.04-0.060.030.070.61
0 头看「」这一行:最关注的是「」(33.5%)。 它只能看 3 个位置(6 个里被挡掉 3 个)。
4 个头各算各的:头与头之间最大差 75.7 个百分点 —— 同一个字,不同的头看的地方明显不一样。它们用的不是同一套矩阵: 每个头都有自己的 Wq / Wk / Wv。
每头宽度
2
2 × 4 = 8 维(和 1 个头时一样宽)
每头 Q/K/V 参数
48 个数
3 张 [8, 2];4 个头加起来 192 个数
整层参数(含输出矩阵)
256 个数
和头数无关:1 / 2 / 4 / 8 头都是 256 个数
打分乘法次数
288
每头 6×6×24 个头加起来还是 6×6×8(与头数无关)
👆 换头数:每张小表里的数字会变(每头宽度变了,投影矩阵也就变了), 但总宽度、总参数、总计算量一个都不变 —— 多头不是「多花算力」,是把同一份算力分成几路并行地看。 真实模型:28 个头 × 128 维 = 3584 维, 一层 29,360,128 个参数(≈ 0.29 亿)、 打分 3,758,096,384 次乘法(≈ 37.58 亿)。
每张小表是一个头自己的「谁看谁」。换头数看两件事: ① 每张表里的数字变了(每头宽度变了,投影矩阵也就变了); ② 下面的参数卡、乘法次数卡一动不动 —— 头数不改变总量。
图 7-3 · H 张 [6, 6] 的注意力表并行算完,拼成 [6, 8],再乘 Wo 回到 [6, 8]

打开「因果掩码」开关,你会看到每张小表右上角都被挖掉一块 ——挡板对每个头一视同仁:不管这个头在看什么,它都不许看未来。

演示里最多 8 个头, 每个头 1 维。如果一路拆下去,每头只剩 1 维,会怎样?
第 3 站

挡板装在哪一步

页首动画里那块挡板,装的位置只有一个正确答案:softmax 之前。这件事看起来像细节,其实是「掩码到底有没有用」的分水岭。

两种装法都试一遍:同一行分数,一种先设 −∞ 再 softmax, 一种先 softmax 再清零。

核心实验加在之前 vs 加在之后 —— 同一行分数,两种顺序
看第几行这一行能看到 1 个位置,被挡 5
① 加在 softmax 之前(代码里的做法)行和 100%
能看
100%
−∞ → 0%
0%
−∞ → 0%
0%
−∞ → 0%
0%
−∞ → 0%
0%
?−∞ → 0%
0%
② 加在 softmax 之后(错的顺序)行和只剩 17.5%
能看
17.5%
被清零
0%
被清零
0%
被清零
0%
被清零
0%
?被清零
0%
同一行分数(第 0 行「」),只是顺序不同:
① 加在之前 —— 被挡的格子是 −∞,softmax 时自动变成 0%,剩下的重新分配,加起来仍然是 100%1.000000)。
② 加在之后 —— 先老老实实分出 100%,再把未来那些格子清零,行和只剩 17.5%0.174669):有 82.5% 的注意力白白扔掉了,位置越靠前扔得越多。
更直接的后果:加权收集的时候,这一行的输出也整体缩水到原来的 0.1747 倍 —— 位置 0 最惨,只剩 17.5%。 所以掩码必须加在 softmax 之前:不是「少看一点」,而是把 100% 重新分给能看的那几个
👆 换一行看看:越靠前的行被挡的格子越多,② 扔掉的也越多。 第 5 行(「?」)没有任何格子被挡, 两栏完全一样 —— 这也是「最后一个位置加不加掩码都一样」的原因。
左栏每行加起来还是 100%(被挡的格子拿 0%,剩下的重新分配); 右栏行和只剩一部分 —— 位置越靠前,白扔得越多
图 7-4 · 加在之前:行和恒 100%;加在之后:第 0 行只剩 17.5%
两栏的账(以第 0 行「」为例)
加在之前:分数 [s₀, −∞, −∞, −∞, −∞, −∞] → softmax → [100%, 0%, 0%, 0%, 0%, 0%]
 行和 = 1.000000 ✅

加在之后:分数 [s₀, s₁, s₂, s₃, s₄, s₅] → softmax → 先分掉 100%, 再把后 5 格清零
 行和 = 0.174669 ❌(白扔掉 82.5
 加权收集时,这一行的输出也缩水到原来的 0.1747

看清楚差别:掩码不是「把某些格子调小」,而是「把这几个格子从这张表里删掉, 然后把 100% 重新分给剩下的」。只有 −∞ 能做到这一点,因为 e^(−∞) = 0。

面试常问 · 掩码的三个常见追问
  • 为什么用 −∞ 不用 0:softmax 先取指数,e⁰ = 1 照样分到一份百分比; e^(−∞) = 0 才是真正的 0%。设成 0 等于「让那个字还能被看见一点点」。
  • 加在 softmax 之前还是之后:之前。之后加会把行和变成小于 100% (第 0 行只剩 17.5%),白扔注意力、输出缩水。
  • −∞ 会不会算出 NaN:如果整行都是 −∞,softmax 就是 0 ÷ 0。 但因果掩码的每一行至少留了自己(j = i),所以永远不会整行被挡 —— 这也是「允许看自己」这条设计的一个附带好处。
第 4 站

不遮会怎样:抄答案

回到页首那个「掀掉挡板」的开关。如果训练时真的不遮,会发生什么? 看看每个位置会分给「未来」多少注意力:

这个字偷看 5偷看 4偷看 3偷看 2偷看 1偷看 0平均偷看掉的比例82.5%66.3%49.9%11.9%17%0%37.9%
第 0 个字「」本来只能看自己,可它把 82.5% 的注意力分给了未来;平均下来每个位置偷看掉 37.9%。 最后一个位置是 0% —— 它本来就没有未来。
核心实验掀掉挡板就会抄答案 —— 拨一下开关,看箭头、读数、loss 曲线和生成结果一起变
掩码每个位置都盯着右边第一个字
没有挡板:每个位置都盯着右边第一个字?把 12.5% 的注意力给了右边第一个字把 16.5% 的注意力给了右边第一个字把 16% 的注意力给了右边第一个字把 4.6% 的注意力给了右边第一个字把 17% 的注意力给了右边第一个字右边没有字生成时这里什么都没有
每个位置分给「右边第一个字」:12.5% + 16.5% + 16% + 4.6% + 17% = 66.6%(5 个有右邻的位置,平均 13.3%)。
每个位置分给「未来」(右边所有字):82.5% + 66.3% + 49.9% + 11.9% + 17% + 0% = 227.6%(6 个位置各 100%,一共 600%;平均 37.9%) —— 这就是上面静态表里那一行「偷看掉的比例」。
注意这两笔账不一样大:「给右边第一个字」只是「未来」里最划算的那一个入口。 把分给未来的 227.6% 全压到它身上,就是抄答案的最优解 —— 因为训练时右边那个字是白送的正确答案。
平均偷看比例(6 个位置平均)
37.9%
227.6% ÷ 6 = 37.9%
直接算 peekAvg(4, 0) 也是 37.9%
被抄走的注意力总量(6 个位置各 100%,共 600%)
227.6%
其中分给「右边第一个字」:66.6%
关掉掩码才有这笔账
分给「右边第一个字」的总量(5 个有右邻的位置)
66.6%
66.6% ÷ 5 = 13.3%(平均一个位置)
随机初始化的头还没训练,这一份还不大
无掩码训练时右边写着答案 —— 抄就完事← 现在选中的模型
0.01.02.03.005千1万1.5万2万0.02
玩具曲线(示意):这不是跑出来的,是「抄答案」这条捷径的样子
第 4000 步:loss 0.13 → 第 2 万步:loss 0.02
生成结果:(点下面的「生成测试」)
有掩码只能靠前文真算
0.01.02.03.005千1万1.5万2万地板 1.651.65
真实训练记录(本课演示模型 30 万步里的前 2 万步)
第 4000 步:loss 1.67 → 第 2 万步:loss 1.65
生成结果:(把开关拨到这一侧,再点「生成测试」)
题目「三加五等于?」 正确答案是「八」
训练的时候,右边那个字就写在卷子上;生成的时候,右边什么都没有(还没写出来)—— 同一个位置,两个完全不同的任务。
现在手上是无掩码的模型,它吐出来的是:(还没跑)
无掩码的模型学会了复读(loss 4000 步就掉到 0.13), 有掩码的模型老老实实答「八」(第 2 万步 loss 1.65)。训练 loss 漂亮 ≠ 模型有用 —— 这是这一课最值钱的一句话。
👆 先拨开关:的时候,5 条琥珀色粗箭头分别从每个字指向右边第一个字, 每格下面写着「把 X% 的注意力给了右边第一个字」,读数跟着算出总量;的时候箭头全没了,换成一块盖住「未来」的灰板, 每格下面改成「能看几个字」,所有读数一起变成 0.0%。 再点「生成测试」:无掩码的模型吐出复读的乱码,有掩码的模型老老实实答「八」。 左边那条玩具曲线掉得飞快、几乎贴地;右边那条是真实训练记录, 走了 2 万步才到 1.65 —— 两条曲线哪个「好看」,你心里有数。
先拨「掩码:关 / 开」:关着的时候,每个位置都伸出一条箭头指向右边第一个字(训练时那就是白送的答案),下面几行读数跟着算出「被抄走的注意力总量」; 一装上挡板,箭头全消失、换成一块盖住「未来」的灰板,所有读数一起归零。 再点「生成测试」:无掩码的模型吐出复读的乱码,有掩码的模型老老实实答出「八」。左边那条 loss 曲线掉得飞快、几乎贴地,右边那条走 2 万步才到 1.65 —— 哪条「好看」,你心里有数。
图 7-5 · 关掉掩码:平均偷看 37.9%、被抄走的注意力总量 227.6%;装上掩码:两笔账都是 0%,100% 重新分给前文
不遮未来去训练,你会发现 loss 掉得飞快(比加了掩码快得多)。这是好消息还是坏消息?
面试常问 · 生成第 1 个字的时候,掩码起作用了吗
  • 严格说:没起作用。生成第 1 个字时手上只有 1 个 token, 它那一行本来就没有未来可挡。
  • 但掩码必须一直在。训练和推理要处理的是同一套规则: 训练时整句话都在手上,只有掩码能保证「第 i 个位置看到的和生成时一样多」。 少了它,训练和推理就是两个不同的任务。
  • 顺手一个细节:掩码对最后一个位置也「没起作用」(它没有未来), 所以那一行加不加掩码输出逐位相同 —— 第 6 题就是这个。
第 5 站

拼起来:一层多头掩码注意力

到这儿零件齐了。把第 6 课那五行加上这两块,就是真实模型里的一层注意力:

完整代码(伪代码,和真实框架一一对应)
for h in range(H):             # ① 多头:每个头各算各的
  Q = x @ Wq[h] K = x @ Wk[h] V = x @ Wv[h]  # [T, d_h]
  s = Q @ K.T / sqrt(d_h)         # 打分 + ÷√d
  s = s.masked_fill(mask, -inf)      # ② 挡板:不许看未来
  o[h] = softmax(s) @ V          # 换算成百分比,再取货
out = concat(o) @ Wo            # ③ 拼回来,混一次

形状一路走一遍:x [6, 8] → 每头 Q/K/V [6, 2] → 每头权重 [6, 6] → 每头输出 [6, 2] → 拼接 [6, 8]Wo [6, 8]进出的形状一模一样,所以它能一层一层叠起来。

演示规模 · 这一层的参数
256 个数
Q/K/V 合计 192 个(恒 3d²)+ 输出矩阵 Wo 8×8 = 64
真实模型 · 一层注意力的参数
29,360,128
2 × 3584² + 2 × 3584×512 ≈ 0.29 亿(K、V 只有 512 宽:28 个头共用 4 组)
演示规模 · 打分乘法次数
288
4 个头 × 6×6×2 = 6×6×8(与头数无关)
真实规模 · 打分乘法次数
37.58 亿
1024 × 3584 × 1024;而且打分、取货各算一遍
面试常问 · 为什么 K、V 可以比 Q 窄(GQA)
  • 多头各看各的,但不必各存一份 K/V:把 28 个头分成几组, 同一组里的头共用一套 K、V。真实模型里 K、V 只有 512 维宽(= 3584 ÷ 7), 于是这两张矩阵的参数直接省掉 6/7。
  • 省的是什么:省的是参数和显存 —— 尤其生成时要缓存的那份 KV Cache (第 10 课会细算),它就是被这个宽度直接决定的。
  • 代价:同一组里的头只能看到同一份 K/V,表达能力略有损失; 工程上认为这笔买卖划算,所以现在的开源模型基本都是这么做的。

到这里,注意力的全部零件就到手了:Q/K/V(第 6 课)+ 多头(这一课)+ 掩码(这一课)。 但它现在还是一个「孤零零的机制」——下一课我们要给它配上残差、 LayerNorm 和 FFN,把它拼成一个能真正堆很多层的完整模块。

第 6 站 · 收官

把这一课钉在墙上

本课核心 · TAKEAWAY

第 6 课那张注意力表只能装下一种关系,而语言里同时存在好几种 —— 于是把宽度拆成 H 份,每个头用自己的 Wq/Wk/Wv 各算一张表, 拼回 [T, d] 再乘 Wo。头数变多只是「每头变窄」,总宽度、总参数(恒 3d² + d²)、总计算量(恒 T×T×d)一个都不变。 另一件事更硬:我们的任务是猜下一个字,所以第 i 个字不许看到未来—— 把 j > i 的格子分数设成 −∞(不是 0,e⁰ = 1 照样分到百分比), 而且必须装在 softmax 之前:掩码不是「少看一点」, 而是把 100% 重新分给能看的那几个。少了这块挡板, 模型会学会抄答案 —— 训练 loss 漂亮得可疑,一生成就崩。

这一课你亲手做完了

  • 看清「一张表不够」:同一句话里三种关系(挨着看 / 数字找运算符 / 答案位找操作数)同时存在,而 6 × 6 的表只能装一种。
  • 拆过宽度1 / 2 / 4 / 8 个头,每头 8 / 4 / 2 / 1 维, 加起来永远是 8 维;Q/K/V 总参数恒 192 个数、总参数 256 个数、 打分乘法恒 288 次。
  • 量过头与头的差别2 头 29.8 个百分点;4 头 78.8 个百分点;8 头 52.1 个百分点 —— 同一个字,不同的头看的地方明显不同。
  • 装过挡板6 × 6 的表里挡掉 15 个格子; 第 0 个字只能看自己(那一行 100% 给自己),最后一个字能看全部 6 个。
  • 比过两种装法:加在 softmax 之前 → 行和恒 100%; 加在之后 → 第 0 行只剩 17.5%、输出缩水到 0.1747 倍。
  • 算过「不遮会怎样」:第 0 个字偷看 82.5%、 平均偷看 37.9% —— 这些注意力本该属于前面的字。

学习小测验

已完成 0 / 60.0%答对 0
还没提交过 —— 每题先选一个选项,再点「提交」,答完就能看到诊断。
Q1为什么需要「多头」?一个头不够用在哪?
Q24 个头 × 2 维 和 1 个头 × 8 维,参数量一样吗?
Q3因果掩码把哪些格子挡住?
Q4掩码为什么必须加在 softmax 之前?
Q5被挡住的格子,为什么要把分数设成 −∞,而不是设成 0?
Q6同一句话、同一个模型,加不加掩码,哪些位置的输出完全一样?
NEXT · 第 8 课

Transformer 拼装:残差 · LayerNorm · FFN

注意力的零件齐了:Q/K/V、多头、掩码。可它现在还只是一个孤零零的机制——如果直接把它叠十几层,信号会在路上衰减成一团噪声。 下一课我们给它配三个保镖:残差连接(一条直通的高速公路)、LayerNorm(把音量调到标准大小)、FFN(每个字回自己座位独立消化)。 拼完这一课,你手上就有了一个能堆 4 层、真正跑起来的 Transformer。

从零手册 —— 下一课:Transformer 拼装:残差 · LayerNorm · FFN