从“一个注意力”到“多个视角”:Multi-Head Attention 论文精读

论文:Attention Is All You Need
章节:3.2.2 Multi-Head Attention
核心问题:既然已经有了 Scaled Dot-Product Attention,为什么还要把它拆成多个并行的 attention heads?

上一小节中,我们已经理解了 Scaled Dot-Product Attention 的基本机制:

Attention⁡(Q,K,V)\operatorname{Attention}(Q,K,V)Attention(Q,K,V)

softmax⁡(QK⊤dk)V \operatorname{softmax} \left( \frac{QK^\top}{\sqrt{d_k}} \right)V softmax(dk QK)V

它先通过 Query 和 Key 计算不同位置之间的匹配分数,再经过 softmax 得到注意力权重,最后使用这些权重对 Value 进行加权汇总。

从功能上看,这似乎已经足够完整:

  • Query 决定“我想找什么”;
  • Key 决定“我可以怎样被匹配”;
  • Value 决定“匹配之后实际取走什么信息”。

那么,一个自然的问题就出现了:

既然单个注意力函数已经可以让每个 token 关注序列中的其他位置,为什么 Transformer 不直接使用一个完整的、大维度注意力,而要设计多个并行的“头”?

这正是 3.2.2 节要回答的问题。

Multi-Head Attention 的价值,并不只是“同时计算了很多次 Attention”。它真正解决的是一个更深层的问题:

单个注意力头只能使用一套表示方式、一套注意力分布和一次加权汇总,同时容纳所有需要表达的关系。

多头注意力则把这一次高度耦合的信息汇总,拆成多次低维、独立、并行的汇总。不同的头可以在不同表示子空间中形成不同的关注模式,最后再将这些结果统一整合。

为了避免一开始就陷入大量符号和矩阵形状,我们先建立整体流程,再回到论文原文,逐句理解作者为什么这样设计。


本节的阅读地图

3.2.2 节整体上采用了一个非常清晰的论证结构:

做法→动机→公式→具体配置→成本论证 \boxed{ \text{做法} \rightarrow \text{动机} \rightarrow \text{公式} \rightarrow \text{具体配置} \rightarrow \text{成本论证} } 做法动机公式具体配置成本论证

作者首先说明 Multi-Head Attention 是怎样计算的:

  1. 对 Query、Key、Value 分别进行多组不同的线性投影;
  2. 在每组投影结果上并行执行 Attention;
  3. 将多个 head 的输出拼接;
  4. 再通过输出矩阵投影回模型主干维度。

接着,作者解释为什么这样做有价值:

多个注意力头允许模型同时关注不同位置、不同表示子空间中的信息。

随后,作者用公式把整个过程精确定义下来,并给出本文所采用的具体配置:

dmodel=512,h=8,dk=dv=64 d_{model}=512,\qquad h=8,\qquad d_k=d_v=64 dmodel=512,h=8,dk=dv=64
最后,作者补上一个非常重要的成本论证:

虽然使用了 8 个注意力头,但因为每个头的维度缩小到原来的 (1/8),所以总体计算成本仍然与一个完整维度的单头注意力相近。

阅读这一节时,建议重点抓住两个句子:

Multi-head attention allows the model to jointly attend to information from different representation subspaces at different positions.

以及:

With a single attention head, averaging inhibits this.

前一句说明多头注意力获得了什么能力,后一句说明单头注意力为什么难以获得这种能力。

它们构成了整节最核心的正反论证。


先看完整过程:Multi-Head Attention 到底在做什么

在逐句阅读原文之前,我们先用一张图建立完整的计算地图。

在这里插入图片描述

这张图展示了一个 Multi-Head Self-Attention 从输入 (X) 到最终输出的完整过程。

为了便于理解,我们使用一句简单的话作为例子:

老白拿起他的锅铲

可以将它划分为四个 token:

老白,拿起,他的,锅铲 \text{老白},\text{拿起},\text{他的},\text{锅铲} 老白,拿起,他的,锅铲

此时序列长度为:

n=4n=4n=4

在原论文的基础配置中,每个 token 使用一个 512 维向量表示,因此输入矩阵 (X) 的形状为:

X∈Rn×512 X\in\mathbb{R}^{n\times512} XRn×512

在这个例子中:

X∈R4×512 X\in\mathbb{R}^{4\times512} XR4×512

这里的两个维度分别表示:

  • (n):序列中的 token 数量;
  • (512):每个 token 的特征维度,也就是 (d_{model})。

整个 Multi-Head Self-Attention 可以概括为六个步骤:

输入→多组投影→计算注意力分数→汇总 Value→拼接多个头→输出投影 \boxed{ \text{输入} \rightarrow \text{多组投影} \rightarrow \text{计算注意力分数} \rightarrow \text{汇总 Value} \rightarrow \text{拼接多个头} \rightarrow \text{输出投影} } 输入多组投影计算注意力分数汇总 Value拼接多个头输出投影

也可以进一步压缩为三个字:

分—算—合 \boxed{\text{分—算—合}}

但这里的“分”,并不是把原始向量直接切割成几段,而是通过不同的可学习矩阵,将同一个输入投影到多个不同的低维表示空间中。


第一步:输入 (X)

设输入序列中有 (n) 个 token,每个 token 的表示维度为 (dmodel=512)(d_{model}=512)(dmodel=512),则:

X∈Rn×512 X\in\mathbb{R}^{n\times512} XRn×512

对于 Self-Attention,Query、Key、Value 都来自同一个输入:

Q=K=V=X Q=K=V=X Q=K=V=X

需要注意,这里的:

Q=K=V=X Q=K=V=X Q=K=V=X

表示它们的原始来源相同,并不意味着最终进入每个 attention head 的 (Qi)(Q_i)(Qi)(Ki)(K_i)(Ki)(Vi)(V_i)(Vi)完全一样。

在进入每个头之前,(X)(X)(X)还要分别乘上不同的投影矩阵。


第二步:每个头进行独立的线性投影

Transformer 使用 (h=8) 个 attention heads。

对于第 (i) 个头,分别有三组可学习矩阵:

WiQ∈R512×64 W_i^Q\in\mathbb{R}^{512\times64} WiQR512×64

WiK∈R512×64 W_i^K\in\mathbb{R}^{512\times64} WiKR512×64

WiV∈R512×64 W_i^V\in\mathbb{R}^{512\times64} WiVR512×64

于是,第 (i) 个头的 Query、Key、Value 分别为:

Qi=XWiQ Q_i=XW_i^Q Qi=XWiQ

Ki=XWiK K_i=XW_i^K Ki=XWiK

Vi=XWiV V_i=XW_i^V Vi=XWiV

对应的张量形状为:

(n,512)×(512,64)=(n,64) (n,512)\times(512,64)=(n,64) (n,512)×(512,64)=(n,64)

因此:

Qi,Ki,Vi∈Rn×64 Q_i,K_i,V_i\in\mathbb{R}^{n\times64} Qi,Ki,ViRn×64

这一步是理解 Multi-Head Attention 的第一个关键点:

Multi-Head Attention 不是把 (X) 的 512 个维度直接切成 8 个固定的 64 维片段。

更准确地说,每个头都从完整的 512 维输入出发,通过自己独立的投影矩阵,构造一个新的 64 维表示。

也就是说,第一个头看到的是:

XW1Q,XW1K,XW1V XW_1^Q,\quad XW_1^K,\quad XW_1^V XW1Q,XW1K,XW1V

第二个头看到的是:

XW2Q,XW2K,XW2V XW_2^Q,\quad XW_2^K,\quad XW_2^V XW2Q,XW2K,XW2V

一直到第八个头。

这些投影矩阵在初始化时通常并没有明确的语义分工,但在训练过程中,不同的头可能逐渐对不同的信息模式更加敏感。

例如,一些头可能更容易表示:

  • 动作和对象之间的关系;
  • 代词和其指代对象之间的关系;
  • 相邻 token 之间的局部关系;
  • 长距离位置之间的依赖;
  • 物品和所属者之间的关系。

需要强调的是,这种分工不是人为预先指定的,也不意味着每个头都一定对应一个清晰、唯一的语言学功能。

“不同表示子空间”的本质,首先是:

不同的投影矩阵对原始特征进行了不同的重组和重新编码。

至于这些空间最终是否对应句法、指代、位置或语义关系,是模型在训练过程中可能形成的功能表现。


第三步:每个头独立计算注意力分数

得到:

Qi,Ki,Vi∈Rn×64 Q_i,K_i,V_i\in\mathbb{R}^{n\times64} Qi,Ki,ViRn×64

之后,第(i)(i)(i)个头先计算:

QiKi⊤ Q_iK_i^\top QiKi

其中:

Qi∈Rn×64 Q_i\in\mathbb{R}^{n\times64} QiRn×64

而:

Ki⊤∈R64×n K_i^\top\in\mathbb{R}^{64\times n} KiR64×n

因此:

QiKi⊤∈Rn×n Q_iK_i^\top \in \mathbb{R}^{n\times n} QiKiRn×n

也就是:

(n,64)×(64,n)=(n,n) (n,64)\times(64,n)=(n,n) (n,64)×(64,n)=(n,n)

这个 ((n,n)) 的矩阵,就是注意力分数矩阵。

矩阵中的每一行对应一个 Query token,每一列对应一个 Key token。

第 ((a,b)) 个元素表示:

第 (a) 个 token 在更新自己的表示时,与第 (b) 个 token 的匹配程度。

仍以“老白拿起他的锅铲”为例。

如果某个头学到了物品与所属者之间的关系,那么当“锅铲”作为 Query 时,它可能会对“老白”这个 Key 产生较高的匹配分数。

此时,注意力矩阵中的对应位置:

Query=锅铲,Key=老白 \text{Query}=\text{锅铲}, \qquad \text{Key}=\text{老白} Query=锅铲,Key=老白

就会具有较高的数值。

计算出点积之后,还要除以:

dk \sqrt{d_k} dk

在本文中:

dk=64 d_k=64 dk=64

所以:

dk=64=8 \sqrt{d_k}=\sqrt{64}=8 dk =64 =8

于是缩放后的注意力分数为:

Si=QiKi⊤64 S_i= \frac{Q_iK_i^\top}{\sqrt{64}} Si=64 QiKi

其形状仍然是:

Si∈Rn×n S_i\in\mathbb{R}^{n\times n} SiRn×n

缩放的作用,是防止当 (d_k) 较大时,点积结果的绝对值过大,使 softmax 进入梯度非常小的饱和区域。


第四步:softmax 后汇总 Value

得到注意力分数矩阵后,沿矩阵的每一行执行 softmax:

Ai=softmax⁡(Si) A_i= \operatorname{softmax}(S_i) Ai=softmax(Si)

也就是:

Ai=softmax⁡(QiKi⊤64) A_i= \operatorname{softmax} \left( \frac{Q_iK_i^\top}{\sqrt{64}} \right) Ai=softmax(64 QiKi)

其中:

Ai∈Rn×n A_i\in\mathbb{R}^{n\times n} AiRn×n

softmax 之后,每一行的元素之和等于 1。

因此,每一行都可以理解为:

当前 Query token 应该以怎样的比例,从所有 Value token 中读取信息。

随后,使用这个注意力权重矩阵对 (V_i) 进行加权汇总:

head⁡i=AiVi \operatorname{head}_i=A_iV_i headi=AiVi

维度变化为:

(n,n)×(n,64)=(n,64) (n,n)\times(n,64)=(n,64) (n,n)×(n,64)=(n,64)

所以:

head⁡i∈Rn×64 \operatorname{head}_i \in \mathbb{R}^{n\times64} headiRn×64

完整公式为:

head⁡i=AiVi=softmax⁡(QiKi⊤64)Vi\operatorname{head}_i=A_iV_i=\operatorname{softmax} \left( \frac{Q_iK_i^\top}{\sqrt{64}} \right)V_i headi=AiVi=softmax(64 QiKi)Vi

这里必须严格区分两个形状:

Ai∈Rn×n A_i\in\mathbb{R}^{n\times n} AiRn×n

和:

head⁡i∈Rn×64 \operatorname{head}_i \in \mathbb{R}^{n\times64} headiRn×64

前者是注意力权重矩阵,它回答:

每个 token 应该看谁、看多少?

后者是一个 head 汇总信息之后得到的新表示,它回答:

每个 token 看完其他位置以后,得到了什么信息?

注意力矩阵不是 head 的最终输出。它只是控制信息如何从 (V_i) 中被读取和组合。


第五步:拼接 8 个 attention heads

每个 attention head 都会产生一个:

(n,64) (n,64) (n,64)

形状的输出。

因此,8 个头分别得到:

head⁡1,head⁡2,…,head⁡8 \operatorname{head}_1, \operatorname{head}_2, \dots, \operatorname{head}_8 head1,head2,,head8

并且:

head⁡i∈Rn×64 \operatorname{head}_i \in \mathbb{R}^{n\times64} headiRn×64

接下来,将 8 个头沿特征维度进行拼接:

Concat⁡(head⁡1,…,head⁡8) \operatorname{Concat} ( \operatorname{head}_1, \dots, \operatorname{head}_8 ) Concat(head1,,head8)

拼接后的形状为:

(n,8×64)=(n,512) (n,8\times64)=(n,512) (n,8×64)=(n,512)

这里改变的是特征维度,而不是 token 数量。

从输入到此处,序列中始终有 (n) 个 token:

n→n→n n\rightarrow n\rightarrow n nnn

变化的是每个 token 的表示方式:

512→8×64→512 512 \rightarrow 8\times64 \rightarrow 512 5128×64512

换句话说,Multi-Head Attention 不会增加或减少 token。它只是让每个 token 的表示,在多个注意力头中分别吸收不同位置的信息。


第六步:通过 (W^O) 进行输出投影

8 个头拼接之后,得到:

H=Concat⁡(head⁡1,…,head⁡8) H= \operatorname{Concat} ( \operatorname{head}_1, \dots, \operatorname{head}_8 ) H=Concat(head1,,head8)

其中:

H∈Rn×512 H\in\mathbb{R}^{n\times512} HRn×512

随后,乘上输出投影矩阵:

WO∈R512×512 W^O\in\mathbb{R}^{512\times512} WOR512×512

于是:

MultiHead⁡(Q,K,V)=HWO \operatorname{MultiHead}(Q,K,V)= HW^O MultiHead(Q,K,V)=HWO

维度变化为:

(n,512)×(512,512)=(n,512) (n,512)\times(512,512)= (n,512) (n,512)×(512,512)=(n,512)

因此,最终输出仍然是:

MultiHead⁡(Q,K,V)∈Rn×512 \operatorname{MultiHead}(Q,K,V) \in \mathbb{R}^{n\times512} MultiHead(Q,K,V)Rn×512

输入和输出的形状完全一致:

(n,512)→(n,512) (n,512)\rightarrow(n,512) (n,512)(n,512)

但二者的含义已经发生了变化。

输入中的每个 token 只有进入注意力层之前的表示;输出中的每个 token,则已经根据多个注意力头学到的不同匹配方式,从其他位置读取并融合了上下文信息。

以“老白拿起他的锅铲”为例,经过 Multi-Head Self-Attention 之后:

  • “老白”的表示可能融合了“拿起”和“锅铲”的信息;
  • “拿起”的表示可能融合了动作执行者和动作对象的信息;
  • “他的”的表示可能融合了指代对象“老白”的信息;
  • “锅铲”的表示可能融合了所属者“老白”和动作“拿起”的信息。

因此,token 本身没有改变,但它所携带的上下文信息变得更加丰富。
在这里插入图片描述


压缩成一句话:

每个头先将完整输入投影到自己的 64 维表示空间,再使用 Q 和 K 决定“看谁、看多少”,按照注意力权重汇总 V;8 个头分别看完以后,将结果拼接,再映射回 512 维。

现在,我们已经知道 Multi-Head Attention 在计算上做了什么。

接下来再回到论文原文,真正需要理解的问题就变成:

作者为什么认为这种设计比一个完整维度的单头 Attention 更好?


逐句拆解

第一句:从最朴素的方案切入

Instead of performing a single attention function with (d_{model})-dimensional keys, values and queries, we found it beneficial to linearly project the queries, keys and values (h) times with different, learned linear projections to (d_k), (d_k) and (d_v) dimensions, respectively.

这句话一开头就使用了一个非常典型的学术论证结构:

Instead of A, we found it beneficial to B.

即:

与其采用 A,我们发现采用 B 更有益。

作者没有直接说“我们设计了 Multi-Head Attention”,而是先把一个最朴素的方案摆出来:

方案 A:单个满维度 Attention

直接让:Q,K,VQ,K,VQ,K,V都保持:dmodeld_{model}dmodel维,然后执行一次完整的 Attention。在本文中:dmodel=512d_{model}=512dmodel=512

因此,最直观的做法就是让一个 attention head 直接处理 512 维的 Query、Key 和 Value。

随后,作者提出方案 B:

方案 B:多组低维投影

将 Query、Key、Value 分别进行 (h) 次不同的线性投影。

每次投影都使用不同的、可学习的矩阵,把它们映射到:dk,dk,dvd_k,\quad d_k,\quad d_vdk,dk,dv维。

这里有三个词非常关键。

different:不同的

每个头的投影矩阵不同。

第一个头拥有:

W1Q,W1K,W1V W_1^Q,W_1^K,W_1^V W1Q,W1K,W1V

第二个头拥有:

W2Q,W2K,W2V W_2^Q,W_2^K,W_2^V W2Q,W2K,W2V

一直到第 (h) 个头。

因此,不同头对同一个输入进行的是不同的特征重组,而不是重复计算完全相同的表示。

learned:通过训练学习的

这些矩阵不是研究者人工设定的,也不是预先固定的。

它们是模型参数,会随着训练过程不断更新。

因此,模型可以根据任务目标,自己学习:

  • 哪些输入特征应该组合在一起;
  • 哪些维度有利于计算 Query;
  • 哪些维度有利于计算 Key;
  • 哪些信息适合作为 Value 被传递。

respectively:分别地

原文写的是:

to (d_k), (d_k) and (d_v) dimensions, respectively

其中:

  • Query 被投影到 (d_k) 维;
  • Key 被投影到 (d_k) 维;
  • Value 被投影到 (d_v) 维。

Query 和 Key 必须具有相同维度,因为它们需要执行点积:QiKi⊤Q_iK_i^\topQiKi

而 Value 不参与这个点积,因此理论上 (d_v) 不必等于 (d_k)。

在原论文的具体配置中,两者都取 64:

dk=dv=64d_k=d_v=64dk=dv=64


第二句:在多个头上并行执行 Attention

On each of these projected versions of queries, keys and values we then perform the attention function in parallel, yielding (d_v)-dimensional output values.

经过投影之后,每个头都获得一组自己的:Qi,Ki,ViQ_i,K_i,V_iQi,Ki,Vi

随后,在每一组投影结果上独立执行一次 Attention:

head⁡i=Attention⁡(Qi,Ki,Vi)\operatorname{head}_i= \operatorname{Attention} ( Q_i,K_i,V_i ) headi=Attention(Qi,Ki,Vi)

原文特别强调:

in parallel

即“并行地”。

多个 attention heads 之间不存在必须按照顺序执行的依赖关系。

第一个头不需要等待第二个头,第二个头也不依赖第三个头的结果。它们可以同时完成各自的矩阵运算。

这延续了 Transformer 贯穿全文的核心思想:

尽可能减少顺序依赖,让大规模矩阵运算可以并行执行。

每个头最终输出一个 (d_v) 维表示。

在本文中:dv=64d_v=64dv=64

因此,每个 token 在每个头中都会得到一个新的 64 维表示。


第三句:拼接后再次投影

These are concatenated and once again projected, resulting in the final values, as depicted in Figure 2.

多个头的输出会先被拼接:

Concat⁡(head⁡1,…,head⁡h) \operatorname{Concat} ( \operatorname{head}_1, \dots, \operatorname{head}_h ) Concat(head1,,headh)

随后,再通过一次线性投影得到最终结果。

这里的:

once again projected

对应的就是:WOW^OWO

因此,整个过程形成一个非常清晰的三段结构:

投影→并行 Attention→拼接并再次投影 \boxed{ \text{投影} \rightarrow \text{并行 Attention} \rightarrow \text{拼接并再次投影} } 投影并行 Attention拼接并再次投影

也就是前面概括的:

分—算—合\boxed{\text{分—算—合}}

对同一个输入使用不同矩阵进行投影:QWiQ,KWiK,VWiVQW_i^Q,\quad KW_i^K,\quad VW_i^VQWiQ,KWiK,VWiV

每个头独立运行 Attention:Attention⁡(QWiQ,KWiK,VWiV)\operatorname{Attention} ( QW_i^Q, KW_i^K, VW_i^V )Attention(QWiQ,KWiK,VWiV)

将多个头的输出拼接,并通过 (W^O) 进行统一整合:Concat⁡(head⁡1,…,head⁡h)WO\operatorname{Concat} ( \operatorname{head}_1,\dots,\operatorname{head}_h ) W^OConcat(head1,,headh)WO

到这里,作者已经讲清了 Multi-Head Attention 是怎样计算的。

接下来,他开始回答更重要的问题:

为什么这样做比一个 attention head 更好?


第四句:Multi-Head Attention 的核心动机

Multi-head attention allows the model to jointly attend to information from different representation subspaces at different positions.

这是整个 3.2.2 节最重要的一句话。

它说明 Multi-Head Attention 带来的并不只是“更多的注意力矩阵”,而是两个维度上的多样性:

  1. 不同的表示子空间;
  2. 不同的序列位置。

什么是不同的表示子空间

原始输入中的每个 token 是一个 512 维向量。

一个 attention head 并不是直接使用这 512 个原始维度,而是通过投影矩阵将其映射为新的 64 维表示:XWiQ,XWiK,XWiVXW_i^Q,\quad XW_i^K,\quad XW_i^VXWiQ,XWiK,XWiV

不同 head 使用不同的投影矩阵,因此它们得到的 64 维表示也不同。

可以将其理解为:

每个头都使用一套不同的坐标系统重新描述同一个 token。

例如,假设原始向量中包含许多混合在一起的信息:

  • 词义信息;
  • 位置信息;
  • 动作信息;
  • 指代信息;
  • 实体信息;
  • 上下文信息。

某个头的投影可能组合出更适合表示动作关系的特征,另一个头的投影可能形成更适合表示指代关系的特征。

这就是:

different representation subspaces

它不是说模型提前知道了哪些维度属于句法、哪些维度属于语义,而是不同的投影矩阵,为不同头提供了不同的信息组织方式。


什么是在不同位置上进行关注

除了表示空间不同,不同头还可以形成不同的注意力分布。

例如,对于 token“锅铲”:

一个头可能重点关注:

"老白"\text{"老白"}"老白"

用来建模物品与所属者之间的关系。

另一个头可能重点关注:

"拿起"\text{"拿起"}"拿起"

用来建模动作和动作对象之间的关系。

还有一个头可能关注:

"他的"\text{"他的"}"他的"

用来建模指代或所属关系。

因此,多个头可以在同一时间:

  • 使用不同的特征表示;
  • 关注不同的位置;
  • 汇总不同类型的信息。

原文中的:

jointly attend

并不是简单地说“看很多地方”,而是强调:

多种关注模式能够同时存在,并在最终输出中被共同整合。


第五句:为什么单头会限制这种能力

With a single attention head, averaging inhibits this.

这句话很短,却是整节最有力量的一句。

作者没有展开解释,但它指出了单头 Attention 的一个根本限制:

单个注意力头最终只能进行一次加权汇总。

对于某个 Query,单头 Attention 会生成一行注意力权重:

α1,α2,…,αn\alpha_1,\alpha_2,\dots,\alpha_nα1,α2,,αn

随后,对所有 Value 做加权求和:

α1V1+α2V2+⋯+αnVn\alpha_1V_1+ \alpha_2V_2+ \dots+ \alpha_nV_nα1V1+α2V2++αnVn

这本质上是一种加权平均。

需要特别注意:

单头注意力不是只能关注一个位置。

softmax 完全可以让多个位置同时具有较高权重。

真正的限制在于:

单头只能使用一套 Query/Key/Value 表示、一套注意力分布,以及一次 Value 加权和,去同时容纳所有关系。

例如,“锅铲”可能同时需要:

  • 关注“老白”,表达所属者关系;
  • 关注“拿起”,表达动作对象关系;
  • 关注“他的”,表达指代和所属关系。

单头可以给这三个位置都分配一定的权重,但它们最终必须进入同一个加权和:

α老白V老白+α拿起V拿起+α他的V他的 \alpha_{\text{老白}}V_{\text{老白}} + \alpha_{\text{拿起}}V_{\text{拿起}} + \alpha_{\text{他的}}V_{\text{他的}} α老白V老白+α拿起V拿起+α他的V他的

这意味着,不同性质的关系会在同一个表示空间中被混合。

因此,原文中的:

averaging inhibits this

并不是简单地说“平均会让信息变模糊”。

更准确的含义是:

一次加权平均迫使多种关系共享同一套表示方式和同一个输出通道,因此限制了模型同时、独立地保留多种关注模式的能力。

Multi-Head Attention 的解决方式是:

  • 开启多条独立的表示通道;
  • 每个通道使用不同的投影;
  • 每个通道形成自己的注意力分布;
  • 每个通道独立进行一次 Value 汇总;
  • 最后再将多个结果拼接并整合。

因此,多头注意力不是把一次平均简单重复多次,而是让多次平均发生在不同的表示空间中。

这可以概括为:

单头是把多种关系先混在一起,再输出一个结果;多头是让多种关系先分别成立,再统一整合。

这才是“averaging inhibits this”背后的深层含义。


公式:把“分—算—合”精确定义下来

论文给出:

MultiHead⁡(Q,K,V)=Concat⁡(head⁡1,…,head⁡h)WO \operatorname{MultiHead}(Q,K,V)=\operatorname{Concat} ( \operatorname{head}_1, \dots, \operatorname{head}_h ) W^O MultiHead(Q,K,V)=Concat(head1,,headh)WO

其中:

head⁡i=Attention⁡(QWiQ,KWiK,VWiV) \operatorname{head}_i=\operatorname{Attention} ( QW_i^Q, KW_i^K, VW_i^V ) headi=Attention(QWiQ,KWiK,VWiV)

第二个公式描述的是第 (i) 个头。

先执行三个投影:QWiQQW_i^QQWiQKWiKKW_i^KKWiKVWiVVW_i^VVWiV

然后,将投影结果送入 Scaled Dot-Product Attention:

Attention⁡(QWiQ,KWiK,VWiV)\operatorname{Attention} ( QW_i^Q, KW_i^K, VW_i^V )Attention(QWiQ,KWiK,VWiV)

得到:head⁡i\operatorname{head}_iheadi

第一个公式描述多个头如何合并。

先拼接:

Concat⁡(head⁡1,…,head⁡h) \operatorname{Concat} ( \operatorname{head}_1,\dots,\operatorname{head}_h ) Concat(head1,,headh)

再乘输出矩阵:

WO W^O WO

最终得到 Multi-Head Attention 的输出。

从文字到公式,两者是严格对应的:

文字动作 公式符号
对 Query 进行投影 (QWiQ)(QW_i^Q)(QWiQ)
对 Key 进行投影 (KWiK)(KW_i^K)(KWiK)
对 Value 进行投影 (VWiV)(VW_i^V)(VWiV)
在每个头上执行 Attention (head⁡i)(\operatorname{head}_i)(headi)
拼接多个头 (Concat⁡)(\operatorname{Concat})(Concat)
再次进行输出投影 (WO)(W^O)(WO)

这体现了论文非常典型的表达方式:

先用自然语言交代机制,再用数学公式消除歧义。

阅读这类论文时,一个很有效的方法是:

每看到一个公式符号,就回到前面的自然语言中寻找它对应的动作;每看到一个文字动作,也检查它是否能在公式中找到对应结构。


投影矩阵的形状

论文继续说明:

WiQ∈Rdmodel×dk W_i^Q \in \mathbb{R}^{d_{model}\times d_k} WiQRdmodel×dk

WiK∈Rdmodel×dk W_i^K \in \mathbb{R}^{d_{model}\times d_k} WiKRdmodel×dk

WiV∈Rdmodel×dv W_i^V \in \mathbb{R}^{d_{model}\times d_v} WiVRdmodel×dv

以及:

WO∈Rhdv×dmodel W^O \in \mathbb{R}^{hd_v\times d_{model}} WORhdv×dmodel

这些形状并不是附属信息,而是理解 Multi-Head Attention 是否真正到位的检验。


(WiQ)(W_i^Q)(WiQ) 和(WiK)和 (W_i^K)(WiK)

WiQ,WiK∈Rdmodel×dk W_i^Q,W_i^K \in \mathbb{R}^{d_{model}\times d_k} WiQ,WiKRdmodel×dk

它们把:dmodeld_{model}dmodel维输入映射到:dkd_kdk维。

在本文中:

512→64512\rightarrow6451264

对应:

(n,512)×(512,64)=(n,64)(n,512)\times(512,64)=(n,64)(n,512)×(512,64)=(n,64)


(WiV)(W_i^V)(WiV)

WiV∈Rdmodel×dvW_i^V \in \mathbb{R}^{d_{model}\times d_v}WiVRdmodel×dv

它将 Value 从:dmodeld_{model}dmodel维映射到:dvd_vdv维。

在本文中同样是:

512→64 512\rightarrow64 51264


(W^O)

拼接 (h) 个 head 后,特征总宽度是:hdvhd_vhdv

因此:
WO∈Rhdv×dmodelW^O \in \mathbb{R}^{hd_v\times d_{model}}WORhdv×dmodel

在本文中:

h=8,dv=64 h=8,\qquad d_v=64 h=8,dv=64

所以:

hdv=8×64=512 hd_v=8\times64=512 hdv=8×64=512

于是:

WO∈R512×512 W^O\in\mathbb{R}^{512\times512} WOR512×512

这个矩阵将多个头拼接后的表示重新映射回:dmodel=512d_{model}=512dmodel=512维。

因此,多头注意力实现了:

(n,dmodel)→(n,dmodel) (n,d_{model}) \rightarrow (n,d_{model}) (n,dmodel)(n,dmodel)

也就是:

(n,512)→(n,512) (n,512)\rightarrow(n,512) (n,512)(n,512)

正是因为输入和输出维度一致,Multi-Head Attention 才能被方便地嵌入 Transformer 的层级结构中,并与后续模块连接。


本文采用的具体配置

论文写道:

In this work we employ (h=8) parallel attention layers, or heads. For each of these we use (d_k=d_v=d_{model}/h=64).

本文采用:

h=8 h=8 h=8

也就是 8 个并行 attention heads。

同时:

dmodel=512 d_{model}=512 dmodel=512

因此:

dk=dv=dmodelh=5128=64 d_k=d_v=\frac{d_{model}}{h}=\frac{512}{8}=64 dk=dv=hdmodel=8512=64

也就是说,每个头都在 64 维空间中执行注意力计算。

这里需要避免一个常见但不够准确的说法:

“Transformer 把 512 维向量切成 8 份,每份 64 维。”

这种说法容易让人误以为原始输入 (X) 的特征被机械地切片。

实际过程是:

每个头都使用一个 (512×64512\times64512×64) 的可学习矩阵,从完整的 512 维输入中生成一个新的 64 维表示。

因此,64 维表示是“投影结果”,不是原始特征的固定切片。

不过,从计算预算的角度看,作者确实将原本集中在一个 512 维注意力空间中的计算,重新分配到了 8 个 64 维空间中。

这个维度设计,正是下一句成本论证成立的基础。


为什么 8 个头没有带来 8 倍计算量

论文最后写道:

Due to the reduced dimension of each head, the total computational cost is similar to that of single-head attention with full dimensionality.

这是整个 3.2.2 节一个非常漂亮的收尾。

读者很自然会提出疑问:

原来只需要计算一次 Attention,现在要计算 8 次,计算量难道不会变成 8 倍吗?

答案是不会。

因为每个头并不是在完整的 512 维空间中计算,而只在 64 维空间中计算。


从注意力核心计算量来看

一次 Attention 的主要计算包括两部分:

QK⊤ QK^\top QK

和:

AV AV AV

设序列长度为 (n)(n)(n),特征维度为 (d)(d)(d)

那么:

QK⊤ QK^\top QK

的主要计算量大致与:

n2d n^2d n2d

成正比。

(AV) 的计算量也大致与:

n2d n^2d n2d

成正比。

如果使用一个完整维度的单头 Attention,且:

d=dmodel=512 d=d_{model}=512 d=dmodel=512

则其核心计算量大致为:

O(n2dmodel) O(n^2d_{model}) O(n2dmodel)

对于 Multi-Head Attention,每个头的维度为:

dk=dv=dmodelh d_k=d_v=\frac{d_{model}}{h} dk=dv=hdmodel

单个头的计算量大致为:

O(n2dmodelh) O \left( n^2\frac{d_{model}}{h} \right) O(n2hdmodel)

共有 (h)(h)(h) 个头,因此总计算量为:

h⋅O(n2dmodelh) h \cdot O \left( n^2\frac{d_{model}}{h} \right) hO(n2hdmodel)

其中 (h) 可以约掉:

O(n2dmodel) O(n^2d_{model}) O(n2dmodel)

也就是说:

头的数量增加了 (h) 倍,但每个头的维度缩小到了原来的 (1/h),二者在主要注意力计算中相互抵消。


从直觉上理解

单头满维 Attention 相当于:

1×512 1\times512 1×512

多头 Attention 相当于:

8×64 8\times64 8×64

而:

8×64=512 8\times64=512 8×64=512

这个等式可以帮助我们建立直觉,但它本身不是完整的计算复杂度证明。

更准确地说:

Multi-Head Attention 并不是把一份完整的 512 维计算复制 8 次,而是把原本集中在一个高维空间中的计算预算,重新分配给 8 个低维注意力空间。

因此,模型在近似相同的计算预算下,获得了:

  • 多组不同的投影;
  • 多套独立的注意力分布;
  • 多种表示子空间;
  • 多路独立的信息汇总结果。

这可以概括为:

用近似相同的计算成本,换取更多并行的表示视角。

这也是 Multi-Head Attention 最有说服力的设计之处。

它不是简单地让模型“算得更多”,而是让同一份计算预算被组织得更加有效。


更深一层:Multi-Head Attention 究竟增加了什么

理解到这里,我们可以进一步追问:

多头注意力究竟增加了什么?

它增加的不是 token 数量,也不是模型主干维度。

输入是:

(n,512) (n,512) (n,512)

输出仍然是:

(n,512) (n,512) (n,512)

它真正增加的是:

信息被组织和汇总的路径数量。

单头注意力只有:

  • 一组 (WQ,WK,WV)(W^Q,W^K,W^V)(WQ,WK,WV)
  • 一个表示空间;
  • 一张注意力矩阵;
  • 一次 Value 加权汇总。

多头注意力则拥有:

  • 多组 (WiQ,WiK,WiV)(W_i^Q,W_i^K,W_i^V)(WiQ,WiK,WiV)
  • 多个不同的表示子空间;
  • 多张独立的注意力矩阵;
  • 多次彼此独立的 Value 汇总。

因此,Multi-Head Attention 的核心优势不是“一个词可以看多个词”。

单头 Attention 本来就可以给多个位置非零权重。

真正的区别是:

多头注意力允许模型使用多套不同的匹配标准,在多个不同的表示空间中,分别完成信息检索和信息汇总。

可以把单头理解为一位使用一套评价标准的专家。

他当然可以同时参考多个证据,但所有证据最终都必须按照同一套判断体系进行整合。

多头则像多位专家:

  • 每位专家拥有不同的专业视角;
  • 每位专家关注不同的证据;
  • 每位专家独立形成自己的中间结论;
  • 最后再将多个结论统一整合。

这个类比的重点并不在“专家数量更多”,而在于:

判断标准和信息通道彼此独立。


这一节真正的设计思想

表面上看,Multi-Head Attention 是一种模型结构。

更深一层看,它体现了一种非常普遍的算法设计思想:

当一次统一计算难以同时容纳多种模式时,可以先将问题投影到多个较小的子空间中独立处理,再把结果组合回来。

这种思路并不局限于 Attention。

它包含三个普遍动作:

1. 解耦

将一个高度耦合的问题拆分成多个相对独立的处理路径。

2. 专门化

让不同路径通过不同参数,学习不同的信息组织方式。

3. 再融合

避免各路径完全孤立,通过统一的融合层整合信息。

Multi-Head Attention 之所以有效,不只是因为“头多”,而是因为它同时完成了:

解耦+专门化+融合 \boxed{ \text{解耦} + \text{专门化} + \text{融合} } 解耦+专门化+融合

如果只有多个头,却没有不同的投影矩阵,那么多个头可能只是重复计算。

如果只有不同投影,却不独立运行 Attention,那么不同关注模式仍然可能在计算过程中过早混合。

如果只有多个独立输出,却没有拼接和 (W^O),那么这些局部结果又无法被统一组合。

因此,Multi-Head Attention 的完整性来自三部分共同作用:

不同投影→独立注意力→统一融合 \boxed{ \text{不同投影} \rightarrow \text{独立注意力} \rightarrow \text{统一融合} } 不同投影独立注意力统一融合

缺少其中任何一个环节,都不能完整实现原论文所描述的能力。


今日知识点

英文表达

Instead of A, we found it beneficial to B

意思是:

与其采用 A,我们发现采用 B 更有益。

这是科研论文中非常常见的对照句式。

它的价值在于,作者不是孤立地提出自己的方法,而是先建立一个最自然的基线方案,再说明自己的设计相对它改进了什么。


different, learned linear projections

意思是:

不同的、通过训练学习得到的线性投影。

其中:

  • different 强调各个头的矩阵不同;
  • learned 强调矩阵参数由训练获得,而非人工设定;
  • linear projections 在深度学习中通常指通过矩阵乘法完成线性映射。

respectively

意思是:

分别地。

它用于将前后两组并列元素一一对应。

例如:

project queries, keys and values to (d_k), (d_k) and (d_v) dimensions, respectively

表示:

  • Query 对应 (d_k);
  • Key 对应 (d_k);
  • Value 对应 (d_v)。

in parallel

意思是:

并行地。

强调多个操作之间不存在必须按照顺序执行的依赖,可以同时计算。


jointly attend to

意思是:

联合地、同时地关注。

这里强调的不是简单关注多个位置,而是多个不同关注模式可以同时存在,并共同参与最终表示的形成。


representation subspace

意思是:

表示子空间。

它指的是原始表示经过某个投影矩阵后形成的新特征空间。

不同 attention heads 使用不同的投影矩阵,因此会在不同的表示子空间中处理信息。


averaging inhibits this

意思是:

加权平均会抑制这种能力。

这里的“抑制”不是说平均操作本身错误,而是说单头的一次加权汇总,会迫使多种关系共享同一表示空间和同一个输出通道,从而限制不同关注模式被独立保留。


Due to …, the total computational cost is similar to that of …

意思是:

由于……,总计算成本与……相近。

其中:

that of

用于代替前面已经出现过的名词短语,避免重复。

例如:

the cost is similar to that of single-head attention

其中 that 代替的是 computational cost。


顺手攒几个小词

project

在日常英语中,project 可以表示“项目”。

但在深度学习和线性代数语境中,作为动词时,它通常表示:

投影、映射。

例如:

project (X) to 64 dimensions

表示:

将 (X) 映射到 64 维空间。

通常对应:XWXWXW

这样的矩阵乘法。


projection

名词形式,表示:

投影或投影变换。

例如:

learned linear projection

表示:

可学习的线性投影。


concatenate

意思是:

拼接。

它不是求和,也不是平均,而是沿某个维度把多个向量或矩阵连接起来。

例如,8 个(n,64)(n,64)(n,64)沿特征维度拼接后得到:(n,512)(n,512)(n,512)


dimensionality

意思是:

维数或维度规模。

例如:

full dimensionality

表示:

完整维度。


科研知识

1. 先建立整体流程,再进入公式细节

面对符号密集的技术段落,不要一开始就逐个解释矩阵。

先回答:

  • 输入是什么;
  • 输出是什么;
  • 中间有几个阶段;
  • 每个阶段完成什么功能;
  • 张量形状如何变化。

建立整体地图后,再读局部公式,理解成本会显著降低。


2. 不要把“投影”误解为“直接切片”

Multi-Head Attention 中的每个 head 都从完整输入出发:

XWiXW_iXWi

而不是直接取原始向量的一段固定维度。

这意味着不同头得到的是不同的特征组合,而不是互不重叠的原始维度片段。


3. 核对维度是验证理解的硬指标

例如:

QiKi⊤:(n,64)×(64,n)=(n,n)Q_iK_i^\top: (n,64)\times(64,n)=(n,n)QiKi:(n,64)×(64,n)=(n,n)

以及:

AiVi:(n,n)×(n,64)=(n,64)A_iV_i: (n,n)\times(n,64)=(n,64)AiVi:(n,n)×(n,64)=(n,64)

再例如:

Concat⁡:8×(n,64)→(n,512)\operatorname{Concat}: 8\times(n,64)\rightarrow(n,512)Concat:8×(n,64)(n,512)

最后:

(n,512)×(512,512)=n,512)(n,512)\times(512,512)=n,512)(n,512)×(512,512)=n,512)

只要某一步形状无法衔接,就说明对机制的理解存在问题。


4. 单头的限制不是“只能关注一个位置”

单头可以同时关注多个位置。

它真正的限制是:

只能使用一套表示方式、一张注意力矩阵和一次 Value 汇总,同时处理所有关系。

多头则让不同关系可以先在不同表示空间中独立处理,再统一融合。


5. 多头增加的是信息路径,而不是主干维度

输入和输出都是:

(n,dmodel) (n,d_{model}) (n,dmodel)

Multi-Head Attention 没有扩大最终向量宽度,而是在内部建立了多个独立的信息处理通道。

因此,它增强的是表示结构,而不是简单增加输出维度。


6. “更复杂”不一定意味着“更昂贵”

Multi-Head Attention 看起来从一个头变成了 8 个头,但每个头的维度同时缩小。

总计算量大致满足:

h⋅O(n2dmodelh)=O(n2dmodel) h \cdot O \left( n^2\frac{d_{model}}{h} \right)= O(n^2d_{model}) hO(n2hdmodel)=O(n2dmodel)

因此,它本质上是在重新分配计算预算,而不是简单叠加计算量。

这是一种非常有说服力的设计:

在相近成本下,获得更丰富的表达能力。


7. 好的论文论证会同时回答“怎么做、为什么、贵不贵”

3.2.2 节虽然很短,却完整回答了三个问题:

怎么做

多次投影、并行 Attention、拼接、输出投影。

为什么

让模型能够同时关注不同位置和不同表示子空间中的信息。

贵不贵

每个头的维度缩小,因此总体计算成本与满维单头相近。

这是非常值得学习的科研写作结构:

方法+动机+成本 \boxed{ \text{方法} + \text{动机} + \text{成本} } 方法+动机+成本

仅仅说明“我做了什么”还不够。

真正有说服力的方法描述,还要说明:

  • 它解决了什么限制;
  • 它为什么有效;
  • 它是否带来了不可接受的代价。

一句话带走

Multi-Head Attention 不是把输入向量机械地切成 (h)(h)(h) 份,而是让 (h)(h)(h) 个头分别使用自己的可学习投影,从完整输入中构造不同的低维表示。每个头在自己的表示子空间中独立决定“看谁、看多少”,再分别汇总 Value;多个头的结果最终被拼接并映射回 (dmodel)(d_{model})(dmodel) 维。它解决的不是“单头看不到多个位置”,而是“单头只能用一套表示方式和一次加权汇总容纳所有关系”的表达瓶颈。通过把同一计算预算重新分配给多个低维注意力空间,模型获得了多种并行的观察视角,而总体计算量仍与一个满维单头处于相近水平。

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐