考察点
这是 Transformer 面试里最高频的一题,背公式谁都会,面试官真正想听的是:你能不能把 Q、K、V 的直觉讲清楚,能不能从数学上解释为什么要除以 √d_k,以及是否知道注意力的复杂度瓶颈和工程优化(FlashAttention)。追问常往 softmax 数值稳定性、多头的设计动机、注意力复杂度优化方向走。
参考答案
Q、K、V 是什么
自注意力的输入是一个 token 序列,每个 token 的 embedding 向量 x 分别乘上三个可学习矩阵 W_Q、W_K、W_V,得到 query、key、value 三个向量。直觉上:query 表示「我在找什么」,key 表示「我能提供什么信息」,value 表示「我实际携带的内容」。
计算一个 token 的输出时,用它的 query 去和序列里所有 token 的 key 做点积,得到一组相关性分数,softmax 归一化成权重,再对所有 token 的 value 加权求和。这样每个位置的输出都融合了全序列的信息,权重是数据驱动学出来的,而不是像 CNN 那样固定的局部卷积核,也不像 RNN 那样靠逐步传递。
完整公式与缩放的原因
$$Attention(Q, K, V) = softmax\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$
Q 是 n×d_k 矩阵(n 个 token,每个 d_k 维),QK^T 得到 n×n 的分数矩阵,缩放后按行 softmax,最后乘 V 得到 n×d_v 的输出。
为什么除以 √d_k:假设 q 和 k 的各分量是独立、均值为 0、方差为 1 的随机变量,那么点积 q·k = Σ q_i·k_i 是 d_k 个独立项之和,均值为 0,但方差是 d_k,也就是点积的量级随 √d_k 增长。d_k = 64 时,点积的标准差大约是 8;d_k 更大时分数会更大。
问题在于 softmax 对大的输入会饱和:分数差距被指数放大,输出的分布接近 one-hot,梯度趋近于 0(softmax 饱和区的导数几乎为零),训练早期就学不动。除以 √d_k 把点积的方差拉回 1 附近,让 softmax 工作在输入有区分度但梯度又健康的区间。这不是拍脑袋的经验值,是把方差归一的标准操作。
多头注意力
与其用一个 d 维的大注意力,不如把 d 拆成 h 个头,每个头独立做 d/h 维的注意力,再把结果拼接后过一层线性变换 W_O。每个头在不同的子空间里学不同的关注模式——有的头关注语法邻近,有的头关注指代关系。典型配置如 7B 模型 32 个头、d_model 4096,每个头 128 维。
复杂度与工程优化
标准注意力的计算量和显存占用都是 O(n²),n 是序列长度。8K 上下文的 attention 矩阵就有 6400 万个元素,这是长文本的主要瓶颈。工程上的应对:
- FlashAttention:不改变数学结果,通过分块(tiling)计算和在线 softmax,避免把 n×n 矩阵物化到显存,把显存从 O(n²) 降到 O(n),同时靠减少 HBM 读写拿到 2-4 倍实际加速。现在几乎是训练标配。
- KV cache:推理时缓存历史 token 的 K、V,每步只计算新 token 的注意力,把自回归生成的每步复杂度从 O(n²) 降到 O(n)。
- 稀疏注意力、滑动窗口注意力(Mistral)、GQA/MQA 等变体,分别从计算模式和 KV 缓存体积上做文章。
数值稳定性细节
实际实现里 softmax 前会减去该行最大值(safe softmax),防止 exp 溢出。混合精度训练时 attention score 的计算和 softmax 通常在 fp32 里做,否则 fp16 的精度损失会明显影响效果。
可能的追问
- 如果 d_k 很大不缩放会怎样? softmax 输入量级过大进入饱和区,输出接近 one-hot,梯度消失,训练初期 loss 几乎不降。
- 为什么多头而不是单头加大维度? 多头在多个子空间并行学不同模式,表达能力更强;计算量与单头基本相当(总维度不变)。
- FlashAttention 改变了注意力数学吗? 没有,结果是精确等价的,只是用分块和在线 softmax 重排了计算顺序,省的是显存读写。