Multi-head · 多头注意力
上一篇我们讲清楚了 QKV——一组"提问、展示、贡献"就能完成一次注意力运算。但一组 QKV 只能同时看一种相关性;而语言里同时并存的相关性有好多种。于是 Transformer 采用了一个简单粗暴又极其有效的办法——"一组不够,就并行搞很多组"。这就是 Multi-Head Attention,多头注意力。
一个复杂病人被送进医院,主治医生看不太准——怎么办?叫多科室会诊。
心内科医生看心电图、看血压;
神经科医生看脑 CT、看反射;
内分泌科医生看血糖、看激素;
影像科医生看片子、看结构异常……
同一个病人,几个专家从不同角度各看一遍——每个人只看自己擅长的那一面,最后主治医生把所有意见汇总,得出一个综合诊断。
Multi-Head Attention 就是这场会诊——每个"头"是一个不同角度的专家。
一个头为什么不够 · 看合同只盯价格的那个人
先讲一个真实到有点扎心的场景。公司让小王去审一份采购合同,小王是财务出身,一拿到合同眼睛就直奔金额那一栏——单价、总价、税率、付款节点,全部核对得一丝不差,报告写得漂亮。三个月后出事了:合同第 17 条写着"乙方逾期交货,甲方不得解除合同,仅可要求延期",供应商晚交了两个月,公司干等着,一分钱违约金都要不到。
小王没有失职,他只是只有一双眼睛,而这双眼睛被训练成专看数字。公司真正需要的是三个人同时看这份合同:财务看金额、法务看违约条款、业务看交付节点和验收标准。三份意见汇总,才叫审完了合同。
单头注意力就是那个只看价格的小王。它的 WQ、WK、WV 三张换算表在训练中只能收敛到一种"关注偏好"上——可能学会了盯语法主谓,那它就顾不上指代;可能学会了追踪指代,那它对语气的敏感度就差。这不是它笨,而是一套参数只能编码一种匹配规则,就像一台只有一个频道的收音机,调到新闻台就听不到音乐台。
那能不能把这套参数做得更大、让它同时学会好几种规则?——理论上可以,实践上不行。原因是 softmax 那一步:它只输出一套百分比,一共 100 分。如果这个头既想把 60 分给主语、又想把 60 分给指代对象,加起来 120 分,softmax 不允许。一个头只有一份 100 分的预算,它必须做取舍。要想同时给两个不同维度各 100 分,唯一的办法就是——再开一个头,各自带一份 100 分的预算。
为什么一组 QKV 不够
回想一下上一篇的例子——"猫追球"里的"追",通过一次 QKV 运算融合了"猫"和"球"两个词的信息。但仔细想想,语言里的相关性其实是多层次并存的:
- 语法关系"追"和"猫"是主谓关系、"追"和"球"是动宾关系。谁是主语、谁是宾语,是纯语法层面的。
- 语义关系"追"是一个动作动词,通常伴随"运动"、"目标"、"猎捕"——这是语义层面的联想。
- 指代关系如果句子里有"它",就要判断"它"指谁——这是指代消解层面的。
- 情感/风格关系"追"这个动作出现在恐怖故事和萌宠故事里,语气完全不同——这是情感层面的。
一组 QKV 的WQ, WK, WV矩阵,训练完之后只能学会一种"看法"——比如可能主要学到了语法关系。如果你希望模型同时懂语义、指代、情感……你就需要好几组独立的 QKV,让它们各自学一个侧面。
一组 QKV 就像一双人眼——只有一个焦点,只能同时清晰地看一个东西。
多头 Attention 就像蜻蜓的复眼——几百个小眼同时看,每个小眼捕捉一个角度,最后合成一个全方位的视野。
机器要理解人类语言的复杂性,需要的正是这种"复眼式"的注意力——同时看语法、语义、指代、情感、语气……
多头是怎么做的 · 三步走
假设向量维度是 d = 512,我们想要 h = 8 个头。Multi-Head Attention 的运作是这样的:
1. 切分维度 · 512 → 8×64
每个头分到一份属于自己的"小视野"——总维度 512 切成 8 份,每份 64 维。
每个头有自己的一套 WQi, WKi, WVi——独立训练、独立学习。
2. 并行做 8 次 Attention
每个头拿自己的 QKV 独立跑一遍完整的 softmax(QKᵀ/√d)·V——各自得到一个 64 维的输出。
注意:这 8 次运算是完全并行的,不互相干扰、不互相等待。
3. Concat + 线性投影
把 8 个 64 维的输出拼接(Concat)回一个 512 维向量。
再乘以一个可学习的输出矩阵 WO——做一次"融合",把 8 个头看到的东西综合成一个统一的表示。
这就是 Multi-Head 的最终输出。
公式化写下来是这样:
MultiHead(Q, K, V) = Concat(head1, ..., headh) · WO
其中 headi = Attention(Q·WQi, K·WKi, V·WVi)
翻译成人话:每个头做一次独立的注意力 → 全部拼起来 → 再用一个线性层融合。
维度切分的算术 · 为什么"参数量一分不多"
多头最反直觉的一点是:请了 8 位专家,工资总额居然没变。这听起来像在骗人,但算一遍就明白了。关键在于——多头不是"复制 8 份 512 维的注意力",而是"把 512 维这块蛋糕切成 8 份,每人拿 64 维"。
【方案 A · 单头】d_model = 512,一个头独享全部 512 维
W_Q: [512, 512] → 262144 个参数
W_K: [512, 512] → 262144
W_V: [512, 512] → 262144
W_O: [512, 512] → 262144
──────────────────────────────
合计 1048576 ≈ 105 万个参数
【方案 B · 8 个头】d_model = 512,每头 d_k = 512/8 = 64
头1: W_Q1[512,64] + W_K1[512,64] + W_V1[512,64] = 3 × 32768 = 98304
头2: 同上 = 98304
...
头8: 同上 = 98304
──────────────────────────────
8 个头小计 = 8 × 98304 = 786432
再加输出融合矩阵 W_O: [512, 512] = 262144
──────────────────────────────
合计 1048576 ≈ 105 万个参数 ← 一模一样!
【为什么一样】
单头的 W_Q 是一张 [512, 512] 的大表;
8 头的 8 个 W_Qi 各是 [512, 64],横着拼起来正好也是 [512, 512]。
也就是说:多头只是「把同一张大表竖着切成 8 条」,
参数一个没多,一个没少 —— 变的只是「怎么用」。
【计算量呢?也一样】
单头: Q[n,512] × K^T[512,n] = n² × 512 次乘法
8 头: 每头 n² × 64,共 8 个 → 8 × n² × 64 = n² × 512 ← 完全相同
所以多头真正改变的是"一起算"还是"分开算"。单头是把 512 维当成一个整体去做一次 softmax,得到一套 100 分;八头是把 512 维分成 8 段,每段独立做一次 softmax,得到八套各 100 分。同样的参数、同样的算力,注意力分配的"自由度"翻了八倍。这就是它被称为"免费的多样性"的原因。
顺便厘清一个术语。dmodel 指的是模型主干上每个词的向量长度(比如 512、4096),所谓 dmodel,说白了就是"这个模型描述一个词用多少个数字";dk 指的是单个头内部的向量长度,说白了就是"每位专家分到多宽的视野"。两者的关系通常是 d_k = d_model / h,其中 h 是头数。工程上一般把 dk 固定在 64 或 128,然后靠模型变宽来增加头数——所以你会看到 GPT-3 那种 12288 ÷ 128 = 96 个头的配置。
| 模型 | dmodel | 头数 h | 每头 dk | 验算 |
|---|---|---|---|---|
| Transformer base | 512 | 8 | 64 | 8 × 64 = 512 ✓ |
| BERT-base | 768 | 12 | 64 | 12 × 64 = 768 ✓ |
| BERT-large | 1024 | 16 | 64 | 16 × 64 = 1024 ✓ |
| GPT-3 175B | 12288 | 96 | 128 | 96 × 128 = 12288 ✓ |
| Llama-2 7B | 4096 | 32 | 128 | 32 × 128 = 4096 ✓ |
| Llama-2 70B | 8192 | 64 | 128 | 64 × 128 = 8192 ✓ |
这张表里每一行都严格满足 h × d_k = d_model,这不是巧合,而是多头设计的硬约束——因为拼接回来的长度必须和主干一致,否则残差连接那一步加不起来。知道这个约束,你以后看任何模型的配置文件都能一眼看出它的头数是怎么定的。
每个头到底在关注什么 · 可视化告诉你
2017 年论文发表后,研究者们做了大量 Attention 权重的可视化实验,得到的结论非常有意思——不同的头,真的自动学会了看不同的东西,而且这完全是模型自己学出来的,没有人告诉它"你这个头去看语法"。
| 头的类型 | 典型学到的模式 | 例子 |
|---|---|---|
| 指代头 | 把"它/他/她/this/that"连到前文的具体名词 | "猫追球,它抓到了"——"它"→"猫" |
| 语法头 | 把动词连到它的主语和宾语 | "我吃苹果"——"吃"→"我"、"吃"→"苹果" |
| 位置头 | 关注邻近词(左右一两个) | 形成局部 N-gram 特征 |
| 语义头 | 把语义相似的词联结起来 | "医生"←→"护士"、"手术"、"病人" |
| 停用词头 | 盯着"the/a/of/的/了"这些高频功能词 | 用于构造句法骨架 |
| 句末头 | 都把注意力扔到句末标点 | 作为"默认锚点",是一种冗余保护 |
让人惊讶的是:这些分工完全是训练过程中自发涌现出来的。研究者们只是把 8 组 QKV 并排放在一起、扔给梯度下降,结果它们就自动分化出了明确的"专业身份"——像一个自组织的小型专家团队。
典型配置 · GPT/BERT/DeepSeek 用多少头
现代大模型的头数已经远远超过 2017 年论文里的 8 个。看一下几个代表模型的配置:
| 模型 | 层数 | 每层的头数 | 每头维度 | 总参数量 |
|---|---|---|---|---|
| Transformer 原论文 | 6 (Enc)+6 (Dec) | 8 | 64 | ~65M |
| BERT-base | 12 | 12 | 64 | 110M |
| BERT-large | 24 | 16 | 64 | 340M |
| GPT-2 | 48 | 25 | 64 | 1.5B |
| GPT-3 | 96 | 96 | 128 | 175B |
| LLaMA-2 70B | 80 | 64 | 128 | 70B |
GPT-3 那个 "96 层 × 96 头"的怪兽——你可以想象一下:每个词都要经过 96 层,每层同时有 96 位"专家"从不同角度看它。而且这 96 层的头数各自不同、独立学习——最终形成了一个近乎万人规模的隐性专家团队,共同解读你输入的每一个 Token。
1 个头 = 一个全科医生看诊;
8 个头 = 一支小型医疗组(原论文规模);
96 个头(GPT-3) = 一家三甲医院整个专家团齐上——心内、神内、内分泌、影像、检验、外科……全都并行看你这一份病历。
每层结束后,主治医生(WO)把各科室的意见综合成一份诊断,然后送进下一层继续再一轮 96 科室会诊——一共 96 层。
这就是为什么 GPT-3 能"什么都懂一点"——它内部就是几千个隐性专家的联合运算。
一个古怪现象 · 所有头都盯着第一个词
研究者在观察真实模型的注意力图时,发现了一个到处都有、但一开始谁也解释不清的现象:相当多的头会把大量注意力扔给序列的第一个 Token——通常是句首标记,一个本身毫无语义的符号。有的头甚至把 80% 以上的分数都给了它。这个现象后来被命名为 Attention Sink(注意力汇聚点,也译作注意力池)。
所谓注意力汇聚点,说白了就是一个"没事就往那儿看"的默认落脚处。为什么会形成这种东西?根源还在 softmax:softmax 强制要求一整行加起来等于 1,它不允许"这一行谁都不重要"这种表达。可现实中一个头经常碰上"这次没什么值得关注的"的情况——比如它专管指代消解,而这句话里压根没有代词。它没法输出"全零",那 100 分总得花掉,于是它学会了把分数倒进一个语义上最无害的垃圾桶里,也就是句首那个符号。
想象一个规定极其死板的委员会:每次表决,每位委员手里的 100 票必须全部投出去,一票都不许留。
大多数议题上,委员们会认真把票分给几个真正相关的提案。
可总有些议题跟某位委员的专业毫不相干——比如让财务委员去表决"办公室植物换哪个品种"。他确实没有意见,但规则不许他弃权。
怎么办?聪明的做法是:把 100 票全投给一个绝对无害的选项,比如"维持现状"。这样既满足了规则,又没有干扰真正的决策。
句首那个符号,就是 Transformer 里那个"维持现状"选项——所有想弃权的头,都把票扔在这里。
知道这一点有个非常实用的推论:做长文本流式处理时,如果你为了省显存把最早的那几个 Token 从缓存里扔掉,模型会突然开始胡言乱语。因为你把那个"废票箱"给撤了,那些想弃权的头被迫把票投给真正的内容词,判断全乱了。StreamingLLM 这篇工作的核心发现就是:无论如何,请把最前面那几个 Token 永久留在缓存里。
这个现象还有一个副作用值得知道:它让模型量化变得更难。所谓量化,说白了就是把模型里的数字从高精度压成低精度以省显存,好比把高清照片压成缩略图。而注意力汇聚点那几个位置的数值往往异常巨大(业内叫 outlier,离群值),一压缩就失真严重,进而拖坏整个模型的输出。于是又催生了一批专门"给离群值开小灶"的量化方法(比如 SmoothQuant、AWQ)。一个由 softmax 归一化约束引出的小怪癖,一路影响到了模型部署的最末端——这就是深度学习工程里典型的连锁反应。
多头带来的三个好处
- 表达能力翻倍同一个词能被同时"多角度"编码——语法、语义、指代、情感一起进入表示。
- 冗余容错一个头没学好没关系,其他头能补上——训练更稳定。
- 几乎不加成本因为总维度不变(512 = 8×64),多头的计算量和一头几乎一样——只是把矩阵重塑一下而已,GPU 并行照跑不误。
一个常见误解 · 头越多越好?
不一定。研究者们做过消融实验:头数并不是越多越好——头数太多时,很多头会变成"冗余的、可以直接砍掉"。实际上很多论文(Voita 等 2019)发现:把 BERT 里 80% 的头剪掉,模型效果几乎不变。
这说明当代大模型的多头设计里有大量"备胎"——它们没有学到独特的模式,只是训练时的随机产物。真正做事的核心头,可能只有那么十几个。这个现象也推动了后来一系列"注意力头剪枝"、"稀疏 Attention"的研究方向。
拼接与 WO 融合 · 主治医生签字的那一步
八个头各算出一个 64 维的结果,接下来要把它们变回一个 512 维的向量。这一步分两个动作,都很容易被略过,但第二个动作其实相当关键。
第一个动作叫 Concat(Concatenation,拼接)——所谓拼接,说白了就是把八段绳子首尾接成一根长绳:头 1 的 64 个数放前面,头 2 的 64 个数接上去……八段接完正好 512 个数。这一步没有任何计算,就是摆放位置,连一次乘法都没有。
第二个动作是乘上输出矩阵 WO,形状是 [512, 512]。为什么必须有这一步?因为拼完的那 512 个数是八份互不相干的意见并排放着,第 1 到 64 位是头 1 的话,第 65 到 128 位是头 2 的话,它们之间还没发生任何交流。WO 干的活儿,相当于会诊结束后主治医生把八份科室意见读一遍,写成一份统一的诊断书——它让每一个输出维度都能同时参考八个头的结论,把"八份并列的报告"揉成"一份综合判断"。
拼接与融合的形状变化(n 个词,d_model=512,h=8)
head_1 输出: [n, 64] ┐
head_2 输出: [n, 64] │
head_3 输出: [n, 64] │
head_4 输出: [n, 64] ├─ Concat 沿最后一维拼接 ─→ [n, 512]
head_5 输出: [n, 64] │ (零计算,纯摆放)
head_6 输出: [n, 64] │
head_7 输出: [n, 64] │
head_8 输出: [n, 64] ┘
↓
乘 W_O [512, 512]
↓
[n, 512] ← 与输入形状一致,可以做残差相加
拼完之后的 512 个数长这样:
[ 头1的64个数 | 头2的64个数 | ... | 头8的64个数 ]
↑ 此时八份意见还是"各说各话",互不通气
过一遍 W_O 之后的 512 个数长这样:
[ 综合了8个头的第1维 | 综合了8个头的第2维 | ... ]
↑ 每一维都是八个头意见的加权混合 —— 这才叫"汇总"
如果去掉 W_O 会怎样?
实验结论:模型效果明显下降。
因为下游的 FFN 层拿到的会是八块「拼贴画」,
而不是一份融合过的表示 —— 相当于让病人自己拿着
八张互相矛盾的科室报告去药房抓药。
并行优势 · 八个头是一次算完,不是排八次队
这里有一个非常容易误解的地方。看到"八个头各自独立计算",很多人脑子里的画面是一个 for 循环跑八遍。真实实现里根本没有循环——八个头是被塞进同一次矩阵乘法里一起算掉的。
做法是利用张量的形状变换(reshape)。所谓 reshape,说白了就是同一批数字换一种摆放方式,数字本身一个不动——就像把 24 颗糖从"一排 24 颗"改摆成"3 行 8 列",糖还是那 24 颗。
真实实现的形状变换(batch=B, 序列长 n, d_model=512, h=8, d_k=64)
1. 一次性算出所有头的 Q(注意:一个大矩阵,不是 8 个小的)
Q_all = X @ W_Q_all # [B, n, 512]
↑ W_Q_all 是 8 个头横向拼起来的大表
2. reshape:把 512 拆成 (8 个头 × 64 维)
Q = Q_all.view(B, n, 8, 64) # [B, n, 8, 64]
3. transpose:把"头"这一维提到前面,让每个头各自成为一个"批"
Q = Q.transpose(1, 2) # [B, 8, n, 64]
4. 打分:GPU 的批量矩阵乘法,8 个头同时算
scores = Q @ K.transpose(-2,-1) # [B, 8, n, n]
↑ 8 张注意力矩阵一次出炉
5. softmax 沿最后一维(每个头各自归一化,互不干扰)
weights = softmax(scores / 8.0) # √64 = 8
6. 加权 V
out = weights @ V # [B, 8, n, 64]
7. transpose 回来 + reshape 成拼接结果
out = out.transpose(1,2).reshape(B, n, 512)
8. 融合
out = out @ W_O # [B, n, 512]
整个过程零循环。GPU 看到的是几个巨大的矩阵乘法,
「8 个头」对它来说只是 batch 维度上多了个 8 —— 顺手的事。
这就是为什么多头几乎不带来速度损失。GPU 最怕的是"很多个小任务排队",最爱的是"一个大任务批量做"——多头恰好被工程实现改造成了后者。如果真的写成 for 循环跑八遍,速度会掉一大截,因为每次启动一个小矩阵乘法的固定开销(kernel launch)比实际计算还贵。这也提醒我们:同一个数学公式,实现方式不同,实际速度可以差十倍。
头数与维度的权衡 · 请多少人、每人看多宽
既然 h × d_k = d_model 是固定的,那给定一个模型宽度,"多请几个人但每人视野窄"和"少请几个人但每人视野宽"就成了一个必须做的取舍。这个取舍有实证结论。
| 配置(dmodel=512) | 每头视野 | 优点 | 缺点 |
|---|---|---|---|
| 1 头 × 512 维 | 极宽 | 单个头的匹配非常精细 | 只有一套 100 分,无法兼顾多种关系 |
| 8 头 × 64 维 | 适中 | 原论文的甜蜜点,多样性与精度平衡 | —— |
| 32 头 × 16 维 | 很窄 | 关注模式极其多样 | 16 维太短,点积的区分度不够,头容易学成噪声 |
| 128 头 × 4 维 | 荒谬地窄 | —— | 4 个数字算出的相似度基本是随机的,效果崩坏 |
原论文里做过这组消融实验,结论是:头太少不行,头太多也不行,8 到 16 是当时的最佳区间。为什么头太多会坏事?因为 dk 变得太小,向量太短,点积能表达的"方向差异"就不够用了。打个比方:用四个指标(身高、体重、年龄、性别)去判断两个人合不合得来,信息量根本不够,判断结果基本靠蒙;换成六十四个指标(还包括爱好、作息、消费观、家庭观……)才有可能判断得靠谱。dk 就是"判断依据的丰富程度",压得太薄,判断就失真。
后来的大模型解决这个矛盾的办法是:不做取舍,直接把模型加宽。GPT-3 把 dmodel 拉到 12288,于是它可以既有 96 个头、又保证每头 128 维——两头都要。这也是"大模型"这个词在结构层面的一层含义:不只是层数多,更是每一层都变得极宽,宽到能塞下近百位互不挤占的专家。
MQA 和 GQA · 现代大模型的多头新变体
2023 年之后,研究者又发现了多头设计的一个新问题——推理时 K/V 缓存太大。生成型任务(GPT 那种)在推理时,会把之前所有 Token 的 K 和 V 保存下来,供后面每一步查询——这就是所谓的 KV Cache。头数一多、序列一长,KV Cache 就爆炸——一次对话动辄占几十 GB 显存。
解决方案是让多头共享 K/V:
- MHA(原始多头)每个头有自己的 Q、K、V——最灵活也最占显存。
- MQA · Multi-Query Attention所有头共用一套 K、V,只有 Q 各自独立——显存占用直接降到 1/头数。PaLM、Falcon 使用。
- GQA · Grouped-Query Attention折中方案——把头分成几组,每组共用一套 K、V。LLaMA-2、Qwen、DeepSeek 广泛使用——是当前的事实标准。
- MLA · Multi-head Latent AttentionDeepSeek-V2/V3 提出的更激进方案——把 KV 压缩到一个低秩隐空间,进一步降低显存占用。
这些变体不改变多头 Attention 的核心思想——多角度并行看句子——但都在共享程度上做了取舍,用少量效果损失换取巨大的推理效率提升。这就是为什么今天的 LLM 能在一张卡上跑 128K 的上下文。
再回到"多科室会诊"那个场景,现在加一个现实约束:病历资料要复印,复印纸很贵。
MHA(原始多头)——八位专家,每人一套专属的病历副本,各自在上面圈画标注。最舒服,但要印八套,纸张成本最高。
MQA(多查询注意力)——八位专家,全体共用一套病历放在会议桌中间,每人拿着自己的问题清单去翻。纸张只印一套,成本降到八分之一。代价是:这套病历上的标注方式必须迁就所有人,谁都不能按自己的习惯改,判断的精细度略有损失。
GQA(分组查询注意力)——八位专家分成两组,内科四人共用一套、外科四人共用一套。印两套,成本是原来的四分之一,但同专业的人共用资料本来就顺手,损失几乎感觉不到。
这正是今天绝大多数开源大模型的选择——共享,但别共享得太狠。
把 Llama 2 的真实配置摆出来,你就能看到这个取舍在工程上到底省了多少。
Llama 2 系列的注意力配置(Q 头数 / KV 头数)
Llama-2-7B : 32 个 Q 头, 32 个 KV 头 → MHA(没省)
Llama-2-13B : 40 个 Q 头, 40 个 KV 头 → MHA
Llama-2-70B : 64 个 Q 头, 8 个 KV 头 → GQA,8 个 Q 头共用 1 套 KV
Llama-2-70B 如果用 MHA 会怎样?算一遍 KV Cache:
每 Token 的 KV 大小 = 2 × 层数 × KV头数 × 每头维度 × 2字节
【假设 MHA,64 个 KV 头】
2 × 80 × 64 × 128 × 2 = 2621440 字节 ≈ 2.5 MB / Token
4096 上下文 → 10.2 GB
32768 上下文 → 82 GB ← 一张 80G 的 A100 装不下
【实际的 GQA,8 个 KV 头】
2 × 80 × 8 × 128 × 2 = 327680 字节 ≈ 0.31 MB / Token
4096 上下文 → 1.3 GB
32768 上下文 → 10.2 GB ← 轻松容纳,还能同时服务好几个用户
省了整整 8 倍。而 Llama 2 论文报告的效果损失:
在多项基准上不到 1 个百分点 —— 这买卖太划算了。
顺便看一眼 DeepSeek 的 MLA 更狠:
它不是"共享 KV",而是把 KV 压缩成一个很小的隐向量再存,
用的时候现场解压。DeepSeek-V2 论文报告 KV Cache 降到
同规模 MHA 的约 1/14,同时效果还略有提升。
为什么"共享 K/V"能省这么多,而"共享 Q"没人做?回到上一篇的结论:K/V 是要长期存起来反复查的档案,Q 是用完就扔的一次性提问。缓存里躺着的全是 K 和 V,所以要省显存,只能从它们身上下手。Q 头数决定"有多少种提问角度",KV 头数决定"要存多少份档案"——GQA 的精髓就是保留全部提问角度,只压缩档案份数。八个人问八种不同的问题,但查的是同一份卷宗,这在现实中也再自然不过。
Flash Attention · 一次不改数学、只改搬运的加速
Flash Attention(2022 年由斯坦福 Tri Dao 等人提出)是过去几年最重要的工程优化之一。它有一个让人意外的性质:它算出来的结果和普通 Attention 一模一样,一位小数都不差,但速度快 2 到 4 倍、显存省 10 到 20 倍。既然数学完全相同,快出来的时间是从哪里省的?
答案是"搬运"。要理解这一点,先得知道 GPU 里有两种存储,速度差得离谱:
- HBM(High Bandwidth Memory,高带宽显存)就是你买卡时看的那个"80GB 显存",容量大但相对慢,带宽约 2 TB/s。相当于地下室的大仓库。
- SRAM(片上高速缓存)紧贴计算核心的一小块超快内存,A100 上只有约 20 MB,但带宽约 19 TB/s,快十倍。相当于你手边的操作台。
普通实现的问题在于:它老老实实按公式分四步走,每一步的中间结果都要写回大仓库、下一步再搬出来。那张 n×n 的注意力矩阵要在仓库和操作台之间往返好几趟——而这张矩阵在长序列下大得吓人(前一篇算过,n=10 万时单张就 20 GB)。搬运这些数据花的时间,远远超过实际做乘法的时间。
普通 Attention 的数据往返(每一步都要读写大仓库 HBM)
1. 读 Q, K → 算 S = QK^T → 把 S 写回 HBM [n×n 大小]
2. 从 HBM 读 S → 算 S/√d → 写回 HBM [n×n]
3. 从 HBM 读 S → 算 softmax → 写回 HBM P [n×n]
4. 从 HBM 读 P, V → 算 PV → 写回 HBM O [n×d]
n×n 大小的东西被读写了约 6 次。n=8192 时,
单个头单层的这张矩阵就是 6700 万个数(134 MB),
乘上 32 层 × 32 头 × 6 次往返 —— 搬运量以 TB 计。
Flash Attention 的做法(分块 + 融合,全程不落地 n×n)
把 Q、K、V 切成能塞进 SRAM 的小块,比如每块 128 行
for 每个 Q 块:
在 SRAM 里初始化一个累加器
for 每个 K/V 块:
把这一小块搬进 SRAM(操作台)
在操作台上算完 打分 → 缩放 → 局部 softmax → 加权
用「在线 softmax」技巧更新累加器,无需保存整块 S
把累加好的结果一次性写回 HBM
关键:那张 n×n 的大矩阵从来没有完整地存在过。
它被切成小块,在操作台上用完即弃。
显存占用从 O(n²) 降到 O(n)。
【打个比方】
普通实现 = 做一道菜,每切完一样食材就跑到地下室仓库
放回去,下一步再跑下去取出来。
Flash = 一次从仓库拿一小筐食材上来,在操作台上
连着切、炒、装盘一气做完,再下去拿下一筐。
菜谱一个字没改,做出来的菜完全一样 —— 只是不再来回跑楼梯。
Flash Attention 里最巧的一环是"在线 softmax"。softmax 按定义需要先知道整行的总和才能归一化,可分块计算时你手里只有一小段。它的解法是边算边修正:先按当前已知的最大值和局部和算出一个临时结果,等新的块来了、发现出现了更大的值,就用一个缩放因子把之前的累加结果整体校正一遍。相当于记账时先按暂估价入账,等发票到了再统一调整,最终账目和一次算完完全一致。
这个技术的影响力有多大?今天你能用到的 128K 上下文,几乎都建立在 Flash Attention 之上——PyTorch 2.0 起把它做成了内置算子,vLLM、SGLang、TensorRT-LLM 等主流推理引擎全部默认启用。它是一个典型的"没有它,长上下文根本不可能"的基础设施。而它的思路——不改数学,只改数据在存储层级间的流动方式——也成了此后一大批优化工作的范式。
注意力头的冗余与剪枝 · 那些拿着工资不干活的头
前面提到"剪掉 80% 的头效果几乎不变",这个结论值得展开——它既有趣,也让人对大模型的内部效率产生怀疑。
最经典的研究是 Michel 等人 2019 年的《Are Sixteen Heads Really Better than One?》和 Voita 等人 2019 年的《Analyzing Multi-Head Self-Attention》。两篇论文用的方法都很朴素:一个一个地把头关掉,看模型效果掉多少——所谓关掉,说白了就是把这个头的输出强行置零,相当于让这位专家当天请假。
| 实验发现 | 具体结果 | 怎么理解 |
|---|---|---|
| 大部分头可以单独去掉 | 关掉任意单个头,BLEU 分数平均只掉 0.1~0.3 | 就像一个几十人的部门,随便请一个人的假,业务照转 |
| 少数头是关键 | 关掉特定的某几个头,效果断崖式下跌 | 部门里总有那么一两个"谁都替不了"的骨干 |
| 可剪枝比例很高 | Voita 等人剪掉 48 个头中的 38 个,BLEU 只掉 0.15 | 近 80% 的头是"备胎" |
| 关键头集中在特定功能 | 存活下来的头多是"位置头"和"句法头" | 基础的语序和语法关系是不可替代的骨架 |
| 层与层不均衡 | 中间层的头冗余最多,第一层和最后几层最不能动 | 入口和出口是咽喉,中间是宽阔的高速路 |
这些结果引出一个自然的问题:既然这么多头没用,为什么不一开始就少设几个头?——答案有点反直觉:那些"没用"的头,在训练过程中是有用的。它们提供了大量随机的探索方向,让梯度下降有机会碰上好的解;等训练收敛,好的方案定型了,它们才显得多余。这就像招聘时多招了些人,最后真正撑起项目的是少数骨干——但如果当初只招那几个人,你根本不知道该招谁。多头提供的是"事后看来浪费、事前必不可少"的搜索广度。
实用层面,这催生了一批技术方向:头剪枝(Head Pruning,训练完把没用的头砍掉,模型变小变快)、结构化稀疏(在硬件支持下让被砍掉的头真正不参与计算)、以及一个更根本的思路——既然参数存在大量冗余,能不能一开始就设计成"参数很多但每次只激活一部分"?这个念头,直接通往第 7.5 篇要讲的 MoE。
直观感受一下 · 三个真实注意力头
放一段 BERT 论文附录里可视化的例子——句子是"I went to the store to buy some milk",看看不同头都在做什么:
- 头 A(介词头)"to" 强烈关注 "went"、"store"、"buy" ——把介词跟它连接的动词、名词绑定。
- 头 B(宾语头)"buy" 强烈关注 "milk"——把动词跟它的宾语连起来,形成动宾结构。
- 头 C(主语头)"went"、"buy" 都强烈关注 "I" ——把动作跟主语连起来。
你看,这三个头之间的分工非常清晰——各管一个语法层面。而这一切都是训练时自动涌现的,没有任何人写过"你这个头去学介词连接"这样的规则。多头 Attention 的自组织能力,是 Transformer 最迷人的性质之一。
最著名的一类头 · 归纳头与"照抄前例"的能力
可视化研究里最重要的一个发现,值得单独拿出来说:归纳头(Induction Head)。这是 Anthropic 团队在 2021 到 2022 年一系列可解释性研究里找到的一种特定功能的注意力头,它被认为是大模型"看几个例子就会做新任务"这一能力的物理基础。
所谓归纳头,说白了就是一个专门负责"照抄前例"的头。它干的活儿用一句话讲清:在前文里找到当前这个词上一次出现的地方,然后看一眼那次它后面跟的是什么,就预测这次也跟这个。
归纳头的工作方式(找 [A][B] ... [A] → 预测 [B])
输入序列: 张 三 是 医 生 。 李 四 是 老 师 。 张 三 是
↑ 当前位置
归纳头做的两步:
第 1 步:在前文里搜索"张三是"这个模式上次出现在哪
→ 找到位置 1~3
第 2 步:看一眼那次之后紧跟的是什么
→ "医"
结论:这次也大概率接"医"
实测中这类头非常明确,注意力权重会精确地
射向"上次同样模式之后的那个位置",亮得像一道激光。
【为什么它这么重要】
这是模型「从上下文里现学现用」的最小单元。
你在 Prompt 里写:
英文: apple 中文: 苹果
英文: dog 中文: 狗
英文: book 中文:
模型能接上"书",靠的正是这类归纳头 ——
它在前文找到了"英文:X 中文:Y"这个模式,
然后把同一个模式套用到新的 X 上。
Anthropic 的研究发现:小模型训练到某个特定阶段,
归纳头会「突然形成」,与此同时模型的
上下文学习能力也在同一时刻突然跃升。
两条曲线的拐点严丝合缝地对齐 ——
这是"涌现能力"少见的、能被定位到具体结构的案例。
归纳头有一个精巧之处:它至少需要两层注意力才能形成。第一层里有一个"上一个词头",负责把每个位置的信息里掺进"我前面那个词是谁";第二层的归纳头才能拿这个信息去做匹配。一层做不到,两层就出现了质变——这也从一个具体角度解释了"为什么深度有用":不是层数多就自动变强,而是某些能力需要多层配合才可能被构造出来。
顺带说明一件事,帮你校准对"可解释性"的期待:像归纳头这样功能干净、可以一句话说清的头,是少数。绝大多数头的注意力图看上去杂乱无章,研究者也说不清它在干什么。所以别把"多头 = 一群分工明确的专家"这个类比推得太远——它更像一个大部门里少数人职责清晰、多数人干着说不清但缺了就不转的杂活儿。这是当前大模型可解释性研究的真实处境。
四十行代码 · 一个完整可跑的多头注意力
把这一篇讲的所有东西合成一段代码,你会发现它短得不可思议——现代大模型最核心的模块,就是下面这几十行。
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, n_heads=8, n_kv_heads=None):
super().__init__()
assert d_model % n_heads == 0, "d_model 必须能被头数整除"
self.n_heads = n_heads
# n_kv_heads 小于 n_heads 就是 GQA;等于 1 就是 MQA
self.n_kv_heads = n_kv_heads or n_heads
self.d_k = d_model // n_heads
# 注意:这里是「一张大表」,不是 8 张小表 —— 效率的关键
self.W_q = nn.Linear(d_model, n_heads * self.d_k, bias=False)
self.W_k = nn.Linear(d_model, self.n_kv_heads * self.d_k, bias=False)
self.W_v = nn.Linear(d_model, self.n_kv_heads * self.d_k, bias=False)
self.W_o = nn.Linear(d_model, d_model, bias=False) # 主治医生签字
def forward(self, x, causal=False):
B, n, _ = x.shape
# 一次算完所有头,再 reshape 成多头形状
q = self.W_q(x).view(B, n, self.n_heads, self.d_k).transpose(1, 2)
k = self.W_k(x).view(B, n, self.n_kv_heads, self.d_k).transpose(1, 2)
v = self.W_v(x).view(B, n, self.n_kv_heads, self.d_k).transpose(1, 2)
# GQA:把少量 KV 头「复制」给多个 Q 头共用(几个人共看一份资料)
if self.n_kv_heads != self.n_heads:
rep = self.n_heads // self.n_kv_heads
k = k.repeat_interleave(rep, dim=1)
v = v.repeat_interleave(rep, dim=1)
# 打分 + 缩放:[B, h, n, n]
scores = (q @ k.transpose(-2, -1)) / (self.d_k ** 0.5)
if causal: # 生成模式:把右上三角遮住
mask = torch.ones(n, n, device=x.device).triu(1).bool()
scores = scores.masked_fill(mask, float('-inf'))
w = F.softmax(scores, dim=-1) # 每个头各自「一共 100 分」
out = w @ v # [B, h, n, d_k]
# 拼接:把 h 个 d_k 摞回 d_model
out = out.transpose(1, 2).contiguous().view(B, n, -1)
return self.W_o(out) # 融合成一份综合判断
# 用起来
mha = MultiHeadAttention(d_model=512, n_heads=8) # 标准 MHA
gqa = MultiHeadAttention(d_model=512, n_heads=8, n_kv_heads=2) # GQA,4个Q共用1套KV
x = torch.randn(2, 10, 512) # 2 个句子,每句 10 个词
print(mha(x).shape, gqa(x, causal=True).shape)
# torch.Size([2, 10, 512]) torch.Size([2, 10, 512]) ← 形状进出一致
三处值得回头对照本篇内容的地方:第一,W_q 是一个输出维度为 n_heads × d_k 的大线性层,不是八个小层——对应前面讲的"一次算完,零循环";第二,n_kv_heads 这一个参数就把 MHA、GQA、MQA 三种方案统一了,改个数字就切换,可见它们本质上是同一个结构的不同配置;第三,返回值形状和输入完全一致,都是 [B, n, 512]——这就是它能被无限堆叠的资格证。
Multi-Head 是一种"廉价的多样性"
回过头来审视 Multi-Head 的设计——你会发现它是深度学习史上少见的"零成本换效果"的巧妙设计。
想想看:如果你要让模型学会关注多种关系,直觉上应该给模型加更多参数、加更多层、加更多计算。但 Multi-Head 做的事却是——在不增加计算量的前提下,让参数"分工"。总维度 512,无论切成 1 头还是 8 头,矩阵乘法的总量是一样的;但切成 8 头之后,模型的"表达角度"翻了 8 倍。
这种"用同样的钱办 8 倍的事",在工程上极其罕见——它是 Transformer 能在 2017 年一夜爆红的关键之一。没有 Multi-Head,Attention 的表达能力会大打折扣,也就没有后来的 GPT 和 BERT。
Multi-Head 是 Attention 的"复眼版本"——把一组 QKV 变成多组并行,每组各看一个侧面,最后拼接融合。它让 Transformer 具备了同时理解语法、语义、指代、情感的能力,而计算成本却几乎不变。
到这里,你已经拿下了 Transformer 的核心运算模块——QKV、缩放、Softmax、多头、拼接投影。下一篇,我们不再谈单层运算,而是把镜头拉远,看看整个 Transformer 分化出了哪三大派系——BERT、GPT、T5。