从零推导 Self-Attention:为什么是 QKV,为什么要除以 √d
关于自注意力的教程已经太多了,但绝大多数在讲完「Q 乘 K 的转置再 softmax」之后就停下了。真正让人卡住的两个问题反而没人回答:为什么要把同一个输入投影成三份?那个 √d 到底是怎么来的?这篇把推导拆开重走一遍。
一、先不要有 QKV
假设我们完全不知道 Transformer 长什么样,只有一个朴素需求:给定一串词向量 x₁…xₙ,希望每个位置的新表示能够"看到"其他位置的信息,并且看多看少由内容本身决定,而不是由固定的位置权重决定。
最直接的写法是让每个位置对其他位置打一个分,然后按分数加权平均:
score 最省事的实现就是内积:score = xᵢ · xⱼ。到这里已经有一个能跑的注意力了,而且没有任何可学参数。问题也随之而来。
问题 1:内积是对称的
xᵢ·xⱼ 恒等于 xⱼ·xᵢ。这意味着「"的" 关注 "苹果"」和「"苹果" 关注 "的"」被迫得到同一个分数。但语言里的依赖关系显然是有方向的:形容词要去找它修饰的名词,代词要去找先行词,反过来的强度完全不同。
解决办法是在打分前对两侧各做一次不同的线性变换:
由于 WQ ≠ WK,这个打分函数不再对称。Q 和 K 的分工到此就解释完了:不是什么"查询和键"的数据库隐喻,而是为了打破内积的对称性。查询/键这个命名是事后的直觉说明,不是设计动机。
问题 2:用来打分的向量,不一定适合被搬运
在上面的写法里,加权求和的对象仍然是原始的 xⱼ。但"判断两个词相不相关"和"从一个词身上取走什么信息"是两件不同的任务。前者可能只需要词性和句法角色,后者需要的是语义内容。让同一个向量同时承担两件事,等于强行把两个子空间压在一起。
所以再加一个投影,专门负责被搬运的内容:
三个矩阵各自的职责就清楚了:WQ 决定「我在找什么」,WK 决定「我能被什么找到」,WV 决定「我被找到之后交出什么」。
二、√d 从哪来
现在处理缩放因子。设 q 和 k 是两个 d 维向量,各分量独立同分布,均值 0、方差 1。它们的内积是:
每一项 qtkt 的期望是 0,方差是 E[q²]E[k²] = 1。d 个独立项相加,方差线性累加,得到 Var(q·k) = d,标准差为 √d。
也就是说,维度越高,内积的取值范围就越宽。d = 128 时,打分的典型量级已经到了 ±11 上下;d = 512 时接近 ±23。把这样的数值送进 softmax 会发生什么?
import numpy as np
def softmax(z):
z = z - z.max()
e = np.exp(z)
return e / e.sum()
logits = np.array([2.0, 1.0, 0.5, 0.0])
print(softmax(logits)) # [0.55 0.20 0.12 0.07] 分布平缓
print(softmax(logits * 11)) # [1.00 0.00 0.00 0.00] 几乎独热
当 softmax 输出接近独热向量时,它对输入的梯度趋近于零——因为 ∂softmax/∂z 的雅可比矩阵是 diag(p) − ppT,p 越接近独热,这个矩阵越接近零矩阵。注意力层在训练早期就会陷入梯度消失,权重几乎不更新。
把内积除以 √d,方差重新归一到 1,打分回到 softmax 的敏感区间。所以缩放因子不是经验超参,而是为了让内积的方差与维度无关而做的方差归一化。同理,如果实现中用了不同的初始化方差 σ²,正确的缩放因子应该是 σ²√d。
三、因果掩码的实现细节
解码器需要保证位置 i 只能看到 ≤ i 的位置。常见做法是把上三角部分加上一个很大的负数,让 softmax 之后接近 0:
scores = q @ k.transpose(-2, -1) / math.sqrt(d_head)
mask = torch.triu(torch.ones(n, n, dtype=torch.bool), diagonal=1)
scores = scores.masked_fill(mask, float("-inf"))
attn = scores.softmax(dim=-1)
这里有个实践中真踩过的坑:如果用 -1e9 而不是 -inf,在 FP16 下会直接溢出成 -inf 还好,但在 BF16 下由于指数位宽足够、尾数位少,-1e9 会被舍入到一个仍然参与计算的有限值。当一行里所有位置都被掩掉(例如 padding 行)时,softmax 分母是 n·exp(-1e9 − max),结果是 NaN 而不是 0,NaN 会顺着残差连接污染整个网络。
assert torch.isfinite(attn).all() 定位到具体层。四、多头是在做什么
把 d_model 切成 h 份,每份独立算注意力,最后拼接再过一个输出投影。参数量和单头几乎相同(因为每个头的维度是 d_model/h),但表达能力不同。
单头注意力的一次输出是值向量的凸组合,本质上是一个秩受限的操作:无论怎么调 softmax,yᵢ 都落在 {vⱼ} 张成的凸包内。多头允许在不同子空间里用不同的关注模式,然后线性组合,可以合成凸包之外的表示。
实证上,训练好的模型里能观察到相当稳定的头分工:
| 头的类型 | 典型行为 | 出现层 |
|---|---|---|
| 位置头 | 几乎只关注前 1–2 个 token | 浅层(1–4) |
| 句法头 | 动词关注主语,介词关注宾语 | 中层(5–16) |
| 归纳头 | 找到重复模式的前一次出现并复制其后继 | 中后层 |
| 汇聚头(sink) | 大部分权重压在序列第一个 token 上 | 各层普遍存在 |
最后一类值得单独说:注意力汇聚现象在几乎所有解码器模型里都存在,第一个 token 会吸走大量权重。一个被广泛接受的解释是,softmax 强制权重和为 1,当某个位置不需要从任何地方取信息时,它必须把权重倾倒到某处,第一个 token 就成了默认垃圾桶。这也是流式推理里不能随便丢弃开头几个 KV 的原因。
五、复杂度与后续
朴素实现的时间和显存都是 O(n²·d)。显存那一项才是真正的瓶颈:n = 8192 时,单头的注意力矩阵就是 8192² × 2 bytes ≈ 128 MB,32 个头再乘上批大小,立刻爆掉。
FlashAttention 的核心思路是根本不物化这个矩阵:把 Q、K、V 分块载入 SRAM,用在线 softmax(维护 running max 和 running sum)分块累加输出。计算量没变,HBM 读写量从 O(n²) 降到 O(n²/M),M 是片上缓存能放下的块大小。实际收益主要来自访存,不是 FLOPs。
下一篇打算写分组查询注意力在长上下文下对 KV Cache 的影响,那部分和推理服务的并发上限直接相关。