Heng's Blog
01 04

注意力机制的共同目标是:根据当前查询为一组候选信息分配权重,再将它们聚合为上下文表示。本文从机器翻译中的 Additive Attention 出发,逐步过渡到 Transformer 使用的 Multi-Head Attention。

本文整理自 hengproject/ML_recall 中的 basics/attention.ipynb。参考了 hwcoder 的手写笔记《Recurrent Models of Visual Attention》精读

Additive Attention#

Bahdanau 等人在 2014 年的论文 Neural Machine Translation by Jointly Learning to Align and Translate 中,将注意力机制用于神经机器翻译。生成每个目标 token 时,解码器会为编码器的全部隐藏状态分配权重,并计算上下文向量。

设解码器上一步的隐藏状态(hidden_state)为 si1\underline{s_{i-1}},编码器的全部隐藏状态为 {h1,,hTx}\{\underline{h_1},\ldots,\underline{h_{T_x}}\}。生成目标 token yiy_i 时,需要判断当前解码状态应该关注源序列中的哪些位置,因此要分别计算 si1\underline{s_{i-1}} 与每个 hj\underline{h_j} 的匹配程度。

记号说明: 下文用单下划线表示 vector,用双下划线表示 matrix;scalar 不加下划线。

1. 计算匹配分数#

eij=vaTtanh(Wasi1+Uahj).e_{ij} = \underline{v_a}^T \tanh\left( \underline{\underline{W_a}}\underline{s_{i-1}} + \underline{\underline{U_a}}\underline{h_j} \right).

Wa\underline{\underline{W_a}}Ua\underline{\underline{U_a}}va\underline{v_a} 都是可训练参数。若 si1Rds\underline{s_{i-1}}\in\mathbb{R}^{d_s}hjRdh\underline{h_j}\in\mathbb{R}^{d_h},并将 attention 隐藏空间的维度记为 dad_a,那么 WaRda×ds\underline{\underline{W_a}}\in\mathbb{R}^{d_a\times d_s}UaRda×dh\underline{\underline{U_a}}\in\mathbb{R}^{d_a\times d_h}vaRda\underline{v_a}\in\mathbb{R}^{d_a}

这个评分函数可视为一个带 tanh 的小型前馈网络,用于估计当前解码状态与每个编码器状态的匹配程度。va\underline{v_a} 的主要作用是将 tanh 输出的 vector 转换成 scalar,作为最终的匹配分数 eije_{ij}

为什么看起来和 Vanilla RNN 这么像?

觉得这个匹配函数像 Vanilla RNN,并不是错觉。它们确实使用了几乎相同的计算模板:两个线性变换相加,再经过 tanh

Vanilla RNN 的典型状态更新是 ht=tanh(Wxxt+Whht1+b)\underline{h_t}=\tanh(\underline{\underline{W_x}}\underline{x_t}+\underline{\underline{W_h}}\underline{h_{t-1}}+\underline{b}),Additive Attention 的匹配函数是 eij=vaTtanh(Wasi1+Uahj)e_{ij}=\underline{v_a}^T\tanh(\underline{\underline{W_a}}\underline{s_{i-1}}+\underline{\underline{U_a}}\underline{h_j})

把两者并排看,对应关系会更明显:

Vanilla RNNAdditive Attention
当前输入 xt\underline{x_t}jj 个编码器状态 hj\underline{h_j}
上一步状态 ht1\underline{h_{t-1}}上一步解码器状态 si1\underline{s_{i-1}}
Wxxt+Whht1\underline{\underline{W_x}}\underline{x_t}+\underline{\underline{W_h}}\underline{h_{t-1}}Uahj+Wasi1\underline{\underline{U_a}}\underline{h_j}+\underline{\underline{W_a}}\underline{s_{i-1}}
输出新的隐藏状态 ht\underline{h_t}再乘 vaT\underline{v_a}^T,输出匹配分数 eije_{ij}

它们看起来相似,是因为两者都在做同一类计算:把两个向量投影到同一个隐藏空间,相加后再经过非线性变换。

在 Additive Attention 中,可以先写出 zij=tanh(Wasi1+Uahj)\underline{z_{ij}}=\tanh(\underline{\underline{W_a}}\underline{s_{i-1}}+\underline{\underline{U_a}}\underline{h_j})。这里的 zij\underline{z_{ij}} 是一个联合表示,表示当前解码需求 si1\underline{s_{i-1}} 与第 jj 个编码器状态 hj\underline{h_j} 组合之后的结果。随后再计算 eij=vaTzije_{ij}=\underline{v_a}^T\underline{z_{ij}},把这个向量压缩成一个标量匹配分数。

因此,它本质上是一个单隐藏层前馈神经网络:[si1,hj]zijeij[\underline{s_{i-1}},\underline{h_j}]\longrightarrow\underline{z_{ij}}\longrightarrow e_{ij}。把两个输入拼接起来后,评分函数也可以写成 eij=vaTtanh([WaUa][si1hj])e_{ij}=\underline{v_a}^T\tanh\left(\begin{bmatrix}\underline{\underline{W_a}}&\underline{\underline{U_a}}\end{bmatrix}\begin{bmatrix}\underline{s_{i-1}}\\\underline{h_j}\end{bmatrix}\right)。论文只是把两个向量分别投影后相加,没有直接写成拼接形式。

不过,它本身不是 RNN。关键区别在于是否存在时间递归。

RNN 计算 ht=f(xt,ht1)\underline{h_t}=f(\underline{x_t},\underline{h_{t-1}}) 后,得到的 ht\underline{h_t} 会继续参与下一时刻的计算,即 ht1htht+1\underline{h_{t-1}}\rightarrow\underline{h_t}\rightarrow\underline{h_{t+1}}

Attention 打分计算的是 eij=f(si1,hj)e_{ij}=f(\underline{s_{i-1}},\underline{h_j})。得到的 eije_{ij} 只是当前位置的匹配分数,不会作为 attention 网络的隐藏状态继续传递。对于固定的解码步骤 iiei1,ei2,,eiTe_{i1},e_{i2},\ldots,e_{iT} 之间不存在 ei1ei2ei3e_{i1}\rightarrow e_{i2}\rightarrow e_{i3} 这样的递归关系,因此这些分数可以并行计算。

更准确地说,Additive Attention 的评分函数长得像 RNN 的状态更新函数,但它是一个前馈网络,而不是循环网络。

这种形式也和它的历史背景有关。Bahdanau Attention 本来就是为 RNN Encoder–Decoder 设计的。当时已经有编码器状态 hj\underline{h_j} 和解码器状态 si1\underline{s_{i-1}},研究者需要一个可训练函数判断二者是否匹配,自然采用了当时常见的 tanh(Wx+Uh)\tanh(\underline{\underline{W}}\underline{x}+\underline{\underline{U}}\underline{h}) 结构。可以把它理解为:用一个结构类似 RNN 单元内部计算的小型 MLP,给两个隐藏状态打分。

vaT\underline{v_a}^T 是 RNN 状态更新公式中没有的部分。RNN 中,tanh 的输出本身就是新的隐藏向量;attention 最终需要的却是一个标量分数。若 zij,vaRda\underline{z_{ij}},\underline{v_a}\in\mathbb{R}^{d_a},那么 vaTzijR\underline{v_a}^T\underline{z_{ij}}\in\mathbb{R}。这个标量越大,表示第 jj 个编码器位置与当前解码步骤越匹配。之后再对所有 eije_{ij} 做 softmax,得到真正用于加权聚合的注意力权重 αij\alpha_{ij}

所以,直觉可以更准确地表述为:它不是一个 RNN,而是它的隐藏层采用了与经典 RNN 单元相同的“两个线性投影相加,再经过 tanh”的计算结构。并且 si1\underline{s_{i-1}} 本身通常就来自 RNN 解码器,因此整套公式在视觉上会更像 RNN。

2. 归一化注意力权重#

注意力分数与权重

eije_{ij} 是未经归一化的注意力分数,αij\alpha_{ij} 是经过 softmax 后得到的注意力权重,满足 jαij=1\sum_j\alpha_{ij}=1,并实际用于编码器状态的加权聚合。

αij=exp(eij)k=1Txexp(eik),j=1Txαij=1.\alpha_{ij} = \frac{\exp(e_{ij})}{\sum_{k=1}^{T_x}\exp(e_{ik})}, \qquad \sum_{j=1}^{T_x}\alpha_{ij}=1.

实现 softmax 时减去最大分数,可以避免指数运算溢出。

这里的指数函数会把任意实数分数转换成正数,并放大不同分数之间的相对差异;除以所有指数分数之和,则把它们归一化成总和为 11 的权重。

3. 聚合上下文#

ci=j=1Txαijhj.\underline{c_i} = \sum_{j=1}^{T_x}\alpha_{ij}\underline{h_j}.

ci\underline{c_i} 是编码器状态的加权和,表示生成第 ii 个目标 token 时需要的源序列信息。

NumPy 实现#

支持 batch 的 PyTorch 实现#

这里用 Linear 实现 Wa\underline{\underline{W_a}}Ua\underline{\underline{U_a}}。它们只变换最后一个维度,因此 sexpandedRB×Tx×dss_{\mathrm{expanded}}\in\mathbb{R}^{B\times T_x\times d_s} 经过 Wa\underline{\underline{W_a}} 后,形状变为 B×Tx×daB\times T_x\times d_a,batch 和序列长度维度保持不变。

从 Additive Attention 到 Scaled Dot-Product Attention#

Additive Attention 与 Scaled Dot-Product Attention 的整体流程没有改变,仍然是“评分、归一化、聚合”。变化主要发生在两处:参与计算的表示经过了更明确的投影,并且匹配分数换了一种计算方式。

对单个解码状态和编码器状态,可以先建立下面的对应关系:

qT=si1TWQ,kjT=hjTWK,vjT=hjTWV.\underline{q}^T = \underline{s_{i-1}}^T\underline{\underline{W_Q}}, \qquad \underline{k_j}^T = \underline{h_j}^T\underline{\underline{W_K}}, \qquad \underline{v_j}^T = \underline{h_j}^T\underline{\underline{W_V}}.

Query q\underline{q} 表示当前需要寻找什么,Key kj\underline{k_j} 用来判断第 jj 个位置是否匹配,Value vj\underline{v_j} 则提供匹配后真正要聚合的内容。也就是说,Query 和 Key 负责计算注意力权重,Value 负责形成输出。

Additive Attention 中,与 Query、Key 角色相近的两个中间 vector 分别是 qadd=Wasi1\underline{q_{\mathrm{add}}}=\underline{\underline{W_a}}\underline{s_{i-1}}kadd,j=Uahj\underline{k_{\mathrm{add},j}}=\underline{\underline{U_a}}\underline{h_j},因此评分函数可以写成 eij=vaTtanh(qadd+kadd,j)e_{ij}=\underline{v_a}^T\tanh(\underline{q_{\mathrm{add}}}+\underline{k_{\mathrm{add},j}})。Scaled Dot-Product Attention 则使用上面定义的 q\underline{q}kj\underline{k_j},把评分函数改为 eij=qTkj/dke_{ij}=\underline{q}^T\underline{k_j}/\sqrt{d_k}

得到注意力权重后,Additive Attention 聚合原始编码器状态 hj\underline{h_j},而 Transformer 聚合投影后的 vj\underline{v_j}。Value 投影使模型可以分别学习“用什么进行匹配”和“匹配后取出什么内容”。

把一个序列中所有位置的 Query、Key 和 Value 分别堆叠起来,就得到下面使用的 matrix Q\underline{\underline{Q}}K\underline{\underline{K}}V\underline{\underline{V}}。在 Transformer 的 self-attention 中,它们通常都由同一个输入序列 matrix X\underline{\underline{X}} 投影得到。

这里把单个 token 的表示写成 row vector,因此使用 qT=si1TWQ\underline{q}^T=\underline{s_{i-1}}^T\underline{\underline{W_Q}}。把所有 token 的 row vector 堆叠成 X\underline{\underline{X}} 后,就得到 Q=XWQ\underline{\underline{Q}}=\underline{\underline{X}}\underline{\underline{W_Q}}。如果改用 column vector 约定,同一个投影应写成 q=WQTsi1\underline{q}=\underline{\underline{W_Q}}^T\underline{s_{i-1}}

Scaled Dot-Product Attention#

Transformer 在 2017 年的 Attention Is All You Need 中使用 Query、Key 和 Value 表示注意力计算。对输入序列矩阵 X\underline{\underline{X}}

Q=XWQ,K=XWK,V=XWV.\underline{\underline{Q}} = \underline{\underline{X}}\underline{\underline{W_Q}}, \qquad \underline{\underline{K}} = \underline{\underline{X}}\underline{\underline{W_K}}, \qquad \underline{\underline{V}} = \underline{\underline{X}}\underline{\underline{W_V}}.

对单个注意力头,暂时省略 batch 维度。若序列长度为 nn,则各 matrix 的形状为:

XRn×dmodel,WQ,WKRdmodel×dk,WVRdmodel×dv,Q,KRn×dk,VRn×dv.\begin{aligned} \underline{\underline{X}} &\in\mathbb{R}^{n\times d_{model}}, \\ \underline{\underline{W_Q}}, \underline{\underline{W_K}} &\in\mathbb{R}^{d_{model}\times d_k}, \\ \underline{\underline{W_V}} &\in\mathbb{R}^{d_{model}\times d_v}, \\ \underline{\underline{Q}}, \underline{\underline{K}} &\in\mathbb{R}^{n\times d_k}, \\ \underline{\underline{V}} &\in\mathbb{R}^{n\times d_v}. \end{aligned}

其中 dmodeld_{model} 是每个 token 表示的维度。它既决定输入 embedding 的宽度,也影响后续投影 matrix 的形状和模型参数量。dkd_k 是每个 Query、Key 的维度,dvd_v 是每个 Value 的维度。

Query 和 Key 决定“关注哪里”,Value 提供被聚合的内容。单个注意力头的计算为:

Attention(Q,K,V)=softmax(QKTdk)V.\operatorname{Attention}\left( \underline{\underline{Q}}, \underline{\underline{K}}, \underline{\underline{V}} \right) = \operatorname{softmax}\left( \frac{ \underline{\underline{Q}}\, \underline{\underline{K}}^T }{\sqrt{d_k}} \right) \underline{\underline{V}}.

维度记号: dkd_k 是每个 Query 和 Key 向量的维度,也就是点积中相加项的数量。从这里开始,后文中的 dxd_x 均表示向量 x\underline{x} 的维度。

为什么要缩放点积#

假设 q,kRdk\underline{q},\underline{k}\in\mathbb{R}^{d_k},并且各分量相互独立且服从标准正态分布,即 qi,kii.i.d.N(0,1)q_i,k_i\overset{\mathrm{i.i.d.}}{\sim}\mathcal{N}(0,1)。先看点积中的单项:

E[qiki]=E[qi]E[ki]=0,Var(qiki)=E[(qiki)2]E[qiki]2=E[qi2]E[ki2]=1.\begin{aligned} \mathbb{E}[q_i k_i] &= \mathbb{E}[q_i]\mathbb{E}[k_i] = 0, \\ \operatorname{Var}(q_i k_i) &= \mathbb{E}\left[(q_i k_i)^2\right] - \mathbb{E}[q_i k_i]^2 \\ &= \mathbb{E}[q_i^2]\mathbb{E}[k_i^2] = 1. \end{aligned}

因此,整个点积的期望与方差为:

E[qk]=i=1dkE[qiki]=0,Var(qk)=i=1dkVar(qiki)=dk.\begin{aligned} \mathbb{E}[\underline{q}\cdot\underline{k}] &= \sum_{i=1}^{d_k}\mathbb{E}[q_i k_i] = 0, \\ \operatorname{Var}(\underline{q}\cdot\underline{k}) &= \sum_{i=1}^{d_k}\operatorname{Var}(q_i k_i) = d_k. \end{aligned}

因此点积的标准差会随维度按 dk\sqrt{d_k} 增长。除以 dk\sqrt{d_k} 后,分数的数值范围更加稳定。如果省略缩放,较大的点积容易让 softmax 输出接近 one-hot,使梯度集中在少数高分位置,其他路径难以得到有效更新。

使用 Value 聚合#

对第 ii 个 Query,softmax 产生它对所有 Key 的注意力权重 αij\alpha_{ij}。随后使用这些权重聚合对应的 Value:

oi=j=1nαijvj.\underline{o_i} = \sum_{j=1}^{n}\alpha_{ij}\underline{v_j}.

这与 Additive Attention 中的 ci=jαijhj\underline{c_i}=\sum_j\alpha_{ij}\underline{h_j} 直接对应。主要区别是 Transformer 聚合经过投影的 vj\underline{v_j},而不是直接聚合 hj\underline{h_j}

Multi-Head Attention#

多头注意力将 dmodeld_{model} 维表示拆成 hh 个头。常见配置是:

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

每个头分别计算注意力,再拼接并投影回模型维度:

headi=Attention(Qi,Ki,Vi),MultiHead(Q,K,V)=Concat(head1,,headh)WO.\begin{aligned} \underline{\underline{\operatorname{head}_i}} &= \operatorname{Attention}\left( \underline{\underline{Q_i}}, \underline{\underline{K_i}}, \underline{\underline{V_i}} \right), \\ \operatorname{MultiHead}\left( \underline{\underline{Q}}, \underline{\underline{K}}, \underline{\underline{V}} \right) &= \operatorname{Concat}\left( \underline{\underline{\operatorname{head}_1}}, \ldots, \underline{\underline{\operatorname{head}_h}} \right) \underline{\underline{W_O}}. \end{aligned}

从单头到多头的维度变化#

原 notebook 还逐步检查了 score、注意力权重、单头输出、拼接结果和最终输出的形状。省略 batch 维度时,这条计算链可以写成:

Si=QiKiTdkRn×n,Ai=softmax(Si)Rn×n,headi=AiViRn×dv,C=Concat(head1,,headh)Rn×(hdv),WOR(hdv)×dmodel,O=CWORn×dmodel.\begin{aligned} \underline{\underline{S_i}} &= \frac{ \underline{\underline{Q_i}} \underline{\underline{K_i}}^T }{\sqrt{d_k}} \in\mathbb{R}^{n\times n}, \\ \underline{\underline{A_i}} &= \operatorname{softmax}\left( \underline{\underline{S_i}} \right) \in\mathbb{R}^{n\times n}, \\ \underline{\underline{\operatorname{head}_i}} &= \underline{\underline{A_i}} \underline{\underline{V_i}} \in\mathbb{R}^{n\times d_v}, \\ \underline{\underline{C}} &= \operatorname{Concat}\left( \underline{\underline{\operatorname{head}_1}}, \ldots, \underline{\underline{\operatorname{head}_h}} \right) \in\mathbb{R}^{n\times(hd_v)}, \\ \underline{\underline{W_O}} &\in\mathbb{R}^{(hd_v)\times d_{model}}, \\ \underline{\underline{O}} &= \underline{\underline{C}} \underline{\underline{W_O}} \in\mathbb{R}^{n\times d_{model}}. \end{aligned}

dk=dv=dmodel/hd_k=d_v=d_{model}/h 时,hdv=dmodelhd_v=d_{model},所以拼接后的宽度回到模型维度。加入 batch 维度并拆分 heads 后,主要张量形状如下:

张量形状
输入 XX(B,n,dmodel)(B,n,d_{model})
分头后的 Q,KQ,K(B,h,n,dk)(B,h,n,d_k)
分头后的 VV(B,h,n,dv)(B,h,n,d_v)
注意力分数与权重(B,h,n,n)(B,h,n,n)
各头的上下文(B,h,n,dv)(B,h,n,d_v)
最终输出(B,n,dmodel)(B,n,d_{model})

Mask#

注意力分数在 softmax 之前加入 mask。被遮蔽位置设为负无穷,softmax 后其权重为零。

Padding mask#

Padding mask 阻止模型关注补齐 token。若 mask 形状为 (B,n)(B,n),扩展为 (B,1,1,n)(B,1,1,n) 后即可广播到所有注意力头和查询位置:

mask = padding_mask[:, None, None, :]
scores = scores.masked_fill(mask.bool(), float("-inf"))
python

Causal mask#

自回归解码要求第 tt 个位置只能看到不晚于自己的 token。上三角位置需要被遮蔽:

[[0, -inf, -inf],
 [0,    0, -inf],
 [0,    0,    0]]
text
causal = torch.triu(
    torch.ones(seq_len, seq_len, device=scores.device, dtype=torch.bool),
    diagonal=1,
)
scores = scores.masked_fill(causal[None, None, :, :], float("-inf"))
python

完整 PyTorch 实现#

下面两张热力图来自原 notebook。第一张展示 causal mask 形成的下三角注意力,第二张还叠加了 padding mask,使后两个 Key 位置不可见。

Causal mask 下的注意力权重

Causal mask 与 padding mask 共同作用

小结#

Additive Attention 与 Scaled Dot-Product Attention 使用不同的匹配函数,但都遵循“评分、归一化、聚合”三个步骤。Multi-Head Attention 进一步让模型在多个表示子空间中并行学习关系,而 mask 则把序列结构和有效长度约束注入注意力计算。

Transformer - 从 Additive Attention 到 Multi-Head Attention
https://heng-blog.pages.dev/blog/from-additive-to-multi-head-attention
Author 杨苏恒
Published at April 19, 2025