考察点
这是八股里的经典题,但面试官用它筛两种人:背公式的和真理解的。期望你能白板推出完整计算流,并用「数值稳定」而不是「防止过大」这种模糊话术解释缩放因子。追问常考:多头的作用、复杂度分析、FlashAttention 为什么快。
参考答案
完整计算链路
输入是 n 个 token 的向量序列。每个向量乘三个可学习矩阵,得到 Query、Key、Value 三组向量。注意力的输出分四步:
- 算相关性分数:每个位置的 Q 和所有位置的 K 做点积,得到 n×n 的分数矩阵,含义是「每个 token 该关注其他 token 多少」。
- 缩放:分数除以根号 d_k(d_k 是 Key 向量的维度)。
- 掩码加 Softmax:Decoder 里先把下三角以外的位置置为负无穷(因果掩码,防止偷看未来),再做 Softmax 把每行归一化成和为 1 的权重。
- 加权求和:权重矩阵乘 Value 矩阵,每个位置得到一个「按相关度聚合了全序列信息」的新向量。
多头(Multi-Head)就是把向量切成 h 份,每份独立跑上面这套流程,最后拼起来再过一层线性变换。切头的意义是让不同子空间学不同的关注模式,相当于把一种注意力变成 h 路并行的小注意力,实践证明比单个大头效果好。
为什么除以根号 dk
这是本题的题眼。假设 Q 和 K 的每个分量是均值为 0、方差为 1 的独立随机变量,那么它们点积的方差约等于 d_k——维度越高,点积结果的绝对值越大,且随维度线性放大。
Softmax 这个函数有个脾气:输入值越大,输出越接近 one-hot,梯度越接近 0(饱和区)。d_k 常见取值 64 或 128,不缩放时点积动不动几十上百,Softmax 几乎把全部权重压到一个位置上,反向传播时梯度消失,训练前期就卡住。除以根号 d_k 之后,点积方差回到 1 附近,Softmax 工作在输入适中、梯度健康的区间。一句话:缩放因子是让 Softmax 的输入尺度不随模型维度变化的归一化手段。
复杂度与工程推论
注意力的时间和显存复杂度都是 O(n²·d)——n 是序列长度。这直接产生三条应用工程师必须知道的推论:
- 长上下文很贵。序列翻倍,注意力计算量翻四倍。128K 上下文的模型,就算支持,单次推理成本也远高于短上下文,业务上能裁剪就别灌全文。
- KV Cache 是推理加速的地基。自回归生成第 t 个 token 时,前面 t-1 个位置的 K、V 算过了就不用重算,缓存起来每步只算新 token 的一行注意力。代价是显存:KV Cache 随 batch 和序列长度线性涨,vLLM 的 PagedAttention 就是为了管这块显存。
- FlashAttention 省的是显存带宽不是计算量。它没有改变 O(n²) 的数学复杂度,而是通过分块计算避免把 n×n 的中间矩阵写回显存,IO 少了就快了好几倍,还顺便让长序列训得起。面试里别把这两点混着说。
一个容易忽略的点
注意力输出的质量不只取决于公式,还取决于 Q/K/V 三个投影矩阵学到了什么。工程上做微调时,LoRA 最常见的注入点就是 Q 和 V 的投影层——因为它们直接决定「关注谁」和「聚合什么」,是注意力机制里可调杠杆最大的位置。
可能的追问
- 为什么用点积而不是加性注意力?点积可以全部落成矩阵乘法,GPU 上高度优化;加性注意力多一层前馈,表达力略强但计算慢,规模上去后点积配合缩放效果足够好。
- Softmax 之前还有什么 mask?除了因果掩码,还有 padding mask(补齐的占位符不能参与注意力)和交叉注意力里的源序列掩码,面试里说出前两个就够。
- 注意力分数异常集中或完全平均说明什么?过度集中可能是温度/缩放异常或训练早期现象;完全平均的头常见,是模型学到的「兜底头」,不一定是 bug。