Chapter 02

注意力机制:可微的软字典查找

第 01 章给了 attention 的软查表直觉和通信/计算框架——本章把直觉变成精确机制:QKV 的矩阵运算、为什么除以 √d_k、因果掩码、多头。

本章交付:scaled dot-product 完整数学 · √d_k 方差论证 · 因果掩码机制 · 多头子空间与 induction heads · self vs cross · 注意力权重≠重要度 | 带走的能力:徒手推导 Attention(Q,K,V) = softmax(QKᵀ/√d_k)·V 的每一步形状,并能解释为什么这四个设计决策缺一不可

读完本章,你脑子里会多出这几条

  • scaled dot-product 的确切公式与形状:Q(n,d_k)·Kᵀ(d_k,n) → (n,n) 相似度矩阵 → softmax 按行 → ×V(n,d_v) → 输出(n,d_v);每一步形状可徒手验证。
  • 除以 √d_k 不是随手一除,是防 softmax 饱和:d_k 维点积的方差等于 d_k,标准差 √d_k;不除,大值把 softmax 推进近 one-hot 区,梯度 σ(1−σ)≈0。
  • 因果掩码 = 未来位置在 softmax 前置 −∞;在 softmax 后置零会改分布;对角线可见(位置 t 看得见自己);错位一格就泄露 t+1。
  • 多头 = 把 d_model 切成 h 个 d_k 维子空间并行运算、concat、Wᴼ 投回;总 FLOPs 近乎不变;induction heads 会真实专门化,相变式突现。

2.1scaled dot-product attention 的完整数学

attention 的全部计算是:把 Q 和 K 做点积得相似度矩阵、除以 √d_k 稳定数值、softmax 归一化成权重、再对 V 加权平均——四步,每步形状唯一确定。

第 01 章的软查表直觉(query 对每个 key 打相似度、softmax 成权重、对 value 加权平均)已经给了框架。现在把它落成矩阵运算,让形状可以逐步验算。

scaled dot-product attention 数学
Attention(Q, K, V) = softmax( Q · Kᵀ / √d_k ) · V

输入形状:
  Q : (n, d_k)    n = 序列长度,d_k = 查询/键的维度
  K : (n, d_k)    与 Q 同形(self-attention 时 Q/K/V 来自同一序列)
  V : (n, d_v)    d_v = 值的维度(原版 Transformer 中 d_k = d_v)

逐步形状:
  Q · Kᵀ        : (n, d_k) × (d_k, n) → (n, n)   注意力分数矩阵
  ÷ √d_k        : (n, n)               → (n, n)   数值缩放
  softmax(行)   : (n, n)               → (n, n)   每行归一,权重 ∈ [0,1],行和 = 1
  × V           : (n, n) × (n, d_v)   → (n, d_v) 加权平均,最终输出

逐符号拆解:

  • Q(query):序列中每个位置发出的"查询",形状 (n, d_k)。第 i 行 Qᵢ 代表位置 i 想要找什么。
  • K(key):序列中每个位置暴露的"键",同样 (n, d_k)。第 j 行 Kⱼ 代表位置 j 愿意被什么样的查询匹配到。
  • QKᵀ:(n, n) 相似度矩阵,第 (i,j) 项 = Qᵢ·Kⱼ,表示位置 i 对位置 j 的未归一化注意力分数。
  • softmax 按行:对第 i 行做 softmax,把位置 i 对所有其他位置的分数归一化成概率分布——这就是"软字典"里的"权重"。
  • ×V:V 的形状是 (n, d_v),第 j 行 Vⱼ 是位置 j 真正携带的信息。加权后,位置 i 的输出是所有位置 value 的加权平均,权重来自刚才的 softmax。
Q (n, d_k) K (n, d_k) Q · Kᵀ (n, n) ÷ √d_k 缩放 (n, n) mask 未来→−∞ (n, n) 可选 softmax 按行归一 (n, n) V (n, d_v) × V 加权平均 (n, d_v)
图 2.1scaled dot-product attention 的完整数据流。Q 和 K 做矩阵乘法得 (n,n) 相似度矩阵,除 √d_k 缩放,因果掩码(decoder 专用,将未来位置置 −∞)可选插入,softmax 按行归一成权重,最后乘 V 得加权平均输出 (n, d_v)。注意:softmax 之前,不是之后,插入掩码——位置不同结果不同。

2.2为什么除以 √d_k

除以 √d_k 是方差控制:d_k 维随机点积的标准差恰好是 √d_k,不除会把 softmax 推进近 one-hot 区,梯度接近零。

从 01 章的软查表直觉出发,softmax 之前的分数越极端,权重就越接近 one-hot(只有一个位置权重接近 1,其余接近 0)。极端分数本身不是问题,但它们是怎么来的?答案在点积的方差统计。

方差论证

假设 Q 和 K 的每个分量独立、零均值、单位方差(这是初始化后合理的近似)。Qᵢ 和 Kⱼ 的点积是 d_k 项之和:

点积方差推导 数学
Qᵢ · Kⱼ = Σₜ₌₁^d_k  Qᵢₜ · Kⱼₜ

每项 Qᵢₜ · Kⱼₜ 的期望 = 0,方差 = Var(Qᵢₜ) · Var(Kⱼₜ) = 1 · 1 = 1

d_k 项独立求和 → 总方差 = d_k · 1 = d_k

→ 标准差 = √d_k

除以 √d_k 后:方差 = d_k / d_k = 1,恢复单位方差

具体数字感受:

  • d_k = 64(原版 Transformer base 配置):点积量级约 ±8,softmax 仍能产生合理的权重分布。
  • d_k = 512(更大的头维度):点积量级约 ±22,softmax 输入中最大值与第二大之差已很悬殊,输出接近 one-hot。
梯度消失的机制

softmax 的梯度是 σ(1−σ),其中 σ 是输出概率。当输出接近 one-hot 时,σ ≈ 1 或 σ ≈ 0,梯度 σ(1−σ) ≈ 0。这不是深层网络的梯度消失,而是在单次 softmax 运算内就发生:注意力权重的梯度信号在训练初期几乎归零,模型学不动。

除以 √d_k 把输入方差拉回 1,softmax 输出保持多峰分布,梯度信号正常传播。

为什么是 √d_k 而不是 d_k

因为控制的是标准差(量级),不是方差。点积量级≈标准差≈√d_k,除以 √d_k 把量级归一,而非除以 d_k(那会过度压缩,让 softmax 输出趋向均匀分布)。

2.3因果掩码:让 decoder 不能偷看未来

因果掩码在 softmax 前把未来位置的分数置为 −∞,使 softmax 后对应权重精确为 0——位置 t 只能看到位置 0…t 的信息。

decoder 在训练时需要同时计算所有位置的输出(并行化,这是 Transformer 相比 RNN 的核心优势之一,第 01 章通信/计算框架里已经给过背景)。但如果位置 t 能看到位置 t+1, t+2, … 的信息,训练时等于"作弊"——模型不需要真正学会预测下一个 token,直接抄答案就行。推理时这些未来 token 根本不存在,训练和推理就彻底不一致了。

具体做法

构造一个下三角矩阵形式的掩码:位置 (i, j) 中,当 j > i 时(位置 i 试图看位置 j > i),在 softmax 前把注意力分数设为 −∞。

因果掩码应用 数学
原始分数矩阵 S = Q · Kᵀ / √d_k    形状 (n, n)

掩码 M(i, j):   M(i, j) = 0   若 j ≤ i(可见)
                M(i, j) = −∞  若 j > i(未来,屏蔽)

应用掩码后:     S_masked(i, j) = S(i, j) + M(i, j)

softmax(S_masked) 的第 (i,j) 项:
  当 j > i:exp(−∞) / Z = 0 / Z = 0  → 权重精确为零
  当 j ≤ i:正常的 softmax 权重       → 对可见位置重新归一
失败模式:在 softmax 后置零

有时会看到把 softmax 输出直接置零的实现。这是错误的:softmax 后置零破坏了"权重之和为 1"的归一化不变量。位置 i 的输出变成所有可见位置 value 的部分加权和,而非完整的加权平均——这相当于让"注意力总量"随可见 token 数量变化,引入位置相关的比例误差。正确做法是在 softmax 前加 −∞,让 softmax 天然地只对可见位置归一。

两个精确性细节

对角线可见:掩码条件是 j > i 才屏蔽,j = i(位置 t 看自己)属于可见范围。位置 t 能且应该关注自身的 key,这是 attention 设计的合理组成部分。

off-by-one 泄露:若掩码条件写成 j ≥ i(把对角线也屏蔽),位置 t 连自身都看不到;若写成 j > i+1,位置 t 会提前看到 t+1。这两种错误都是沉默的——训练不会报错,但模型会学到与推理不一致的分布,表现为推理时质量下降或生成重复。

因果掩码矩阵(n=4) j=0 j=1 j=2 j=3 ← key 位置 j → i=0 i=1 i=2 i=3 ← query 位置 i → 可见 −∞ −∞ −∞ 可见 可见 −∞ −∞ 可见 可见 可见 −∞ 可见 可见 可见 可见 可见(softmax 前保留分数) −∞(softmax 后权重=0)
图 2.34×4 因果掩码矩阵:深色格(下三角含对角线)为位置 i 可见的 key,朱红格(上三角)在 softmax 前被置为 −∞。注意:对角线(i=j)属于可见范围——位置 t 看得见自己;掩码写在 softmax 前,确保权重精确为零而不破坏归一化。

2.4多头注意力:多个表示子空间

单头只能产生一个加权平均(只承诺一种混合方式);多头把 d_model 切成 h 个子空间并行运算,让模型同时学习不同类型的依赖关系——总 FLOPs 近乎不变,表达能力翻倍。

单头 attention 的局限:给定位置 i,softmax 输出的是一个概率分布,最终输出是所有 value 的一个加权平均。这意味着在一次 attention 里,模型只能"承诺"一种混合方式——或者关注句法依赖,或者关注语义相似,或者关注指代关系,无法同时做到。

多头的切分方式

multi-head attention 数学
MultiHead(Q, K, V) = Concat(head₁, head₂, …, headₕ) · Wᴼ

其中:
  headᵢ = Attention(Q · WᵢQ,  K · WᵢK,  V · WᵢV)

投影矩阵形状:
  WᵢQ : (d_model, d_k)   d_k = d_model / h
  WᵢK : (d_model, d_k)
  WᵢV : (d_model, d_v)   d_v = d_model / h  (原版中 d_k = d_v)
  Wᴼ  : (h · d_v, d_model)   将 concat 结果投回 d_model

原版 Transformer base:d_model = 512,h = 8,d_k = d_v = 64
总 FLOPs ≈ 单头 attention(d_model=512)的计算量(各头维度缩小 1/h,h 头并行)

每个头 headᵢ 先用自己的线性投影 (WᵢQ, WᵢK, WᵢV) 把输入投到 d_k 维子空间,在那个子空间里做独立的 scaled dot-product attention,然后 h 个头的输出 concat 拼接成 h·d_v 维,最后通过 Wᴼ 投回 d_model。

为什么说"近乎免费"

单头用 d_model 维做 attention,FLOPs 主要来自 Q·Kᵀ,量级 O(n²·d_model)。切成 h 头后每头只有 d_model/h 维,每头 FLOPs ≈ O(n²·d_model/h),h 头相加恰好还是 O(n²·d_model)。参数量上投影矩阵增加了,但 d_k·h = d_model,整体参数量同阶。代价近乎为零,但表达空间从一个 d_model 维视角扩展到 h 个 d_k 维的独立视角。

induction heads:真实的专门化

多头真的会学到不同功能,而不只是理论上可以。Olsson 等(2022)通过机制可解释性研究发现 induction heads(归纳头):给定模式 [A][B]…[A],一个 head 会 attend 到当前 token [A] 上次出现后面的那个 [B],完成"模式补全"。这是 in-context learning 的基础机制之一。

更关键的是,induction heads 在训练过程中不是渐进形成的,而是相变式突现(类似 grokking 现象):某个训练步之前几乎不存在,某个训练步之后突然出现、与 in-context learning 能力的跃升同步。这说明多头的"专门化"是有真实机制根基的,不只是数学上的自由度。

输入 d_model W₁Q,K,V W₂Q,K,V WₕQ,K,V head₁ d_k = d_model/h head₂ d_k = d_model/h ⋮ headₕ d_k = d_model/h Concat h × d_v = d_model × Wᴼ → d_model 输出 d_model 各头学不同依赖关系 句法 / 语义 / 指代 / 模式补全…
图 2.2多头注意力:输入经 h 组不同线性投影映射到 d_k = d_model/h 维子空间,各头独立运行 scaled dot-product attention,输出 concat 后通过 Wᴼ 投回 d_model。注意:h 头并行运算的总 FLOPs 与单头(维度=d_model)几乎相同,但各头可以自由专门化——induction heads 是真实观测到的专门化案例。

2.5self-attention vs cross-attention

self-attention 的 Q/K/V 来自同一序列,cross-attention 的 Q 来自 decoder、K/V 来自 encoder——差异只在 QKV 的来源,公式形式相同。

scaled dot-product attention 的公式 softmax(QKᵀ/√d_k)·V 对两种形式完全通用,区别只在输入从哪来:

表 2.5 · self-attention vs cross-attention
类型Q 来源K/V 来源典型用途
self-attention 当前序列 当前序列(同 Q) encoder 双向建模;decoder 因果自回归建模
cross-attention decoder 当前状态 encoder 输出(固定,与 decoder 无关) 经典 encoder-decoder(翻译、摘要);decoder 提取 encoder 表示

在第 01 章介绍的通信/计算框架中,self-attention 让序列内每个位置交换信息。cross-attention 的 K/V 来自 encoder 的最终输出,decoder 的每一步都通过它"查询" encoder 对源序列的理解。完整 encoder-decoder 架构的细节在第 03 章覆盖——那里会把 cross-attention 放进完整的 block 堆叠里讲。值得注意的是,现代 decoder-only 模型(GPT、Llama 等)没有 cross-attention:每个 block 只有因果 self-attention + FFN,不存在 encoder 输出可查。

2.6一个反直觉:注意力权重 ≠ 重要度

输出是 value 向量的加权平均,高注意力权重 × 低范数 value = 输出贡献极小;热力图只展示权重,永远看不到 value 的范数——这是读注意力热力图最常见的误解来源。

attention 热力图直觉上像"模型在看哪些词",但它只展示了 softmax(QKᵀ/√d_k) 的输出,即权重矩阵 W。最终输出是 W · V,位置 i 对位置 j 的实际贡献是 W(i,j) × ‖Vⱼ‖ 的方向投影。一个很高的权重 W(i,j) 配上一个几乎为零范数的 Vⱼ,对输出的贡献趋近于零;相反,一个中等权重配上高范数 value,可能主导输出。

想一想

在 attention 热力图中,位置 i 对位置 j 的权重是 0.9(极高),位置 i 对位置 k 的权重是 0.1(很低)。能否断言位置 j 对位置 i 输出的影响远大于位置 k?为什么?

展开答案(先停 10 秒再点)

不能断言。输出是加权平均 W · V,位置 j 对输出的实际贡献量级约为 W(i,j) × ‖Vⱼ‖ = 0.9 × ‖Vⱼ‖,位置 k 的约为 0.1 × ‖Vₖ‖。若 ‖Vₖ‖ ≫ ‖Vⱼ‖(例如 Vₖ 范数是 Vⱼ 的 10 倍以上),位置 k 的贡献反而更大。

权重矩阵只是注意力分布,value 的范数是缺失的另一半。热力图误导的根源恰在于此——它展示的是"模型把注意力集中到哪",但注意力集中到的位置的 value 范数小时,该位置对实际输出几乎没有影响。这是 Jain & Wallace(2019,arxiv 2004.10102 的先驱工作)以及后续研究反复验证的结论。

失败模式:用热力图做因果解释

把注意力权重当作"模型决策依据",用它来解释"为什么模型给出了这个输出",是对 attention 机制的误用。权重只描述了 key-query 相似度分布,不描述 value 对输出的实际影响量。正确的可解释性方法需要同时考虑 value 的范数,或直接用梯度/积分梯度等输出相关方法。

§本章 self-check

先合上教程,把能想到的答案写在纸上或编辑器里。写完再对照——直接点开等于把这一节当再读一遍。

  1. Attention(Q,K,V) = softmax(QKᵀ/√d_k)·V,若 n=128、d_k=64、d_v=64,写出 Q、K、V、QKᵀ、最终输出的形状。
  2. 为什么 d_k=512 时不除 √d_k 会导致梯度消失?从点积方差到 softmax 导数,用三句话讲清楚。
  3. 因果掩码为什么必须在 softmax 前置 −∞ 而不是在 softmax 后置零?两种做法的数学结果有何不同?
  4. 注意力热力图上某个词的权重接近 1,是否说明该词对模型输出"很重要"?缺少什么信息才能做出完整判断?
答案(先做完再展开)
  1. Q: (128,64);K: (128,64);V: (128,64);QKᵀ: (128,128);最终输出: (128,64)。每步矩阵乘法的形状规则:(n,d_k)×(d_k,n)→(n,n),(n,n)×(n,d_v)→(n,d_v)。
  2. ①d_k=512 时,零均值单位方差的 d_k 维点积方差=512,标准差=√512≈22.6,点积量级约±22。②softmax 输入出现极端大值时,最大项的 exp 值远大于其他项,输出接近 one-hot。③softmax 的梯度 σ(1−σ):当 σ≈1 时导数≈0,当 σ≈0 时导数也≈0——梯度在 softmax 内就消失了,模型无法从 attention 层学到任何有用信号。
  3. softmax 前置 −∞:exp(−∞)=0,该位置在归一化分母中贡献 0,其余可见位置的权重重新对自身归一,行和始终=1。softmax 后置零:强制让某些权重变 0,但分母已经包含了这些位置的 exp 值(分母被高估),剩余权重的行和 < 1——输出是所有可见位置 value 的"部分加权",不再是完整的加权平均,引入与可见 token 数量相关的比例误差。
  4. 不能。输出是加权平均 W·V,该词对输出的实际贡献 = 权重 × ‖V‖(value 范数)× 方向对齐程度。缺少该词的 value 向量范数信息,无法判断实际影响大小。