10.3 Scaled Dot-Product 与 Multi-Head Attention:先把 Mask 语义写对
RNN 把历史压进一个连续更新的状态。情报官提出另一种查询方式:处理当前任务时,直接向档案中的所有相关位置发出 query,再按匹配程度汇总 value。
Attention 就是一个可微检索/加权汇总。真正容易出错的不是 Q/K/V 比喻,而是 shape、缩放、softmax 维度和 mask 语义。
本课目标
- 推导 scaled dot-product attention 的 shape 与缩放;
- 区分 self、cross、causal 与 padding attention;
- 正确构造和测试 Boolean/additive masks;
- 理解 multi-head projection、拼接与输出映射;
- 识别 attention 的复杂度和解释限制。
1. Q、K、V 是 Learned Projections
输入表示 $X\in\mathbb R^{B\times T\times d_{model}}$。Self-attention 常计算:
$$ Q=XW_Q,qquad K=XW_K,qquad V=XW_V. $$
Query/key 决定匹配分数,value 提供被汇总内容。这些角色由训练目标学出,不保证等同数据库 key/value 或人类语义标签。
Cross-attention 中 Q 来自 target/decoder states,K/V 来自 source/encoder states,因此 query length $L$ 与 source length $S$ 可不同。
2. Scaled Dot-Product Attention
单头:
$$ A=\operatorname{softmax}\left(\frac{QK^T}{\sqrt{d_k}}+M\right), $$
$$ O=AV. $$
若:
Q: [B, L, d_k]
K: [B, S, d_k]
V: [B, S, d_v]
scores/A: [B, L, S]
O: [B, L, d_v]Softmax 沿 key/source 维 $S$。每个 query row 的 allowed weights 和为 1。
若 Q/K components 近似零均值、单位方差且独立,dot-product variance 随 $d_k$ 增长;除以 $\sqrt{d_k}$ 把典型尺度拉回,减少 softmax 过早饱和。这是近似动机,不保证实际 trained representations 满足独立假设。
3. Mask 必须在 Softmax 前应用
Additive mask 对禁止位置加 $-\infty$/足够负值,使 softmax 权重为 0。Softmax 后再把权重乘 0 会让剩余权重不再归一,除非重新归一化。
常见 masks:
- key padding mask:所有 queries 忽略 source padding;
- causal mask:位置 $t$ 不看未来 $>t$;
- local/block/sparse mask:只允许指定连接;
- cross-attention source mask:忽略 encoder padding。
Query padding 也要在 loss/output 处 mask;只屏蔽 padded keys 不会自动删除 padded query outputs。
4. Boolean Mask 语义并不统一
这是工程高危点。当前 PyTorch API 中:
nn.MultiheadAttention的 Booleankey_padding_mask/attn_mask:True表示不允许/忽略;torch.nn.functional.scaled_dot_product_attention的 Booleanattn_mask:True表示允许参与。
从一个 API 迁移到另一个可能需要取反。不要凭变量名猜,写一个 $3\times3$ 人工例子验证 forbidden weights 为 0。
若某个 query 的所有 keys 都被 mask,softmax 没有合法分布,可能得到 NaN 或未定义行为。Batch 构造必须保证至少一个合法 key,或显式定义 empty-row 结果。
5. 从零写一个单头版本
import math
import torch
def scaled_dot_product_attention(query, key, value, allowed=None):
# query: [B, L, D], key: [B, S, D], value: [B, S, Dv]
scores = query @ key.transpose(-2, -1) / math.sqrt(query.shape[-1])
if allowed is not None:
# allowed 应可 broadcast 到 [B, L, S];True 表示允许。
allowed = torch.broadcast_to(allowed, scores.shape)
if not torch.all(allowed.any(dim=-1)):
raise ValueError("Every query must have at least one allowed key")
scores = scores.masked_fill(~allowed, float("-inf"))
weights = torch.softmax(scores, dim=-1)
output = weights @ value
return output, weights
q = torch.randn(2, 4, 8)
k = torch.randn(2, 6, 8)
v = torch.randn(2, 6, 5)
out, weight = scaled_dot_product_attention(q, k, v)
assert out.shape == (2, 4, 5)
assert torch.allclose(weight.sum(dim=-1), torch.ones(2, 4))生产代码优先用 framework optimized kernels;这个版本用于建立 shape/mask tests,不处理 fused precision、dropout、GQA 和高性能布局。
6. Causal Mask
长度 $T$ 的 decoder self-attention,allowed matrix:
q0: k0
q1: k0 k1
q2: k0 k1 k2
...即下三角(含 diagonal)。Causal mask 防止训练时 token $t$ 读取未来 targets。仅把 labels shift 一位却忘记 mask,模型会通过 hidden representation 偷看答案,训练 loss 异常低。
Prefix-LM、bidirectional encoder、sequence-to-sequence decoder 的 mask 不同。不要把“Transformer mask”当单一模板。
7. Multi-head Attention
将 $d_{model}$ 投影成 $H$ 组 head subspaces:
$$ head_h=Attention(XW_Q^{(h)},XW_K^{(h)},XW_V^{(h)}), $$
$$ MHA(X)=Concat(head_1,\ldots,head_H)W_O. $$
常见实现要求 $d_{model}$ 可被 heads 整除,$d_h=d_{model}/H$。增加 heads 不一定增加总 projection 参数,但会改变每头维度、kernel efficiency 和表示分解。
“每头分别学习语法、指代、位置”只是某些分析观察,不是训练约束。Heads 可冗余、混合或不稳定,不能给每头强行命名。
8. Framework API 的性能与 Dropout
nn.MultiheadAttention(batch_first=True) 输入可用 [B,T,D]。只需输出、不需要 weights 时设 need_weights=False,框架更可能使用优化的 scaled-dot-product path。
直接调用当前 F.scaled_dot_product_attention 时,dropout_p>0 会按参数应用 dropout,不会自动读取 module 的 train/eval;module 通常要传:
dropout_p = self.dropout if self.training else 0.0Fused Flash/memory-efficient/math kernels 可能因 dtype、device、mask shape 而切换,并有浮点差异。性能测试要记录 backend,不要只从公式估 latency。
9. Attention 没有位置信息
若同时置换 X 的 token 顺序,未加 position signal 的 self-attention 输出也相应置换;它对集合式输入是 permutation equivariant,不知道第 1 和第 10 个位置的差别。
位置可通过 absolute embedding、sinusoidal、relative bias、rotary 等机制注入。不同方法影响长度外推和 cache 实现,第 10.4 课展开。
10. 复杂度不只有 $O(T^2)$ 一行
Self-attention score matrix 大小约 [B,H,T,T],标准 dense attention 的 score 计算/存储随 $T^2$ 增长;projection 和 FFN 还依赖 $Td_{model}^2$。
在短序列/大 hidden 下 projection/FFN 可能主导;长序列时 attention matrix 常成为瓶颈。FlashAttention 类算法可减少 materialized memory/I/O,并保持 exact attention 语义(允许浮点实现差异),但不会把所有理论工作都自动变线性。
Sparse/local/linear attention 改变连接或 kernel approximation,要验证任务质量和实际加速。
11. Attention Weight 不是因果解释
高 weight 表示当前 head/layer 中 value mixture 的系数,不等于 token 对最终输出的唯一贡献:
- value vectors 和 output projection 会改变效果;
- residual/FFN/后续 layers 有其他路径;
- 多组 weights 可产生相近输出;
- 干预 token 会同时改变 Q/K/V。
可把 weights 用作诊断,但应结合 ablation、gradient/perturbation 与 counterfactual analysis,仍不能直接宣称因果。
常见误区
- Q/K/V 是固定语义字段:它们是 learned projections。
- Mask 后乘 0 就够:应在 softmax 前屏蔽并测试归一化。
- PyTorch 所有 Boolean mask 都是 True=保留:不同 API 语义可相反。
- 每个 head 会自动分工成人类概念:没有这个约束。
- Attention weight 就是解释:最终输出还有 value、projection 和其他路径。
练习
- 为 cross-attention 写出 Q/K/V/scores/output shapes。
- 构造 $3\times3$ causal mask,验证未来 weights 为 0。
- 对比
MultiheadAttention与 SDPA Boolean mask 语义。 - 令某 query 全部被 mask,观察并修复失败。
- 估算不同 $T,D,H$ 下 attention score 与 FFN 的内存/运算。
小结
Scaled dot-product attention 用 Q/K 相似度产生每个 query 对 values 的归一化混合。缩放控制 softmax 输入尺度,mask 定义信息流,多头在不同 learned subspaces 并行计算。它擅长直接连接位置,也引入二次矩阵、API 语义和解释风险。
下一课把 attention 放进 residual、normalization 和 FFN,加入位置与 causal objective,组成 encoder、decoder 和完整 Transformer。