跳到内容

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. $$

若:

text
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 的 Boolean key_padding_mask/attn_maskTrue 表示不允许/忽略;
  • torch.nn.functional.scaled_dot_product_attention 的 Boolean attn_maskTrue 表示允许参与。

从一个 API 迁移到另一个可能需要取反。不要凭变量名猜,写一个 $3\times3$ 人工例子验证 forbidden weights 为 0。

若某个 query 的所有 keys 都被 mask,softmax 没有合法分布,可能得到 NaN 或未定义行为。Batch 构造必须保证至少一个合法 key,或显式定义 empty-row 结果。

5. 从零写一个单头版本

python
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:

text
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 通常要传:

python
dropout_p = self.dropout if self.training else 0.0

Fused 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 和其他路径。

练习

  1. 为 cross-attention 写出 Q/K/V/scores/output shapes。
  2. 构造 $3\times3$ causal mask,验证未来 weights 为 0。
  3. 对比 MultiheadAttention 与 SDPA Boolean mask 语义。
  4. 令某 query 全部被 mask,观察并修复失败。
  5. 估算不同 $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。

Built with VitePress | Software Systems Atlas