博客 · 关于 · 友链 · English · RSS

计算多头注意力

2026年7月5日

多头注意力(Multi-Head Attention)的计算是 transformer 大模型的基石。然而,其计算却因为矩阵颠三倒四而变得不好直观理解。所以我写了这篇文章,做一下总结。

因为抽象的字母 a,b,ca, b, c 会增加思考成本,所以这里凡是涉及到矩阵维数的时候,我都会用一个独一无二的质数来表示,这样给出的例子更具体,更方便理解。因为是独一无二的质数,因此其本身不一定要被看成是具体的数字,也可以看成只是一种编码。

线性代数复习

当我们说一个 a×ba \times b 矩阵的时候,我们说的是这个矩阵有 aa 行和 bb 列。这个矩阵可以看成由 aa 个行向量组成,也可以看成由 bb 个列向量组成。

一个 a×ba \times b 矩阵 AA 乘以 b×cb \times c 矩阵 BB ,会得到 a×ca \times c 的矩阵 RRRR 的第 xx 行第 yy 列元素 RxyR_{xy}AA 的第 xx 行向量和 BB 的第 yy 列向量的内积。

从输入到输出

首先我们看一下输入输出。这里我们假设输入 3 个词元,每个词元加上位置编码之后都“嵌入”为一个 5 维向量。因此这里的输入是一个 3×53 \times 5 的矩阵,我们记为 XX。每一个词元是一个行向量,记为 X1,X2,X3X_1, X_2, X_3

而输出,需要和输入是大小相同的,也是一个 3×53 \times 5 的矩阵,记为 YYYY 的每一行记为 Y1,Y2,Y3Y_1, Y_2, Y_3

Q、K、V

我们先看 QQKK

假设对于每个词元,QQKK 都是 7 维向量;那么对于 3 个词元,QQKK 就是 3×73 \times 7 的矩阵。为了得到 QQKK,定义参数 WQW_QWKW_K5×75 \times 7 的矩阵。然后做矩阵乘法:

Q=XWQK=XWK\begin{aligned} Q &= X W_Q \\ K &= X W_K \end{aligned}

KK 转置之后和 QQ 相乘,QKTQ K^T 也就是 3×73 \times 7 的矩阵乘以 7×37 \times 3 的矩阵,得到一个 3×33 \times 3 的矩阵。

对这个矩阵缩放,每一个数字都除以 7\sqrt{7},然后加上因果掩码(Causal Mask),上三角变成负无穷。为了方便理解什么是因果掩码,这里举个例子,假如有一个 3×33 \times 3 矩阵:

(123456789)\begin{pmatrix} 1 & 2 & 3 \\ 4 & 5 & 6 \\ 7 & 8 & 9 \end{pmatrix}

那么,加上因果掩码之后就会变成:

(145789)\begin{pmatrix} 1 & -\infty & -\infty \\ 4 & 5 & -\infty \\ 7 & 8 & 9 \end{pmatrix}

然后对每一行进行 softmax\operatorname{softmax} 操作,此时负无穷会变成0,结果将是一个上三角为 0 的 3×33 \times 3 矩阵。我们将其记作 SS:

S=softmax(QKT7)S = \operatorname{softmax}\left(\frac{Q K^T}{\sqrt{7}}\right)

如果你不知道怎么是softmax操作,你可以先暂时把它理解成一种归一化,把 -\infty++\infty 之间的所有数都映射到0到1之间,并且保证该行所有数字加起来等于1。

最后,我们假设对于每个词元,VV 都是 11 维向量。那么 WVW_V 应该是一个 5×115 \times 11 的矩阵。根据:

V=XWVV = X W_V

得到 VV 是一个 3×113 \times 11 的矩阵。

SSVV 相乘,得到:

H=SVH = S V

HH 也是一个 3×113 \times 11 的矩阵。

多头注意力

我们假设有 2 个头。也就是说,上面的 Q,K,VQ, K, V 实际上都有 2 个,WQ,WK,WVW_Q, W_K, W_V 也都相应有两个。也就是:

Q1=XWQ1K1=XWK1V1=XWV1Q2=XWQ2K2=XWK2V2=XWV2\begin{aligned} Q_1 &= X W_{Q1} \\ K_1 &= X W_{K1} \\ V_1 &= X W_{V1} \\ Q_2 &= X W_{Q2} \\ K_2 &= X W_{K2} \\ V_2 &= X W_{V2} \end{aligned}

因此,得到的 HH 也有两个: H1,H2H_1, H_2。这两个 HH 都是 3×113 \times 11 的矩阵。将这两个拼接起来,可以得到一个 3×223 \times 22 的矩阵:

Concat(H1,H2)\operatorname{Concat}(H_1, H_2)

输出

通过多头注意力,对于每一个输入的词元,我们都得到了一个 22 维的向量。而我们的输出需要和输入一样是 5 维。因此我们加上一个 22×522 \times 5 的参数矩阵 WOW_O3×223 \times 22 矩阵乘以 22×522 \times 5 的矩阵,得到 3×53 \times 5 的矩阵 YY

Y=Concat(H1,H2)WOY = \operatorname{Concat}(H_1, H_2) W_O

输入输出关系分析

假如我们把上面算法当成黑箱。把所有参数统称为 WW。也就是说:

W=(WQ1,WK1,WV1,WQ2,WK2,WV2,WO)W = (W_{Q1}, W_{K1}, W_{V1}, W_{Q2}, W_{K2}, W_{V2}, W_O)

然后逐行分析输入参数和输出参数的抽象关系,可以写成:

Y1=f1(X1;W)Y2=f2(X1,X2;W)Y3=f3(X1,X2,X3;W)\begin{aligned} Y_1 &= f_1(X_1; W) \\ Y_2 &= f_2(X_1, X_2; W) \\ Y_3 &= f_3(X_1, X_2, X_3; W) \end{aligned}

其中,函数 f1,f2,f3f_1, f_2, f_3 都是由算法和超参数确定的,不会变动,WW 是可训练的参数。可以看到,输出的第 NN 行,都由输入的前 NN 行所确定。

KV Cache

我们假设在算出 YY 之后,给输入 XX 加一个第4行 X4X_4,而 X1,X2,X3X_1, X_2, X_3 都保持不变。根据前述分析,此时,Y1,Y2,Y3Y_1, Y_2, Y_3 也是不变的,只需要计算 Y4=f4(X1,X2,X3,X4;W)Y_4 = f_4(X_1, X_2, X_3, X_4; W)

可以看到所有前述中间结果中,3 行的矩阵,都会变成 4 行,而前三行则都保持不变。因此,如果能把前三行缓存下来,就只需要计算第四行就可以了。这样可以大大节省计算量。更进一步,我们可以注意到,矩阵 HH 的第四行,甚至和 QQ 的前三行毫无关系,因此事实上矩阵 QQ 的内容用完就可以直接扔掉。这也正是 KV Cache 省钱的原理。

尾声

最后放上一张GPT原始论文里面的架构框图:

GPT架构框图

可以看到,唯一比较复杂的就是本文中提到的多头注意力。其它没有涉及的部分主要是残差连接、layer norm和FFN,不过这几个都没有什么难度,所以不再赘述。


Email: i (at) mistivia (dot) com