上一篇结尾我留了一个问题:一组 \(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_modelMulti-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 是同一个数量级。
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]最后那步 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
2 │
3 ┌─────────┼─────────┐
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 的结构可以读成两个阶段:前半段刻意做 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
4 ↓
5 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。
标准 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?这些是下一站的事了。
参考资料#
- Vaswani et al., Attention Is All You Need, 2017 — https://arxiv.org/abs/1706.03762
- MIT 6.S191, Introduction to Deep Learning — https://introtodeeplearning.com/
- 从动态读取到 QKV:Self-Attention 的公式是怎么长出来的
- 从 RNN 到 Self-Attention:信息为什么一定要沿着 State 一步步传递?