Tensor NotesNotes on machine learning
Transformer

从零推导 Self-Attention:为什么是 QKV,为什么要除以 √d

从零推导 Self-Attention:为什么是 QKV,为什么要除以 √d题图

关于自注意力的教程已经太多了,但绝大多数在讲完「Q 乘 K 的转置再 softmax」之后就停下了。真正让人卡住的两个问题反而没人回答:为什么要把同一个输入投影成三份?那个 √d 到底是怎么来的?这篇把推导拆开重走一遍。

一、先不要有 QKV

假设我们完全不知道 Transformer 长什么样,只有一个朴素需求:给定一串词向量 x₁…xₙ,希望每个位置的新表示能够"看到"其他位置的信息,并且看多看少由内容本身决定,而不是由固定的位置权重决定。

最直接的写法是让每个位置对其他位置打一个分,然后按分数加权平均:

yᵢ = Σⱼ softmax( score(xᵢ, xⱼ) )ⱼ · xⱼ

score 最省事的实现就是内积:score = xᵢ · xⱼ。到这里已经有一个能跑的注意力了,而且没有任何可学参数。问题也随之而来。

问题 1:内积是对称的

xᵢ·xⱼ 恒等于 xⱼ·xᵢ。这意味着「"的" 关注 "苹果"」和「"苹果" 关注 "的"」被迫得到同一个分数。但语言里的依赖关系显然是有方向的:形容词要去找它修饰的名词,代词要去找先行词,反过来的强度完全不同。

解决办法是在打分前对两侧各做一次不同的线性变换:

score(xᵢ, xⱼ) = (WQ xᵢ) · (WK xⱼ)

由于 WQ ≠ WK,这个打分函数不再对称。Q 和 K 的分工到此就解释完了:不是什么"查询和键"的数据库隐喻,而是为了打破内积的对称性。查询/键这个命名是事后的直觉说明,不是设计动机。

问题 2:用来打分的向量,不一定适合被搬运

在上面的写法里,加权求和的对象仍然是原始的 xⱼ。但"判断两个词相不相关"和"从一个词身上取走什么信息"是两件不同的任务。前者可能只需要词性和句法角色,后者需要的是语义内容。让同一个向量同时承担两件事,等于强行把两个子空间压在一起。

所以再加一个投影,专门负责被搬运的内容:

yᵢ = Σⱼ αᵢⱼ · (WV xⱼ)

三个矩阵各自的职责就清楚了:WQ 决定「我在找什么」,WK 决定「我能被什么找到」,WV 决定「我被找到之后交出什么」。

小结QKV 不是三个平行的概念,而是两次独立设计决策的产物:Q/K 分离是为了打破对称性,V 分离是为了解耦"匹配空间"与"内容空间"。
解码器层结构:RMSNorm、分组查询注意力、SwiGLU 前馈与残差连接
图 1 · 现代解码器层的一个典型配置(前置归一化 + GQA + SwiGLU)

二、√d 从哪来

现在处理缩放因子。设 q 和 k 是两个 d 维向量,各分量独立同分布,均值 0、方差 1。它们的内积是:

q · k = Σt=1..d qt kt

每一项 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 会顺着残差连接污染整个网络。

排查建议训练中途出现 NaN 且梯度范数在前一步还正常,优先怀疑掩码与 padding 的交互,而不是学习率。可以在 forward 里临时插一个 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 的影响,那部分和推理服务的并发上限直接相关。

上一篇:MoE 稀疏路由的工程真相:负载均衡、容量因…

继续阅读

Related