精选·模型与微调基础

Self-Attention 完整计算过程:为什么 Softmax 前要除以根号 d_k

91学AI·2026/7/13·5 阅读

考察点

这是八股里的经典题,但面试官用它筛两种人:背公式的和真理解的。期望你能白板推出完整计算流,并用「数值稳定」而不是「防止过大」这种模糊话术解释缩放因子。追问常考:多头的作用、复杂度分析、FlashAttention 为什么快。

参考答案

完整计算链路

输入是 n 个 token 的向量序列。每个向量乘三个可学习矩阵,得到 Query、Key、Value 三组向量。注意力的输出分四步:

  1. 算相关性分数:每个位置的 Q 和所有位置的 K 做点积,得到 n×n 的分数矩阵,含义是「每个 token 该关注其他 token 多少」。
  2. 缩放:分数除以根号 d_k(d_k 是 Key 向量的维度)。
  3. 掩码加 Softmax:Decoder 里先把下三角以外的位置置为负无穷(因果掩码,防止偷看未来),再做 Softmax 把每行归一化成和为 1 的权重。
  4. 加权求和:权重矩阵乘 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。

评论 (0)

暂无评论,快来抢沙发吧!

91学AI

© 2026 91学AI · 按岗位学 AI 与大数据. All rights reserved.