文章目录

注意力机制中缩放点积除以根号dk的方差稳定化推导

发布于 2026-07-28 08:46:39 · 浏览 23 次 · 评论 0 条

注意力机制中缩放点积除以根号 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_ik_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_ik_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 输入过于极端(避免梯度饱和)
  • 保持梯度尺度稳定,有利于深层网络训练
  • 在数学上完全可推导,非经验猜测

当你实现或使用注意力机制时,永远不要省略这个除法

评论 (0)

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

扫一扫,手机查看

扫描上方二维码,在手机上查看本文