注意力机制的共同目标是:根据当前查询为一组候选信息分配权重,再将它们聚合为上下文表示。本文从机器翻译中的 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)为 si−1,编码器的全部隐藏状态为 {h1,…,hTx}。生成目标 token yi 时,需要判断当前解码状态应该关注源序列中的哪些位置,因此要分别计算 si−1 与每个 hj 的匹配程度。
记号说明: 下文用单下划线表示 vector,用双下划线表示 matrix;scalar 不加下划线。
1. 计算匹配分数#
eij=vaTtanh(Wasi−1+Uahj).
Wa、Ua 和 va 都是可训练参数。若 si−1∈Rds、hj∈Rdh,并将 attention 隐藏空间的维度记为 da,那么 Wa∈Rda×ds、Ua∈Rda×dh、va∈Rda。
这个评分函数可视为一个带 tanh 的小型前馈网络,用于估计当前解码状态与每个编码器状态的匹配程度。va 的主要作用是将 tanh 输出的 vector 转换成 scalar,作为最终的匹配分数 eij。
2. 归一化注意力权重#
αij=∑k=1Txexp(eik)exp(eij),j=1∑Txαij=1.
实现 softmax 时减去最大分数,可以避免指数运算溢出。
这里的指数函数会把任意实数分数转换成正数,并放大不同分数之间的相对差异;除以所有指数分数之和,则把它们归一化成总和为 1 的权重。
3. 聚合上下文#
ci=j=1∑Txαijhj.
ci 是编码器状态的加权和,表示生成第 i 个目标 token 时需要的源序列信息。
NumPy 实现#
import numpy as np
def additive_attention(s_prev, h, W_a, U_a, v_a):
"""
s_prev: (hidden_size,)
h: (source_length, hidden_size)
W_a: (attention_dim, hidden_size)
U_a: (attention_dim, hidden_size)
v_a: (attention_dim,)
"""
source_length = h.shape[0]
s_expanded = np.broadcast_to(s_prev, (source_length, s_prev.size))
energy = np.tanh(s_expanded @ W_a.T + h @ U_a.T) @ v_a
weights = np.exp(energy - np.max(energy))
weights /= np.sum(weights)
context = weights @ h
return context, weights
python
支持 batch 的 PyTorch 实现#
这里用 Linear 实现 Wa 和 Ua。它们只变换最后一个维度,因此 sexpanded∈RB×Tx×ds 经过 Wa 后,形状变为 B×Tx×da,batch 和序列长度维度保持不变。
import torch
import torch.nn as nn
import torch.nn.functional as F
class AdditiveAttention(nn.Module):
def __init__(self, hidden_size, attention_dim):
super().__init__()
self.W_a = nn.Linear(hidden_size, attention_dim, bias=False)
self.U_a = nn.Linear(hidden_size, attention_dim, bias=False)
self.v_a = nn.Parameter(torch.randn(attention_dim))
def forward(self, s_prev, h):
# s_prev: (batch, hidden_size)
# h: (batch, source_length, hidden_size)
source_length = h.size(1)
s_expanded = s_prev.unsqueeze(1).expand(-1, source_length, -1)
energy = torch.tanh(self.W_a(s_expanded) + self.U_a(h))
scores = torch.matmul(energy, self.v_a)
weights = F.softmax(scores, dim=-1)
context = torch.bmm(weights.unsqueeze(1), h).squeeze(1)
return context, weights
python
从 Additive Attention 到 Scaled Dot-Product Attention#
Additive Attention 与 Scaled Dot-Product Attention 的整体流程没有改变,仍然是“评分、归一化、聚合”。变化主要发生在两处:参与计算的表示经过了更明确的投影,并且匹配分数换了一种计算方式。
对单个解码状态和编码器状态,可以先建立下面的对应关系:
qT=si−1TWQ,kjT=hjTWK,vjT=hjTWV.
Query q 表示当前需要寻找什么,Key kj 用来判断第 j 个位置是否匹配,Value vj 则提供匹配后真正要聚合的内容。也就是说,Query 和 Key 负责计算注意力权重,Value 负责形成输出。
Additive Attention 中,与 Query、Key 角色相近的两个中间 vector 分别是 qadd=Wasi−1 和 kadd,j=Uahj,因此评分函数可以写成 eij=vaTtanh(qadd+kadd,j)。Scaled Dot-Product Attention 则使用上面定义的 q 和 kj,把评分函数改为 eij=qTkj/dk。
得到注意力权重后,Additive Attention 聚合原始编码器状态 hj,而 Transformer 聚合投影后的 vj。Value 投影使模型可以分别学习“用什么进行匹配”和“匹配后取出什么内容”。
把一个序列中所有位置的 Query、Key 和 Value 分别堆叠起来,就得到下面使用的 matrix Q、K 和 V。在 Transformer 的 self-attention 中,它们通常都由同一个输入序列 matrix X 投影得到。
这里把单个 token 的表示写成 row vector,因此使用 qT=si−1TWQ。把所有 token 的 row vector 堆叠成 X 后,就得到 Q=XWQ。如果改用 column vector 约定,同一个投影应写成 q=WQTsi−1。
Scaled Dot-Product Attention#
Transformer 在 2017 年的 Attention Is All You Need ↗ 中使用 Query、Key 和 Value 表示注意力计算。对输入序列矩阵 X:
Q=XWQ,K=XWK,V=XWV.
对单个注意力头,暂时省略 batch 维度。若序列长度为 n,则各 matrix 的形状为:
XWQ,WKWVQ,KV∈Rn×dmodel,∈Rdmodel×dk,∈Rdmodel×dv,∈Rn×dk,∈Rn×dv.
其中 dmodel 是每个 token 表示的维度。它既决定输入 embedding 的宽度,也影响后续投影 matrix 的形状和模型参数量。dk 是每个 Query、Key 的维度,dv 是每个 Value 的维度。
Query 和 Key 决定“关注哪里”,Value 提供被聚合的内容。单个注意力头的计算为:
Attention(Q,K,V)=softmax(dkQKT)V.
维度记号: dk 是每个 Query 和 Key 向量的维度,也就是点积中相加项的数量。从这里开始,后文中的 dx 均表示向量 x 的维度。
为什么要缩放点积#
假设 q,k∈Rdk,并且各分量相互独立且服从标准正态分布,即 qi,ki∼i.i.d.N(0,1)。先看点积中的单项:
E[qiki]Var(qiki)=E[qi]E[ki]=0,=E[(qiki)2]−E[qiki]2=E[qi2]E[ki2]=1.
因此,整个点积的期望与方差为:
E[q⋅k]Var(q⋅k)=i=1∑dkE[qiki]=0,=i=1∑dkVar(qiki)=dk.
因此点积的标准差会随维度按 dk 增长。除以 dk 后,分数的数值范围更加稳定。如果省略缩放,较大的点积容易让 softmax 输出接近 one-hot,使梯度集中在少数高分位置,其他路径难以得到有效更新。
使用 Value 聚合#
对第 i 个 Query,softmax 产生它对所有 Key 的注意力权重 αij。随后使用这些权重聚合对应的 Value:
oi=j=1∑nαijvj.
这与 Additive Attention 中的 ci=∑jαijhj 直接对应。主要区别是 Transformer 聚合经过投影的 vj,而不是直接聚合 hj。
Multi-Head Attention#
多头注意力将 dmodel 维表示拆成 h 个头。常见配置是:
dk=dv=hdmodel.
每个头分别计算注意力,再拼接并投影回模型维度:
headiMultiHead(Q,K,V)=Attention(Qi,Ki,Vi),=Concat(head1,…,headh)WO.
从单头到多头的维度变化#
原 notebook 还逐步检查了 score、注意力权重、单头输出、拼接结果和最终输出的形状。省略 batch 维度时,这条计算链可以写成:
SiAiheadiCWOO=dkQiKiT∈Rn×n,=softmax(Si)∈Rn×n,=AiVi∈Rn×dv,=Concat(head1,…,headh)∈Rn×(hdv),∈R(hdv)×dmodel,=CWO∈Rn×dmodel.
当 dk=dv=dmodel/h 时,hdv=dmodel,所以拼接后的宽度回到模型维度。加入 batch 维度并拆分 heads 后,主要张量形状如下:
| 张量 | 形状 |
|---|
| 输入 X | (B,n,dmodel) |
| 分头后的 Q,K | (B,h,n,dk) |
| 分头后的 V | (B,h,n,dv) |
| 注意力分数与权重 | (B,h,n,n) |
| 各头的上下文 | (B,h,n,dv) |
| 最终输出 | (B,n,dmodel) |
Mask#
注意力分数在 softmax 之前加入 mask。被遮蔽位置设为负无穷,softmax 后其权重为零。
Padding mask#
Padding mask 阻止模型关注补齐 token。若 mask 形状为 (B,n),扩展为 (B,1,1,n) 后即可广播到所有注意力头和查询位置:
mask = padding_mask[:, None, None, :]
scores = scores.masked_fill(mask.bool(), float("-inf"))
python
Causal mask#
自回归解码要求第 t 个位置只能看到不晚于自己的 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 实现#
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, num_heads):
super().__init__()
if d_model % num_heads != 0:
raise ValueError("d_model must be divisible by num_heads")
self.d_model = d_model
self.num_heads = num_heads
self.d_k = d_model // num_heads
self.W_q = nn.Linear(d_model, d_model)
self.W_k = nn.Linear(d_model, d_model)
self.W_v = nn.Linear(d_model, d_model)
self.W_o = nn.Linear(d_model, d_model)
def forward(
self,
query,
key=None,
value=None,
padding_mask=None,
causal=False,
):
key = query if key is None else key
value = key if value is None else value
batch_size, query_len, _ = query.shape
key_len = key.size(1)
if value.size(1) != key_len:
raise ValueError("key and value must have the same sequence length")
if causal and query_len != key_len:
raise ValueError("causal attention requires equal query and key lengths")
def split_heads(projection, seq_len):
# Linear 一次生成所有 heads,再 reshape:
# (B, n, d_model) -> (B, n, h, d_k) -> (B, h, n, d_k)
return projection.view(
batch_size, seq_len, self.num_heads, self.d_k
).transpose(1, 2)
Q = split_heads(self.W_q(query), query_len)
K = split_heads(self.W_k(key), key_len)
V = split_heads(self.W_v(value), key_len)
scores = Q @ K.transpose(-2, -1) / math.sqrt(self.d_k)
if padding_mask is not None:
scores = scores.masked_fill(
padding_mask[:, None, None, :].bool(), float("-inf")
)
if causal:
causal_mask = torch.triu(
torch.ones(
query_len,
key_len,
device=query.device,
dtype=torch.bool,
),
diagonal=1,
)
scores = scores.masked_fill(
causal_mask[None, None, :, :], float("-inf")
)
weights = F.softmax(scores, dim=-1)
context = weights @ V
# transpose 后内存通常不连续,先调用 contiguous() 再 view
context = context.transpose(1, 2).contiguous().view(
batch_size, query_len, self.d_model
)
return self.W_o(context), weights
python
下面两张热力图来自原 notebook。第一张展示 causal mask 形成的下三角注意力,第二张还叠加了 padding mask,使后两个 Key 位置不可见。


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