Learn

多头注意力

实际 Transformer 很少只用一套 Q、K、V,而是使用多头,并行计算多次注意力,再拼接成完整的自注意力矩阵

1. 多头注意力的原理

单头只有一套匹配方式。三个 token 的关系可能同时包含:

  • 语法:谁修饰谁
  • 指代:代词指向哪个名词
  • 位置:谁离当前 token 更近

单头自注意力 WQ,WK,WVW_Q,W_K,W_V 很难同时学好这些关系。多头把 dmodeld_{model} 拆成若干个子空间,每个头各自投影、各自打分、各自加权;最后再拼回去。

2. 多头注意力计算

假设输入的矩阵:

X∈R3×4X \in \mathbb{R}^{3 \times 4}

三个 token,每个 dmodel=4d_{model}=4。取 h=2h=2 个头,每个头的宽度:

dk=dmodel/h=2d_k = d_{model}/h = 2

第 ii 个头有自己的三组参数:

Qi=XWQi,Ki=XWKi,Vi=XWViQ_i = XW_Q^i,\quad K_i = XW_K^i,\quad V_i = XW_V^i WQi,WKi,WVi∈R4×2W_Q^i,W_K^i,W_V^i \in \mathbb{R}^{4 \times 2}

因此每个头得到:

Qi,Ki,Vi∈R3×2Q_i,K_i,V_i \in \mathbb{R}^{3 \times 2}

实现上常常先用更大的矩阵一次投影,再按头切开:

Q=XWQ∈R3×4→Q1,Q2∈R3×2Q = XW_Q \in \mathbb{R}^{3 \times 4} \quad \rightarrow \quad Q_1,Q_2 \in \mathbb{R}^{3 \times 2}

每个头分别计算自注意力:

head⁡i=softmax⁡(QiKiTdk+M)Vi\operatorname{head}_i = \operatorname{softmax} \left( \frac{Q_iK_i^T}{\sqrt{d_k}}+M \right) V_i

两个头并行算出:

O1,O2∈R3×2O_1,O_2 \in \mathbb{R}^{3 \times 2}

多个头的输出沿特征维拼接:

Concat⁡(O1,O2)=[O1∣O2]∈R3×4\operatorname{Concat}(O_1,O_2) = \begin{bmatrix} O_1 & | & O_2 \end{bmatrix} \in \mathbb{R}^{3 \times 4}

再乘输出投影,把多个子空间混合成全局信息:

O=Concat⁡(O1,O2) WO,WO∈R4×4O = \operatorname{Concat}(O_1,O_2)\,W_O, \quad W_O \in \mathbb{R}^{4 \times 4}

多头注意力

完整写法:

MHA⁡(X)=Concat⁡(head⁡1,…,head⁡h) WO\operatorname{MHA}(X) = \operatorname{Concat}(\operatorname{head}_1,\ldots,\operatorname{head}_h)\,W_O

4. MHA、MQA、GQA

为了降低复杂度,工程上多头还可以让若干个 Query 头共享同一组 K/V:

名称Query 头Key/Value 组典型用途
MHAhhhh原始 Transformer
GQAhhgg,且 1<g<h1<g<h多数现代 LLM
MQAhh11进一步压缩 KV Cache