Skip to main content
  1. Posts/

一组 QKV 就够了吗:Multi-Head Attention 到底在拆什么

·5296 words·11 mins
Note: This article is available in Chinese only. 本文暂无英文版本。 View original

上一篇结尾我留了一个问题:一组 \(W_Q, W_K, W_V\) 只定义了一个 matching / routing space,如果不同类型的信息需求需要不同的 routing 策略,一个 space 够吗?

这个问题让我停了下来。因为到 single-head self-attention 为止,每个 token 已经可以访问整个 sequence 了——任意两个位置之间的通路都是通的。如果表达能力不够,加大 \(d_k\) 不就行了?为什么还要搞出 Multi-Head?

这篇想把这件事掰开。

一套 routing space 到底缺什么
#

先回顾一下 single-head 在做什么。一个 token 拿到自己的 \(q_i\),和所有 candidate 的 \(k_j\) 做 matching,得到一组 routing weights,再按这些 weights 从 Values 中做 weighted read。整个过程只有一套 Q/K/V projection——一套 compatibility geometry、一套 routing distribution、一次 weighted read。

问题不是它看不到某个 token。它看到了整个 sequence。

问题在于:sequence 里不同位置之间可能同时存在很多种不同性质的关系。局部搭配、远距离依赖、句法结构、实体关联、指代关系——这些关系在 compatibility 的意义上完全不同。一个 token 对另一个 token 的"需要",可能同时发生在好几个完全不同的维度上。

而 single-head 把所有这些潜在的 relational needs 全压进了同一个 Q/K matching space。对于每个 candidate,只能算出一个 scalar compatibility score;对于每个 requester,只能形成一个 routing distribution;然后只做一次 weighted read。

一个足够高维的 single head 在理论上可以表达很复杂的东西。但它要求所有不同性质的 relational computation 共享同一个 compatibility space 和同一个 routing policy。这就像一个 service 只有一条 API 接口,所有类型的 request 都走同一条路径——也不是不能工作,但所有 caller 的需求被迫在同一个 matching 空间里竞争、妥协、折中。

Multi-Head 扩展的不是 receptive field(它本来就能看到整个 sequence),而是 routing space 的自由度。

每个 Head 是不是看不同 token?
#

我第一次听到"Multi-Head 就是多个 head 并行做 attention"时,下意识以为:每个 head 看一部分 token,不同 head 分工覆盖不同区域。这个直觉后来被推翻了。

事实是:每个 head 都看完整的 \(n\) 个 token。

一旦接受这件事,下一个反应就是:那 attention 的计算量不是要乘 \(h\) 倍?

后来发现这里有两条不同的 axis,不能混:

1sequence axis(token positions):n
2feature axis(representation dimension):d_model

Multi-Head 并不是在 sequence axis 上切分。它在 feature axis 上做 factorization。经典配置:

1d_model = 512
2num_heads = 8
3d_head = 64

不是 8 个各自完整 512 维的 attention——那确实就是 8 倍开销了。而是 8 个各自 64 维的 attention heads。每个 head 在一个较低维度的子空间里做 matching 和 read。

所以 attention 的主计算量:

18 个 head × n² × 64 = n² × 512

和一个 512 维的 single-head attention 是同一个数量级。

Single Head vs Multi Head:routing space 的数量变了,receptive field 没变

Multi-Head 不是把一个宽 attention 粗暴复制 \(h\) 份,而是把一次宽的 relational computation factorize 成多个较窄的 relational computations。 拆的是 representation width,不是 sequence。

是不是把 hidden state 直接切成几段?
#

理解了 \(d_{\text{model}} = 512\)、8 个 head、\(d_{\text{head}} = 64\) 之后,我的第二个错误直觉是:

1Head 1 看 hidden state 的前 64 维
2Head 2 看 65~128 维
3Head 3 看 129~192 维
4...

也就是把原始 representation 机械地切成 8 段,每段交给一个 head。

这是错的。

真正发生的是:每个 head 都从完整的 \(h_i \in \mathbb{R}^{512}\) 出发,通过自己独立的 learned projection 矩阵,得到一个 64 维的 view:

1完整 h_i ∈ R^512
2      ├── W_Q^(1), W_K^(1), W_V^(1) → Head 1 的 64-d view
3      ├── W_Q^(2), W_K^(2), W_V^(2) → Head 2 的 64-d view
4      ├── W_Q^(3), W_K^(3), W_V^(3) → Head 3 的 64-d view
5      ...

每个 head 的 64 维不是原始 512 维里的一段固定切片,而是 512 维的一个 learned linear combination。和上一篇里写过的一样——projection 不是挑选原始特征,而是学习一种新的 representation。

Head 是 learned subspace,不是原始 tensor 的机械 slicing。

工程实现中经常能看到的操作是:

1[n, 512] → 一次 projection → [n, 8 × 64] → reshape → [n, 8, 64]
每个 Head 从完整 hidden state 投影,不是切 feature 也不是切 token

最后那步 reshape 看起来像是"把输出切成 8 段",但语义上,这个 projection 矩阵本身就是 8 套独立 \(W\) 拼起来的——先做 learned projection,再 reshape 成多个 head。不要把 implementation 层面的 tensor layout 和 architecture 层面的语义混为一谈。

每个 Head 到底在算什么
#

到这里可以完整走一遍每个 head 的 computation 了。全文固定用这套 shape:

1sequence length = n
2d_model = 512
3num_heads = 8
4d_head = 64

输入是整个 sequence 的 representation \(H \in \mathbb{R}^{n \times 512}\)。

对第 \(r\) 个 head,先做三组 projection:

$$ Q_r = H W_Q^{(r)}, \quad K_r = H W_K^{(r)}, \quad V_r = H W_V^{(r)} $$

其中 \(W_Q^{(r)}, W_K^{(r)} \in \mathbb{R}^{512 \times 64}\),\(W_V^{(r)} \in \mathbb{R}^{512 \times 64}\)。得到 \(Q_r, K_r, V_r \in \mathbb{R}^{n \times 64}\)。

然后和 single-head 完全一样——在这个 head 自己的 64 维空间里跑一遍 scaled dot-product attention:

$$ A_r = \operatorname{softmax}\!\left(\frac{Q_r K_r^T}{\sqrt{d_k}}\right) $$

\(Q_r K_r^T \in \mathbb{R}^{n \times n}\):第 \(r\) 套 matching space 下的 routing score table。softmax 后还是 \(\mathbb{R}^{n \times n}\):第 \(r\) 套 routing distribution。

最后做 weighted read:

$$ O_r = A_r V_r $$

\(O_r \in \mathbb{R}^{n \times 64}\)。对于单个 token \(i\),\(o_i^{(r)} \in \mathbb{R}^{64}\) 就是它在第 \(r\) 套 learned relational space 下,从整个 sequence 中动态读取回来的 information。

和上一篇的心智模型完全对齐:Q/K 负责 routing decision,V 是 payload,\(A_r V_r\) 是真正的 read result。只不过现在有 8 套这样的 routing + read 在并行运行,每套工作在自己的 learned subspace 里。

一个 token 走哪个 Head?
#

我在这里冒出过一个自然的疑问:一个 requester token 是只走一个 head,还是像 MoE 那样通过某种 router 选择几个 head?

答案比想象中简单:标准 Multi-Head Attention 中,每个 token 同时经过所有 heads。

不是:

1token → router → 选择 Head 3 → 只走 Head 3

而是:

1            h_i
23   ┌─────────┼─────────┐
4   ↓         ↓         ↓
5 Head 1    Head 2    Head 3  ... Head 8
6   ↓         ↓         ↓
7 read 1    read 2    read 3  ... read 8

没有 router,没有 selection,所有 head 全部参与。这和 Mixture of Experts 结构有本质区别——MoE 通常由 router 选择一个 subset 的 experts 激活。Multi-Head Attention 是 all heads participate, always。

Concat 做了什么
#

8 个 head 各自输出 \(O_r \in \mathbb{R}^{n \times 64}\)。下一步是把它们 concat 起来:

1Head 1: [n, 64]
2Head 2: [n, 64]
3...
4Head 8: [n, 64]
5       ↓ Concat
6    [n, 512]

对于单个 token \(i\):

1[ Head 1 read result | Head 2 read result | ... | Head 8 read result ]

但 concat 本身没有做任何智能的信息融合。它只是把 8 个不同 routing channel 获得的 read results 并排摆在一起——collect,不是 mix。

我学到这里时的第一反应是:怎么可能只是 concat?不同 head 之间的信息什么时候真正融合?

为什么 Concat 后面还有 W_O
#

这个疑问的答案就是 \(W_O\):

$$ O = \operatorname{Concat}(O_1, \dots, O_h) \, W_O $$

shape 对一下。\(\operatorname{Concat}(O_1, \dots, O_h) \in \mathbb{R}^{n \times (h \cdot d_v)}\),也就是 \(\mathbb{R}^{n \times 512}\)。\(W_O \in \mathbb{R}^{512 \times 512}\)。输出 \(O \in \mathbb{R}^{n \times 512}\)。

很多资料把 \(W_O\) 描述成"把维度投影回 \(d_{\text{model}}\)"。这没错,但不是最值得理解的地方——在 \(h \cdot d_v = d_{\text{model}}\) 的标准配置下,concat 出来已经是 512 了,根本不需要"投影回来"。

\(W_O\) 真正重要的作用是:跨 head 的 feature-level learned fusion。

concat 后的 representation 是:

1[ Head 1 的 64 个 features | Head 2 的 64 个 features | ... | Head 8 的 64 个 features ]

这些 feature 来自不同 head,目前只是并排存在。\(W_O\) 是一个 dense linear layer,它允许最终 output 的任意一个 feature 同时依赖多个 head 的信息——比如 output feature 第 3 维可以由 Head 1 feature 7、Head 2 feature 34、Head 6 feature 3、Head 8 feature 51 共同决定。

所以这两步的分工是:

Concat = preserve / collect:把多个 routing channel 的 read results 收集起来。W_O = mix / fuse:在 feature level 做跨 head 的 learned integration。

完整 Multi-Head Attention 数据流

Multi-Head 的结构可以读成两个阶段:前半段刻意做 decomposition——把一个 relational problem 拆成多个独立 routing channels;后半段做 integration——通过 \(W_O\) 把多路 read results 融合回统一的 representation。

能不能每个 Head 先投影再 Concat?
#

看到 \(\operatorname{Concat} \to W_O\) 之后,我自然想到一个问题:既然最终需要融合,能不能每个 head 自己先乘一个 output projection,再 concat?数学上等价吗?

把 \(W_O\) 按输入维度对应的 head 切成 \(h\) 个 block:

1W_O = [ W_O^(1) ; W_O^(2) ; ... ; W_O^(h) ]

其中 \(W_O^{(r)} \in \mathbb{R}^{d_v \times d_{\text{model}}}\)(对应第 \(r\) 个 head 的那一段输入)。

那么:

$$ \operatorname{Concat}(O_1, \dots, O_h) \, W_O = O_1 W_O^{(1)} + O_2 W_O^{(2)} + \cdots + O_h W_O^{(h)} $$

注意:这是每个 head 的 \(O_r \in \mathbb{R}^{n \times 64}\) 通过各自的 \(W_O^{(r)} \in \mathbb{R}^{64 \times 512}\) 映射到完整 output space 后再 sum

这和"每个 head 各自 \(64 \to 64\) 再 concat"完全不同。后者的结果仍然是:

1[ Head 1 的 64 维 | Head 2 的 64 维 | ... | Head 8 的 64 维 ]

不同 head 在这个阶段还没有发生过任何 cross-head mixing——output feature 第 1 维只依赖 Head 1,第 65 维只依赖 Head 2,互不干涉。

而统一的 \(W_O\) 做的是:

1[ 所有 head 的 features ] → dense linear mixing → output

任意 output feature 可以同时使用多个 head 的信息。这才是 cross-head integration。

能不能更早融合
#

顺着"融合"这条线,我进一步想过:如果最后反正要融合多个 head 的结果,那是不是可以在更早的位置融合?比如先把多个 head 的 attention weights 合并成一个,再乘 Value?

想了一下,这在标准 MHA 中通常不等价。

每个 head 不只是 attention weights \(A_r\) 不同,Value 也不同:

$$ V_r = H W_V^{(r)} $$

\(V_1 \ne V_2 \ne V_3\)。所以标准 Multi-Head 是:

1A_1 × V_1 → read result 1
2A_2 × V_2 → read result 2
3A_3 × V_3 → read result 3
45    fusion

如果先融合 attention weights:

1A_1, A_2, A_3 → 某种合并 → A
2A × V → ???

第一个问题就是:\(V\) 用哪个?\(V_1, V_2, V_3\) 是不同 projection 的结果。

更根本的问题是:过早融合 attention pattern 会丢掉不同 routing channels 的 distinction。看一个简单例子:

1A_1 = [0.9, 0.1]   → Head 1 主要读 token A
2A_2 = [0.1, 0.9]   → Head 2 主要读 token B

两个 head 分别做 read,一个读到的主要是 token A 的信息,另一个主要是 token B 的信息——两路信息都保留了。

如果先平均:

1A = [0.5, 0.5]   → 平均读取两个 token

这就把"通过两个不同 routing policy 分别读取不同信息"压成了"用一个 routing policy 平均读取两个 token"。原来两路 read 各有侧重、互不干扰;融合之后变成一个平淡无奇的 uniform read。这不是同一个 representational structure。

标准 MHA 的 late fusion vs 过早融合 attention weights

标准 Multi-Head Attention 是一种 late-fusion architecture:先让每个 head 独立完成 route → read,再融合 read results。

有一个 caveat 要补:如果在某些特殊限制下,所有 head 共用同一个 \(V\),并且最终 fusion 只是 scalar linear combination,比如 \(\lambda_1 A_1 V + \lambda_2 A_2 V\),那因为线性性,\(= (\lambda_1 A_1 + \lambda_2 A_2) V\),先融合 attention weights 和后融合 read results 确实等价。但标准 MHA 不满足这些条件——不同 head 有独立的 \(V\) projection,\(W_O\) 也是 feature-level 的 dense mixing。

“Head 1 学语法"只是可能的结果
#

很多教程会这样写:

1Head 1 learns syntax
2Head 2 learns coreference
3Head 3 learns long-range dependencies

这种说法有启发性,但要搞清楚边界。

Architecture 提供的保证是:独立的 \(W_Q^{(r)}, W_K^{(r)}, W_V^{(r)}\) → 独立的 compatibility space → 独立的 routing distribution → 独立的 Value transformation。因此它允许不同 head 在训练过程中 specialize 到不同类型的 relational computation 上。

Architecture 没有保证的是:哪个 head 学什么。训练后的实际表现可能是各种情况——有的 head 确实 specialize 了,有的 head 之间行为高度重叠,有的 head 接近冗余,一个 head 可能同时捕获多种关系。

Attention visualization 里看到的 pattern 可以用来做 post-hoc 分析,但不要把它反过来写成定义。

Multi-Head 提供的是 specialization 的 structural freedom。specialization 本身是 empirical training outcome,不是 architectural guarantee。

计算量和 Head 之间的并行
#

回到最开始的直觉——h 个 head,FLOPs 是不是要乘 h?

前面已经算过了。如果 \(d_{\text{model}} = 512\),\(h = 8\),\(d_{\text{head}} = 64\),那么每个 head 的 score computation 是 \(O(n^2 \times 64)\),8 个 head 总共 \(O(8 \times n^2 \times 64) = O(n^2 \times 512)\),和一个 512 维的 single-head attention 同数量级。

Multi-Head 的设计价值之一就在这里:用 representation factorization 换取更多 relational diversity,而不是单纯扩大 FLOPs。

另一个值得提的点是 heads 之间的并行性。在 attention computation 阶段,每个 head 独立做自己的 \(Q_r K_r^T\)、softmax、\(A_r V_r\),彼此之间没有依赖关系,逻辑上可以并行。但不要脑补成"8 个 head = 8 个 CPU thread 或 8 个 CUDA stream”——GPU 实现中通常把 heads 作为 tensor 的一个 dimension:

1Q, K, V: [B, H, N, D]
2Q @ K^T → [B, H, N, N]

由 GPU 的 batched / fused computation 并行执行。architecture 的 branch 是逻辑结构;实际 implementation 是 batched tensor computation。

写在最后
#

走到这里,Multi-Head Attention 在我脑子里的 picture 是这样的:

同一个 token 的 representation 被投影到多个独立的 learned relational subspaces 里。每个 head 都在完整 sequence 上,通过自己的 Q/K matching space 产生自己的 routing pattern,在自己的 Value space 中完成一次 dynamic read。多个 read results 保持独立直到 read 全部完成,再通过 concat 收集、\(W_O\) 做跨 head 的 feature mixing,重新形成统一的 token representation。

压短一点:Multi-Head Attention 不是让同一个 attention 重复很多遍,而是把一次大的 relational computation factorize 成多个独立的 routing channels,分别 route、分别 read,再 late-fuse。

但 Multi-Head Attention 到这里仍然只是一个 information routing / read mechanism。它负责的是"从 sequence 中动态读取信息"。一个完整的 Transformer layer 还包含 FFN、residual connection、layer normalization 等组件——attention 读回来的信息怎么和原有 representation 整合?读完之后还做了什么 transformation?这些是下一站的事了。

参考资料
#

Related

从 RNN 到 Self-Attention:信息为什么一定要沿着 State 一步步传递?

上一篇《从 RNN / LSTM / GRU 到早期 Attention》结尾我留了一句话:早期的 attention 并不是为了取代 RNN 而出现的,它只是夹在 encoder–decoder 中间解决 fixed-context bottleneck;但“根据当前需求去读取一组 representation”这种机制,并不一定要依附在 RNN 上。这篇接着讲那个“另一个故事”。