多头注意力


文档摘要

多头注意力 本节摘要:一个注意力头一次只能学一种关系。八个头能学八种。单个自注意力头只算出一张注意力矩阵,把主谓一致、共指、长程篇章、句法分块全搅在一起,softmax 把它们糊成一坨,丢掉一半信号。2017 年 Vaswani 论文的修法是:并行跑多个注意力函数,每个有自己的 Q/K/V 投影,输出拼接。每个头在更小的子空间( 维)里独立工作,总参数量不变,表达能力却上去了。本节将带你从零实现多头切分(一次 reshape、一次 transpose,无循环),搞懂为什么 GPU 把它看成一次批矩阵乘( ),并梳理 2026 年的变体谱系:MHA、MQA、GQA、MLA——以及为什么 GQA 成了现代默认。

多头注意力

本节摘要:一个注意力头一次只能学一种关系。八个头能学八种。单个自注意力头只算出一张注意力矩阵,把主谓一致、共指、长程篇章、句法分块全搅在一起,softmax 把它们糊成一坨,丢掉一半信号。2017 年 Vaswani 论文的修法是:并行跑多个注意力函数,每个有自己的 Q/K/V 投影,输出拼接。每个头在更小的子空间(d_model / n_heads 维)里独立工作,总参数量不变,表达能力却上去了。本节将带你从零实现多头切分(一次 reshape、一次 transpose,无循环),搞懂为什么 GPU 把它看成一次批矩阵乘(heads 是免费加的),并梳理 2026 年的变体谱系:MHA、MQA、GQA、MLA——以及为什么 GQA 成了现代默认。读完本节,你能为新模型选对头数,并说清 DeepSeek 的 MLA 为什么能用低秩压缩把 KV 缓存再砍一刀。

学习目标

阅读完本节,你应当能够:

  1. 从零实现多头注意力的切分、并行计算、拼接与输出投影,并解释为什么这在 GPU 上是一次批矩阵乘。
  2. 说明每个头为什么能在不争夺表示预算的情况下专精一种关系(位置头、前驱 token 头、复制头、归纳头)。
  3. 区分 MHA / MQA / GQA / MLA 四种变体,以及它们各自在 KV 缓存显存上的代价。
  4. 给定参数预算与上下文长度,为新 Transformer 选择头数与 d_head

一、问题与直觉

单头注意力只算一张矩阵,这张矩阵只能捕捉一种关系——通常是让训练损失最小的那种。但真实语言里,主谓一致、共指、长程篇章、句法分块全绞在一起,一个 softmax 分布把它们糊成一坨,信号丢一半。

Vaswani 的修法:并行跑多个注意力函数,每个用独立的 Q/K/V 投影,输出拼接。每个头在 d_model / n_heads 维的子空间里工作,总参数量不变,表达能力上去。到 2026 年,多头是每个 Transformer 的默认配置,争议只在于几个头,以及 K/V 是否共享投影(GQA、MQA、MLA)。

三步成型:切分、并行、拼接投影

切分(Split)。 取形状 (N, d_model) 的 X,投影成 Q/K/V(仍为 (N, d_model)),reshape 成 (N, n_heads, d_head),再转置成 (n_heads, N, d_head)

并行注意(Attend)。 在每个头内部跑缩放点积注意力,各产出 (N, d_head)。各头在嵌入的不同子空间上工作,注意力计算期间彼此不通信。

拼接投影(Concat & Project)。 把头叠回 (N, d_model),乘一个可学习的输出矩阵 Wo((d_model, d_model))。Wo 是各头混合的舞台。

💡 为什么有效:每个头能在不与其他头竞争表示预算的情况下专精。探针研究(2019~2024)发现不同头有清晰分工:位置头、关注前一个 token 的头、复制头、命名实体头、归纳头(Induction Head)——后者正是上下文学习(In-Context Learning)背后的电路。

2026 年的变体谱系

变体 Q 头数 K/V 头数 代表模型
多头 MHA N N GPT-2、BERT、T5
多查询 MQA N 1 PaLM、Falcon
分组查询 GQA N G(如 N/8) Llama 2 70B、Llama 3+、Qwen 2+、Mistral
多头潜在 MLA N 压缩到低秩 DeepSeek-V2、V3

GQA 是现代默认,因为它把 KV 缓存显存砍掉 N/G 倍,质量几乎不掉。MLA 更激进——把 K/V 压进低秩潜空间,计算时再投影回来,花点算力省下大量显存。

二、从零实现

Step 1:从头切分

在第 02 节的 SelfAttention 外包一层切分/拼接。完整代码见原课程 phases/07-transformers-deep-dive/03-multi-head-attention/code/main.py

def split_heads(X, n_heads): n, d = X.shape d_head = d // n_heads return X.reshape(n, n_heads, d_head).transpose(1, 0, 2) # (heads, n, d_head) def combine_heads(H): h, n, d_head = H.shape return H.transpose(1, 0, 2).reshape(n, h * d_head)

一次 reshape 加一次 transpose,没有循环。PyTorch 的 nn.MultiheadAttention 底层也是这么干的。

Step 2:每个头并行跑缩放点积注意力

每个头拿到自己的 Q/K/V 切片,注意力变成批矩阵乘:

def mha_forward(X, W_q, W_k, W_v, W_o, n_heads): Q, K, V = X @ W_q, X @ W_k, X @ W_v Qh = split_heads(Q, n_heads) # (heads, n, d_head) Kh = split_heads(K, n_heads) Vh = split_heads(V, n_heads) scores = Qh @ Kh.transpose(0, 2, 1) / np.sqrt(Qh.shape[-1]) weights = softmax(scores, axis=-1) out = weights @ Vh # (heads, n, d_head) concat = combine_heads(out) return concat @ W_o, weights

真实硬件上 Qh @ Kh.transpose(...) 是一次 bmm。GPU 看到的是形状 (heads, N, d_head) × (heads, d_head, N) -> (heads, N, N) 的单次批矩阵乘。加头几乎不花额外成本

Step 3:分组查询注意力(GQA)变体

只改 K/V 投影:Q 有 n_heads 组,K/V 只有 n_kv_heads < n_heads 组,通过重复来匹配:

def gqa_project(X, W, n_kv_heads, n_heads): kv = split_heads(X @ W, n_kv_heads) # (kv_heads, n, d_head) repeat = n_heads // n_kv_heads return np.repeat(kv, repeat, axis=0) # (n_heads, n, d_head)

推理时只有 n_kv_heads 份副本进 KV 缓存,而不是 n_heads 份。Llama 3 70B 用 64 个查询头配 8 个 KV 头——8 倍缓存缩减

Step 4:探针看每个头学了什么

对一句短句跑 4 头 MHA,打印每个头的 (N, N) 注意力矩阵。即便随机初始化,不同头也会挑出不同结构——这既有一部分真信号,也有一部分子空间旋转对称性造成的偶然。

三、框架对比:PyTorch 与 GQA

PyTorch 一行版:

import torch.nn as nn mha = nn.MultiheadAttention(embed_dim=512, num_heads=8, batch_first=True)

PyTorch 2.5+ 的 GQA,通过 scaled_dot_product_attention 自动派发 Flash Attention:

from torch.nn.functional import scaled_dot_product_attention # Q 形状 (B, n_heads, N, d_head);K,V 形状 (B, n_kv_heads, N, d_head) out = scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=True)

头数怎么选? 2026 年生产模型的经验值:

模型规模 d_model n_heads d_head
小(~125M) 768 12 64
基础(~350M) 1024 16 64
大(~1B) 2048 16 128
前沿(~70B) 8192 64 128

d_head 几乎总是落在 64 或 128——它是「一个头能看多少」的单位。低于 32,头会与缩放因子 √d_head 打架;高于 256,「许多小专家」的好处就没了。

四、可复用产物

原课程产出 outputs/skill-mha-configurator.md:一个配置器 Skill,给定新 Transformer 的参数预算、序列长度、部署目标,推荐头数、KV 头数、投影策略(MHA / GQA / MLA)。

五、练习

  1. (Easy)n_heads 从 1 扫到 16(d_model=64 固定),在一个合成复制任务上画单层小模型的损失曲线。更多头是有帮助、是平台期、还是变差?
  2. (Medium) 实现 MQA(所有查询头共享一个 KV 头)。测量参数量比全 MHA 少多少;N=2048 时推理 KV 缓存缩小多少。
  3. (Hard) 实现一个迷你版 MLA:把 K/V 压成秩 r 的潜变量,缓存里只存潜变量,注意力时再解压。问 r 多大时缓存显存能降到全 MHA 的 1/8 以下,同时验证集困惑度(perplexity)下降不超过 1 bit?

本节要点回顾

  1. 单头糊成一坨:主谓一致、共指、长程、句法混进一个 softmax,信号丢一半——这就是为什么需要多头。
  2. 切分/拼接就是 reshape+transpose:无循环,GPU 看成单次批矩阵乘,加头近乎免费。
  3. 每个头专精一种关系:位置头、前驱 token 头、复制头、命名实体头、归纳头——归纳头是上下文学习背后的电路。
  4. 总参数量不变,表达力上升:每头在 d_head = d_model / n_heads 子空间工作。
  5. Wo 是头的混音台:拼接后乘 Wo(d_model×d_model),各头在此混合。
  6. 变体谱系:MHA(全独立)、MQA(共享 1 个 KV)、GQA(共享 G 个,现代默认)、MLA(低秩压缩,DeepSeek)。
  7. GQA 是 2026 默认:Llama 3 70B 用 64 查询头 + 8 KV 头,KV 缓存缩 8 倍,质量几乎不掉。
  8. d_head 几乎总是 64 或 128:低于 32 与缩放因子打架,高于 256 失去「小专家」好处。

下一节,我们将解决注意力的另一个先天缺陷——对顺序无感知,用正弦编码、RoPE、ALiBi 三种方案把位置信息注入模型。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U