注意力机制中缩放点积除以根号 d_k 的方差稳定化推导
注意力机制的核心公式为 Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V。其中除以 sqrt(d_k) 不是一个随意选择,而是一个保持方差稳定的关键操作。本文将从零推导这一缩放因子的必要性,并展示它如何避免梯度消失/爆炸。
1. 明确问题:点积的方差失控
假设查询向量 q 和键向量 k 的每个元素独立同分布,均值为 0,方差为 1(常见标准化结果)。即:
$$q_i \sim \mathcal{N}(0, 1), \quad k_i \sim \mathcal{N}(0, 1), \quad i = 1, 2, \dots, d_k$$
那么两个向量的点积为:
$$a = q \cdot k = \sum_{i=1}^{d_k} q_i k_i$$
因为 q_i 和 k_i 独立且均值为 0,所以:
- 期望:$E[a] = \sum E[q_i]E[k_i] = 0$
- 方差:每个乘积项
q_i k_i的方差为:
$$Var(q_i k_i) = E[q_i^2]E[k_i^2] - (E[q_i]E[k_i])^2 = (1)(1) - 0 = 1$$
由于各项独立,和的方差为各项方差之和:
$$Var(a) = \sum_{i=1}^{d_k} 1 = d_k$$
因此,点积 a 的方差等于向量维度 d_k。当 d_k 很大(如 512、1024)时,点积的绝对值也会很大,导致 softmax 的输入跨度极大,使 softmax 输出趋向于 one-hot 形式(一个接近 1,其余接近 0),从而梯度变得极小(饱和区)。
2. 标准差缩放:将方差稳定到 1
为了消除维度影响,我们除以 sqrt(d_k),使缩放后的点积方差恢复为 1:
$$a' = \frac{a}{\sqrt{d_k}}$$
计算方差:
$$Var(a') = Var\left(\frac{a}{\sqrt{d_k}}\right) = \frac{Var(a)}{d_k} = \frac{d_k}{d_k} = 1$$
此时 a' 的方差与 d_k 无关,保持在 1。softmax 的输入分布不会随维度膨胀,从而梯度更稳定。
3. 严格推导:从概率分布到方差归一化(可选深度理解)
若你想从更基础的随机变量角度验证,可以按以下步骤手算:
3.1 计算单个乘积的方差
设 X = q_i k_i,由于 q_i 和 k_i 独立且服从标准正态分布:
$$E[X] = E[q_i]E[k_i] = 0 \times 0 = 0$$
$$E[X^2] = E[q_i^2]E[k_i^2] = 1 \times 1 = 1$$
所以:
$$Var(X) = E[X^2] - (E[X])^2 = 1 - 0 = 1$$
3.2 计算和的方差
对于独立随机变量,和的方差等于方差之和:
$$Var\left(\sum_{i=1}^{d_k} X_i\right) = \sum_{i=1}^{d_k} Var(X_i) = d_k \times 1 = d_k$$
3.3 除以 sqrt(d_k) 后的方差
缩放因子 1 / sqrt(d_k) 的平方为 1 / d_k,所以:
$$Var\left(\frac{a}{\sqrt{d_k}}\right) = \frac{1}{d_k} Var(a) = \frac{1}{d_k} \cdot d_k = 1$$
推导完成。方差稳定为 1。
4. 扩展到实际应用:为什么这对训练至关重要
如果不进行缩放,softmax 输入 [a_1, a_2, ...] 的方差为 d_k。当 d_k 较大时,最大值与最小值差距很大,softmax 输出趋向于 0 或 1(极端分布)。反向传播时,softmax 的梯度在饱和区域几乎为 0,导致梯度消失。
缩放后,softmax 输入方差为 1,处于相对非饱和区域(softmax 函数在输入接近 0 时导数较大),梯度流动更顺畅,模型训练更快更稳定。
5. 实现验证:简单 Python 模拟
你可以用以下代码亲自验证方差变化:
import numpy as np
d_k = 128
num_samples = 100000
# 生成随机 q 和 k (标准正态)
q = np.random.randn(num_samples, d_k)
k = np.random.randn(num_samples, d_k)
# 计算点积
a = np.sum(q * k, axis=1) # shape: (num_samples,)
print("缩放前方差:", np.var(a)) # 约等于 128
# 除以 sqrt(d_k)
a_scaled = a / np.sqrt(d_k)
print("缩放后方差:", np.var(a_scaled)) # 约等于 1
运行结果会显示缩放前方差接近 d_k,缩放后方差接近 1。
6. 最终结论
缩放因子 1 / sqrt(d_k) 使得注意力分数(点积)的方差与维度无关,始终为 1。这一设计:
- 防止 softmax 输入过于极端(避免梯度饱和)
- 保持梯度尺度稳定,有利于深层网络训练
- 在数学上完全可推导,非经验猜测
当你实现或使用注意力机制时,永远不要省略这个除法。

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