QKV 与多头注意力
从相似度、缩放项、掩码和张量形状推导注意力,并写出可与 PyTorch 对齐的多头实现。
- 建议时长
- 8–10 小时
- 难度
- 核心机制
- 课程位置
- 2 / 12
学完这一课,你应该能:
- 能解释 Q、K、V 分别承担的匹配与聚合职责
- 能推导除以 √dₖ 的方差原因
- 能正确构造 causal mask 与 padding mask
- 能实现多头注意力并验证前向和梯度
1. Attention 在解决什么问题
当前位置 (i) 要从所有可见位置 (j) 收集信息。过程分两步:
- 匹配:当前位置需要什么,历史位置能提供什么?
- 聚合:按匹配权重,对历史信息加权求和。
Q = XWq 当前 token 想找什么
K = XWk 每个 token 如何被匹配
V = XWv 被找到后交付的内容
Q 与 K 用来“算权重”,V 用来“传内容”。分开投影让模型可以独立学习检索标准与信息载荷。
2. 为什么要除以 (\sqrt{d_k})
单个 query 与 key 的相似度是点积:
sⱼ = q · kⱼ
若 q、k 各分量独立、均值为 0、方差约为 1,点积是 (d_k) 项乘积之和,方差随 (d_k) 增长。维度变大时 logits 绝对值变大,softmax 更容易饱和,梯度变得尖锐。缩放后:
Attention(Q,K,V) = softmax(QKᵀ / √dₖ + M)V
缩放使 logits 方差回到较稳定量级。(M) 是 mask:允许位置加 0,不允许位置加极大负数。必须在 softmax 之前加,因为被屏蔽位置不应参与归一化。
3. 一次写清所有张量形状
设 batch 为 B、序列长 T、模型维 D、头数 H、每头维度 Dh=D/H:
Q,K,V projection: [B,T,D]
split heads: [B,H,T,Dh]
Q @ Kᵀ: [B,H,T,T]
softmax(scores)@V: [B,H,T,Dh]
merge heads: [B,T,D]
调试注意力时,先打印这六个形状,再看热力图。大多数错误来自 transpose 轴次序、mask 广播或 head 合并。
4. Causal mask 与 padding mask
自回归模型位置 (i) 不能看到未来位置 (j>i):
[[0, -∞, -∞, -∞],
[0, 0, -∞, -∞],
[0, 0, 0, -∞],
[0, 0, 0, 0]]
Causal mask 常为 [1,1,T,T],沿 batch 和 head 广播。Padding mask 从 [B,T] 变成 [B,1,1,T],屏蔽作为 key/value 来源的补齐 token。注意不同 API 对布尔 mask 的定义可能相反,不能凭习惯猜 True 是保留还是屏蔽。
5. 最小多头注意力实现
import math
import torch
from torch import nn
class MultiHeadSelfAttention(nn.Module):
def __init__(self, d_model: int, n_heads: int):
super().__init__()
assert d_model % n_heads == 0
self.n_heads = n_heads
self.d_head = d_model // n_heads
self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)
self.out = nn.Linear(d_model, d_model, bias=False)
def forward(self, x, padding_mask=None):
b, t, d = x.shape
q, k, v = self.qkv(x).chunk(3, dim=-1)
def split(z):
return z.view(b, t, self.n_heads, self.d_head).transpose(1, 2)
q, k, v = map(split, (q, k, v))
scores = q @ k.transpose(-2, -1) / math.sqrt(self.d_head)
future = torch.triu(
torch.ones(t, t, device=x.device, dtype=torch.bool),
diagonal=1,
)
scores = scores.masked_fill(future[None, None], float("-inf"))
if padding_mask is not None:
scores = scores.masked_fill(
padding_mask[:, None, None, :],
float("-inf"),
)
weights = torch.softmax(scores, dim=-1)
context = weights @ v
context = context.transpose(1, 2).contiguous().view(b, t, d)
return self.out(context), weights
transpose 通常只改变 stride。随后使用 view 前调用 contiguous(),是为了让底层布局和目标形状一致。使用 reshape 也可自动处理,但理解 stride 有助于排查性能问题。
6. 不止比较输出形状:对齐官方实现
复制相同权重,比较前向与输入梯度:
torch.manual_seed(17)
b, t, d, h = 2, 5, 16, 4
x1 = torch.randn(b, t, d, dtype=torch.float64, requires_grad=True)
x2 = x1.detach().clone().requires_grad_(True)
mine = MultiHeadSelfAttention(d, h).double()
ref = nn.MultiheadAttention(d, h, bias=False, batch_first=True).double()
with torch.no_grad():
ref.in_proj_weight.copy_(mine.qkv.weight)
ref.out_proj.weight.copy_(mine.out.weight)
mask = torch.triu(torch.ones(t, t, dtype=torch.bool), diagonal=1)
y1, _ = mine(x1)
y2, _ = ref(x2, x2, x2, attn_mask=mask, need_weights=False)
print("forward diff", (y1 - y2).abs().max().item())
y1.sum().backward()
y2.sum().backward()
print("gradient diff", (x1.grad - x2.grad).abs().max().item())
验收标准不是“能运行”,而是相同输入、权重和 mask 下,前向与梯度最大误差都足够小。
7. 复杂度与系统瓶颈
标准注意力分数矩阵大小为 [T,T],计算近似 (O(T^2D)),中间概率内存近似 (O(T^2))。但不要只背复杂度:
- 短序列、大 D 时,QKV 与 FFN 线性层可能占更多 FLOPs。
- 训练要保存反向中间量,Attention 概率显著占显存。
- 自回归 decode 用 KV Cache 复用历史 K/V,但要持续读取增长的缓存。
- FlashAttention 不改变精确注意力数学结果,主要通过分块与 IO-aware 算法减少显存读写。
8. 注意力权重不自动等于解释
热力图适合发现 mask、位置偏好与特殊 token 行为,但高权重不等于对最终输出因果重要:
- Value 与输出投影会继续改变影响。
- 多层、多头与残差共同作用。
- 相关性不能替代干预实验。
真正分析重要性应结合梯度、特征归因、attention rollout 或遮蔽干预。
9. 本课实验
| 变量 | 对照 | 必须观察 |
|---|---|---|
| 缩放项 | 除以 √dₖ / 不缩放 | logits 方差、entropy、grad norm |
| causal mask | 正确 / 关闭 | 是否偷看未来 |
| head 数 | 1 / 2 / 4 / 8 | 吞吐、效果、单头维度 |
| 实现 | 自研 / PyTorch | forward diff、gradient diff |
再注入一个错误:把 padding mask 写成 [B,T,1,1]。它可能可以广播,但屏蔽轴完全错误。用只有两个有效 token 的样本证明这一点。
10. 练习与答案提示
- D=768、H=12 时 Dh?**答:**64,分数最后两维是
[T,T]。 - 为什么不能 softmax 后再乘 mask?**答:**概率和不再为 1;重新归一化等价于前置 mask。
- 总 D 不变时,头数增加会让 QKV 参数量增长吗?**答:**不会,矩阵只是被重排为多个头。
- causal LM 的第 0 个位置能看什么?**答:**只能看到自身和模型约定的前缀。
- 为什么点积适合作为相似度?**答:**能高效批量矩阵乘法,并与可学习投影自然组合。