2.2 缩放点积注意力:为什么要除以根号 d
缩放点积注意力(Scaled Dot-Product Attention)是 Transformer 中最基础的注意力计算单元。它的公式看起来很简洁,但每一项都有明确的设计理由。
上一节把单头的 定在 64 或 128,那是 按头数切分的结果;Q、K 必须同维,也只是因为两者要做点积。但 一旦定下来就不再只是形状约束——它还会出现在下面这个公式的分母里,而且是以平方根的形式。
2.2.1 完整公式与计算流程
缩放点积注意力的计算公式为:
整个计算可以分解为四个步骤:
计算注意力分数:,结果矩阵 中的每个元素 是第 个查询与第 个键的点积,衡量两者的“匹配度”
缩放:,将分数除以键向量维度的平方根,保持点积的方差为常数 1
归一化:,对每一行应用 Softmax,得到注意力权重矩阵,每行的权重之和为 1
加权求和:,用注意力权重对值向量进行加权求和
[!NOTE] 行向量与矩阵打包:按照《Attention Is All You Need》原论文的设定与深度学习的实现惯例,单个词的查询向量 被视作行向量(维度为 )。序列中 个词的查询向量被逐行“打包”(packed together)拼接成完整的查询矩阵 。这使得矩阵乘法 可以一次性并行计算出所有词对之间的点积分数。
整个过程可以完全用矩阵运算表达,因此能够在 GPU 上高效执行。
2.2.2 点积为什么能衡量相关性
点积 本质上衡量的是两个向量的方向一致性。根据向量的几何关系:
其中 是两个向量之间的夹角。当两个向量方向一致时(),点积最大;方向正交时(),点积为零;方向相反时(),点积最小。
这种几何属性使得点积成为衡量语义相关性的天然工具:经过投影后,语义相关的查询和键应该在向量空间中方向接近,从而产生较大的点积分数。
相比 1.3 节中介绍的加性注意力(需要额外的权重矩阵和 tanh 激活),点积注意力不需要额外的可学习参数,且可以直接利用优化过的矩阵乘法库(如 cuBLAS),在实践中快得多。这也是 Transformer 选择点积而非加性注意力的主要原因。
2.2.3 为什么必须缩放:Softmax 的饱和问题
现在来到这个公式中最关键的设计细节:为什么要除以 ?
这是因为当维度 较大时,点积的数值也会变大,导致 Softmax 函数进入饱和区。
Softmax 函数 对输入的尺度非常敏感。当输入值很大时,指数函数会使最大值对应的输出接近 1,其余接近 0——这就是“饱和(saturation)”。在饱和区,Softmax 的梯度趋近于零,导致反向传播时梯度消失,模型难以学习。
顺带一提工程实现中的一个经典技巧:Softmax 对输入整体平移不变(),因此实际计算时会先按行减去最大值再取指数(),避免 数值上溢——PyTorch 等框架的 Softmax 内部默认就这样做。注意它与缩放因子解决的是两个不同的问题:减最大值防的是数值溢出,除以 防的是梯度消失。
假设查询和键的每个分量都独立地服从标准正态分布 ,那么它们的点积:
是 个独立随机变量乘积的和。根据概率论:
均值:(因为 和 独立,)
方差:(因为当 时, 的方差为 1,共 项独立)
因此,点积的标准差为 。当 时,标准差为 8;当 时,标准差约为 22.6。这意味着点积值的数值范围和量级会随维度增长而扩大——典型的点积值(约在均值周围 范围内)会变得越来越大。
下表直观展示了不同 下,未缩放点积的方差增长。Softmax 输出和梯度并不由 单独决定,还取决于整行 logits、序列长度和分数分布;这里保留方差和典型尺度,避免把某个 toy logits 设置误读成通用数值规律。
点积方差
点积标准差
典型点积范围()
缩放后方差
16
16
4.0
1
64
64
8.0
1
128
128
11.3
1
512
512
22.6
1
[!NOTE] 关键观察:未缩放时, 增大会使点积分数的尺度扩大,Softmax 更容易进入饱和区域,梯度变小。而除以 后(最右列),方差恢复为常数 1,使 Softmax 更可能保持在健康的梯度区域工作。这就是为什么缩放因子对训练稳定性至关重要。
除以 正是将点积的理论标准差重新缩放为 1(实际模型中约为常数级 1),使 Softmax 保持在梯度较大的“活跃区”,从而确保训练的稳定性。
这不是一个随意的超参数选择,而是在常见初始化假设下可以由方差分析解释的稳定化设计。
2.2.4 Softmax 的作用:从分数到概率
Softmax 在注意力计算中扮演着双重角色:
第一,归一化:将任意实数范围的注意力分数转换为非负且总和为 1 的权重,使其可以解释为“概率分布”——每个位置被关注的程度。
第二,竞争机制:Softmax 的指数特性使得较大的分数被进一步放大,较小的分数被进一步压缩。这创造了一种赢者通吃的效果,使模型能够将注意力集中在最相关的少数位置上,而非均匀分散。
下面的示例直观说明了这一效果。假设某个查询对五个键的点积分数为 ,经过 Softmax 后变为:
1
2.0
0.606
2
1.0
0.223
3
0.1
0.091
4
-0.5
0.050
5
-1.0
0.030
可以看到,分数最高的位置 1 获得了约 60.6% 的注意力,而分数较低的位置几乎被忽略。这正是注意力机制“选择性关注”的体现。
2.2.5 完整的计算示例
以一个简化的例子来完整展示缩放点积注意力的计算过程。假设序列长度 ,维度 :
计算点积 :
缩放(除以 ):
Softmax 归一化(按行操作,得到注意力权重):
加权求和计算输出 :
每个输出位置的向量是所有位置值向量的加权组合,权重由查询和键的匹配度决定。在一般形式中,注意力输出维度是 ;在标准 Transformer 实现里,多头拼接后再通过输出投影回到 ,因此可以直接传递给下一层。
下面用 PyTorch 代码完整实现上述计算过程,使用与正文相同的 Q、K 矩阵,读者可以直接运行验证每一步的结果:
运行上述代码后,可以得到如下结果。
把这个权重矩阵绘制为热力图(heatmap),从颜色的深浅可以直观感受每个查询对各键的关注程度。

图 2-1:缩放点积注意力权重热力图(生成脚本)
在热力图中,颜色越深的单元格表示注意力权重越高——即该查询位置更“关注”对应的键位置。读者可以观察到,每一行(每个查询)的权重之和为 1,且不同查询关注的位置分布差异明显,这正是注意力机制“选择性关注”能力的直观体现。
最后更新于
