读 Transformer 时,最容易卡住的地方是:Q、K、V 到底是什么,为什么要分成多个头?先看一个直觉:处理句子里某个词时,我们希望它按需要读取其他词的信息。当前位置提出问题(query),各位置提供可匹配的索引(key)和被读取的内容(value);匹配程度越高,读取的比例越大。这三个名字描述的是作用,不意味着它们必须来自三份不同的输入。
从一个位置到一整段序列
设批大小为 B,查询长度为 Lq,键和值的长度为 Lk。投影后 Q ∈ R^(B×Lq×dk)、K ∈ R^(B×Lk×dk)、V ∈ R^(B×Lk×dv)。矩阵乘法先得到每个查询与每个键的分数,再在键所在的维度做 softmax:
S = Q Kᵀ / √dk [B, Lq, Lk]
A = softmax(S + M, dim=-1) [B, Lq, Lk]
O = A V [B, Lq, dv]
M 可选;允许关注的位置加 0,禁止的位置加 −∞。因此每行 A 的权重和为 1,输出是 V 的加权和。除以 √dk 是为了减轻维度增大时点积分数变大、softmax 容易饱和的问题;它不是把向量归一化。这个计算顺序可以对照 Vaswani 等人的原论文图 2 左侧。
Self-attention 指 Q、K、V 都由同一段输入 X ∈ R^(B×L×dmodel) 分别线性投影得到。虽然来源相同,三组投影参数通常不同。若查询来自解码器、键和值来自编码器输出,则是 cross-attention;此时 Lq 与 Lk 可以不同。注意力输出是上下文表示,不等于“选中”某一个词。
为什么要有多个头
单个头只产生一套权重。多头注意力用不同的投影参数建立 h 套表示,分别计算注意力,再把结果拼接并投影回模型宽度:
head_i = Attention(Q WᵢQ, K WᵢK, V WᵢV)
MultiHead(Q,K,V) = Concat(head_1, …, head_h) Wᴼ
常见设置令 dk = dv = dmodel / h,所以每个头比较窄;拼接后的宽度仍是 dmodel。多个头有机会关注不同位置或表示子空间,但不能保证每个头都有固定、可解释的语法职责。原论文图 2 右侧展示了“分别投影 → 并行注意力 → 拼接 → 输出投影”的流程。
两种掩码不要混淆
- Padding mask:一批句子补齐长度后,禁止读取填充出来的键位置。它通常随样本变化,形状可为
[B, Lk]。 - Causal mask:自回归生成时,查询位置
i不能读取未来位置j > i,通常是[L, L]的上三角。编码器的双向 self-attention 一般不使用它。
两者可以同时存在。掩码必须在 softmax 之前加到分数上;算完权重再把未来位置乘零,剩余权重并不会自动重新归一化。每个有效查询还应至少保留一个可读的键,否则一整行都是 −∞,softmax 无法给出有效概率。
下面的实现只展示核心运算。blocked 中的 True 表示禁止连接;对于多头分数,它需要能广播到 [B, H, Lq, Lk]。例如 causal mask 用 [L, L],padding mask 用 [B, 1, 1, Lk]。
import math
import torch
from torch import nn
class TinyMHA(nn.Module):
def __init__(self, d_model: int, heads: int):
super().__init__()
assert d_model % heads == 0
self.heads = heads
self.d_head = d_model // heads
self.q = nn.Linear(d_model, d_model)
self.k = nn.Linear(d_model, d_model)
self.v = nn.Linear(d_model, d_model)
self.out = nn.Linear(d_model, d_model)
def split(self, x): # [B, L, D] → [B, H, L, D/H]
b, length, _ = x.shape
return x.reshape(b, length, self.heads, self.d_head).transpose(1, 2)
def forward(self, query, memory, blocked=None):
q = self.split(self.q(query))
k = self.split(self.k(memory))
v = self.split(self.v(memory))
score = q @ k.transpose(-2, -1) / math.sqrt(self.d_head)
if blocked is not None:
score = score.masked_fill(blocked, float("-inf"))
weight = torch.softmax(score, dim=-1)
context = weight @ v
b, _, length, _ = context.shape
context = context.transpose(1, 2).contiguous().reshape(b, length, -1)
return self.out(context)
调用 layer(x, x) 是 self-attention;调用 layer(decoder_state, encoder_state) 是 cross-attention。教学代码省略了 dropout、残差和归一化,它们属于完整 Transformer 层。尤其注意:blocked 的布尔语义和不同 PyTorch API 的 attn_mask 不总相同,换 API 时应查对应文档,而不是直接复用掩码。
延伸阅读: Attention Is All You Need(原论文,§3.2、图 2);PyTorch MultiheadAttention 文档。