§ 7.3 · Section

Multi-head · 多头注意力

Multi-Head Attention

上一篇我们讲清楚了 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,让它们各自学一个侧面

Analogy · 一双眼睛 vs 复眼

一组 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 的最终输出。

公式化写下来是这样:

公式 · Multi-Head Attention

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 base5128648 × 64 = 512 ✓
BERT-base768126412 × 64 = 768 ✓
BERT-large1024166416 × 64 = 1024 ✓
GPT-3 175B122889612896 × 128 = 12288 ✓
Llama-2 7B40963212832 × 128 = 4096 ✓
Llama-2 70B81926412864 × 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)864~65M
BERT-base121264110M
BERT-large241664340M
GPT-24825641.5B
GPT-39696128175B
LLaMA-2 70B806412870B

GPT-3 那个 "96 层 × 96 头"的怪兽——你可以想象一下:每个词都要经过 96 层,每层同时有 96 位"专家"从不同角度看它。而且这 96 层的头数各自不同、独立学习——最终形成了一个近乎万人规模的隐性专家团队,共同解读你输入的每一个 Token。

Analogy · 会诊团队规模

1 个头 = 一个全科医生看诊;
8 个头 = 一支小型医疗组(原论文规模);
96 个头(GPT-3) = 一家三甲医院整个专家团齐上——心内、神内、内分泌、影像、检验、外科……全都并行看你这一份病历。
每层结束后,主治医生(WO)把各科室的意见综合成一份诊断,然后送进下一层继续再一轮 96 科室会诊——一共 96 层。
这就是为什么 GPT-3 能"什么都懂一点"——它内部就是几千个隐性专家的联合运算。

一个古怪现象 · 所有头都盯着第一个词

研究者在观察真实模型的注意力图时,发现了一个到处都有、但一开始谁也解释不清的现象:相当多的头会把大量注意力扔给序列的第一个 Token——通常是句首标记,一个本身毫无语义的符号。有的头甚至把 80% 以上的分数都给了它。这个现象后来被命名为 Attention Sink(注意力汇聚点,也译作注意力池)

所谓注意力汇聚点,说白了就是一个"没事就往那儿看"的默认落脚处。为什么会形成这种东西?根源还在 softmax:softmax 强制要求一整行加起来等于 1,它不允许"这一行谁都不重要"这种表达。可现实中一个头经常碰上"这次没什么值得关注的"的情况——比如它专管指代消解,而这句话里压根没有代词。它没法输出"全零",那 100 分总得花掉,于是它学会了把分数倒进一个语义上最无害的垃圾桶里,也就是句首那个符号。

Analogy · 必须投票的会议与那张废票

想象一个规定极其死板的委员会:每次表决,每位委员手里的 100 票必须全部投出去,一票都不许留
大多数议题上,委员们会认真把票分给几个真正相关的提案。
可总有些议题跟某位委员的专业毫不相干——比如让财务委员去表决"办公室植物换哪个品种"。他确实没有意见,但规则不许他弃权。
怎么办?聪明的做法是:把 100 票全投给一个绝对无害的选项,比如"维持现状"。这样既满足了规则,又没有干扰真正的决策。
句首那个符号,就是 Transformer 里那个"维持现状"选项——所有想弃权的头,都把票扔在这里。
知道这一点有个非常实用的推论:做长文本流式处理时,如果你为了省显存把最早的那几个 Token 从缓存里扔掉,模型会突然开始胡言乱语。因为你把那个"废票箱"给撤了,那些想弃权的头被迫把票投给真正的内容词,判断全乱了。StreamingLLM 这篇工作的核心发现就是:无论如何,请把最前面那几个 Token 永久留在缓存里

这个现象还有一个副作用值得知道:它让模型量化变得更难。所谓量化,说白了就是把模型里的数字从高精度压成低精度以省显存,好比把高清照片压成缩略图。而注意力汇聚点那几个位置的数值往往异常巨大(业内叫 outlier,离群值),一压缩就失真严重,进而拖坏整个模型的输出。于是又催生了一批专门"给离群值开小灶"的量化方法(比如 SmoothQuant、AWQ)。一个由 softmax 归一化约束引出的小怪癖,一路影响到了模型部署的最末端——这就是深度学习工程里典型的连锁反应。

多头带来的三个好处

一个常见误解 · 头越多越好?

不一定。研究者们做过消融实验:头数并不是越多越好——头数太多时,很多头会变成"冗余的、可以直接砍掉"。实际上很多论文(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:

这些变体不改变多头 Attention 的核心思想——多角度并行看句子——但都在共享程度上做了取舍,用少量效果损失换取巨大的推理效率提升。这就是为什么今天的 LLM 能在一张卡上跑 128K 的上下文。

Analogy · 几个人共用一份资料

再回到"多科室会诊"那个场景,现在加一个现实约束:病历资料要复印,复印纸很贵
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 里有两种存储,速度差得离谱:

普通实现的问题在于:它老老实实按公式分四步走,每一步的中间结果都要写回大仓库、下一步再搬出来。那张 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",看看不同头都在做什么:

你看,这三个头之间的分工非常清晰——各管一个语法层面。而这一切都是训练时自动涌现的,没有任何人写过"你这个头去学介词连接"这样的规则。多头 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

Recap · 收束

Multi-Head 是 Attention 的"复眼版本"——把一组 QKV 变成多组并行,每组各看一个侧面,最后拼接融合。它让 Transformer 具备了同时理解语法、语义、指代、情感的能力,而计算成本却几乎不变。
到这里,你已经拿下了 Transformer 的核心运算模块——QKV、缩放、Softmax、多头、拼接投影。下一篇,我们不再谈单层运算,而是把镜头拉远,看看整个 Transformer 分化出了哪三大派系——BERT、GPT、T5。

☰ 主页
Xue Hai Wu Ya · § 7.3 · Multi-Head