多头注意力 本节摘要:一个注意力头一次只能学一种关系。八个头能学八种。单个自注意力头只算出一张注意力矩阵,把主谓一致、共指、长程篇章、句法分块全搅在一起,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 缓存再砍一刀。
阅读完本节,你应当能够:
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)背后的电路。
| 变体 | 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 压进低秩潜空间,计算时再投影回来,花点算力省下大量显存。
在第 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 底层也是这么干的。
每个头拿到自己的 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) 的单次批矩阵乘。加头几乎不花额外成本。
只改 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 倍缓存缩减。
对一句短句跑 4 头 MHA,打印每个头的 (N, N) 注意力矩阵。即便随机初始化,不同头也会挑出不同结构——这既有一部分真信号,也有一部分子空间旋转对称性造成的偶然。
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)。
n_heads 从 1 扫到 16(d_model=64 固定),在一个合成复制任务上画单层小模型的损失曲线。更多头是有帮助、是平台期、还是变差?r 的潜变量,缓存里只存潜变量,注意力时再解压。问 r 多大时缓存显存能降到全 MHA 的 1/8 以下,同时验证集困惑度(perplexity)下降不超过 1 bit?d_head = d_model / n_heads 子空间工作。Wo 是头的混音台:拼接后乘 Wo(d_model×d_model),各头在此混合。d_head 几乎总是 64 或 128:低于 32 与缩放因子打架,高于 256 失去「小专家」好处。下一节,我们将解决注意力的另一个先天缺陷——对顺序无感知,用正弦编码、RoPE、ALiBi 三种方案把位置信息注入模型。