For the complete documentation index, see llms.txt. This page is also available as Markdown.

2.2 缩放点积注意力:为什么要除以根号 d

缩放点积注意力(Scaled Dot-Product Attention)是 Transformer 中最基础的注意力计算单元。它的公式看起来很简洁,但每一项都有明确的设计理由。

上一节把单头的 dkd_k 定在 64 或 128,那是 dmodeld_{\text{model}} 按头数切分的结果;Q、K 必须同维,也只是因为两者要做点积。但 dkd_k 一旦定下来就不再只是形状约束——它还会出现在下面这个公式的分母里,而且是以平方根的形式。

2.2.1 完整公式与计算流程

缩放点积注意力的计算公式为:

Attention(Q,K,V)=softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V

整个计算可以分解为四个步骤:

  1. 计算注意力分数S=QKTS = QK^T,结果矩阵 SRn×nS \in \mathbb{R}^{n \times n} 中的每个元素 SijS_{ij} 是第 ii 个查询与第 jj 个键的点积,衡量两者的“匹配度”

  2. 缩放S=S/dkS' = S / \sqrt{d_k},将分数除以键向量维度的平方根,保持点积的方差为常数 1

  3. 归一化A=softmax(S)A = \text{softmax}(S'),对每一行应用 Softmax,得到注意力权重矩阵,每行的权重之和为 1

  4. 加权求和Output=AV\text{Output} = AV,用注意力权重对值向量进行加权求和

[!NOTE] 行向量与矩阵打包:按照《Attention Is All You Need》原论文的设定与深度学习的实现惯例,单个词的查询向量 q\vec{q} 被视作行向量(维度为 1×dk1 \times d_k)。序列中 nn 个词的查询向量被逐行“打包”(packed together)拼接成完整的查询矩阵 QRn×dkQ \in \mathbb{R}^{n \times d_k}。这使得矩阵乘法 QKTQK^T 可以一次性并行计算出所有词对之间的点积分数。

整个过程可以完全用矩阵运算表达,因此能够在 GPU 上高效执行。

2.2.2 点积为什么能衡量相关性

点积 qk=i=1dkqikiq \cdot k = \sum_{i=1}^{d_k} q_i k_i 本质上衡量的是两个向量的方向一致性。根据向量的几何关系:

qk=qkcosθq \cdot k = \|q\| \|k\| \cos\theta

其中 θ\theta 是两个向量之间的夹角。当两个向量方向一致时(θ=0\theta = 0),点积最大;方向正交时(θ=90°\theta = 90°),点积为零;方向相反时(θ=180°\theta = 180°),点积最小。

这种几何属性使得点积成为衡量语义相关性的天然工具:经过投影后,语义相关的查询和键应该在向量空间中方向接近,从而产生较大的点积分数。

相比 1.3 节中介绍的加性注意力(需要额外的权重矩阵和 tanh 激活),点积注意力不需要额外的可学习参数,且可以直接利用优化过的矩阵乘法库(如 cuBLAS),在实践中快得多。这也是 Transformer 选择点积而非加性注意力的主要原因。

2.2.3 为什么必须缩放:Softmax 的饱和问题

现在来到这个公式中最关键的设计细节:为什么要除以 dk\sqrt{d_k}

这是因为当维度 dkd_k 较大时,点积的数值也会变大,导致 Softmax 函数进入饱和区。

Softmax 函数 softmax(zi)=ezi/jezj\text{softmax}(z_i) = e^{z_i} / \sum_j e^{z_j} 对输入的尺度非常敏感。当输入值很大时,指数函数会使最大值对应的输出接近 1,其余接近 0——这就是“饱和(saturation)”。在饱和区,Softmax 的梯度趋近于零,导致反向传播时梯度消失,模型难以学习。

顺带一提工程实现中的一个经典技巧:Softmax 对输入整体平移不变(softmax(z)=softmax(zc)\text{softmax}(z) = \text{softmax}(z - c)),因此实际计算时会先按行减去最大值再取指数(softmax(zmaxizi)\text{softmax}(z - \max_i z_i)),避免 ezie^{z_i} 数值上溢——PyTorch 等框架的 Softmax 内部默认就这样做。注意它与缩放因子解决的是两个不同的问题:减最大值防的是数值溢出,除以 dk\sqrt{d_k} 防的是梯度消失

假设查询和键的每个分量都独立地服从标准正态分布 N(0,1)\mathcal{N}(0, 1),那么它们的点积:

qk=i=1dkqikiq \cdot k = \sum_{i=1}^{d_k} q_i k_i

dkd_k 个独立随机变量乘积的和。根据概率论:

  • 均值E[qk]=i=1dkE[qiki]=0E[q \cdot k] = \sum_{i=1}^{d_k} E[q_i k_i] = 0(因为 qiq_ikik_i 独立,E[qi]=E[ki]=0E[q_i] = E[k_i] = 0

  • 方差Var(qk)=i=1dkVar(qiki)=dk\text{Var}(q \cdot k) = \sum_{i=1}^{d_k} \text{Var}(q_i k_i) = d_k(因为当 qi,kiN(0,1)q_i, k_i \sim \mathcal{N}(0,1) 时,qikiq_i k_i 的方差为 1,共 dkd_k 项独立)

因此,点积的标准差为 dk\sqrt{d_k}。当 dk=64d_k = 64 时,标准差为 8;当 dk=512d_k = 512 时,标准差约为 22.6。这意味着点积值的数值范围和量级会随维度增长而扩大——典型的点积值(约在均值周围 ±2σ\pm 2\sigma 范围内)会变得越来越大。

下表直观展示了不同 dkd_k 下,未缩放点积的方差增长。Softmax 输出和梯度并不由 dkd_k 单独决定,还取决于整行 logits、序列长度和分数分布;这里保留方差和典型尺度,避免把某个 toy logits 设置误读成通用数值规律。

dkd_k

点积方差 Var(qk)\text{Var}(q \cdot k)

点积标准差

典型点积范围(±2σ\approx \pm 2\sigma

缩放后方差

16

16

4.0

±8\pm 8

1

64

64

8.0

±16\pm 16

1

128

128

11.3

±22.6\pm 22.6

1

512

512

22.6

±45.2\pm 45.2

1

[!NOTE] 关键观察:未缩放时,dkd_k 增大会使点积分数的尺度扩大,Softmax 更容易进入饱和区域,梯度变小。而除以 dk\sqrt{d_k} 后(最右列),方差恢复为常数 1,使 Softmax 更可能保持在健康的梯度区域工作。这就是为什么缩放因子对训练稳定性至关重要。

除以 dk\sqrt{d_k} 正是将点积的理论标准差重新缩放为 1(实际模型中约为常数级 1),使 Softmax 保持在梯度较大的“活跃区”,从而确保训练的稳定性。

这不是一个随意的超参数选择,而是在常见初始化假设下可以由方差分析解释的稳定化设计。

2.2.4 Softmax 的作用:从分数到概率

Softmax 在注意力计算中扮演着双重角色:

第一,归一化:将任意实数范围的注意力分数转换为非负且总和为 1 的权重,使其可以解释为“概率分布”——每个位置被关注的程度。

第二,竞争机制:Softmax 的指数特性使得较大的分数被进一步放大,较小的分数被进一步压缩。这创造了一种赢者通吃的效果,使模型能够将注意力集中在最相关的少数位置上,而非均匀分散。

下面的示例直观说明了这一效果。假设某个查询对五个键的点积分数为 [2.0,1.0,0.1,0.5,1.0][2.0, 1.0, 0.1, -0.5, -1.0],经过 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 完整的计算示例

以一个简化的例子来完整展示缩放点积注意力的计算过程。假设序列长度 n=3n=3,维度 dk=4d_k = 4

Q=[101001011100],K=[110000111010],V=[100101101100]Q = \begin{bmatrix} 1 & 0 & 1 & 0 \\ 0 & 1 & 0 & 1 \\ 1 & 1 & 0 & 0 \end{bmatrix}, \quad K = \begin{bmatrix} 1 & 1 & 0 & 0 \\ 0 & 0 & 1 & 1 \\ 1 & 0 & 1 & 0 \end{bmatrix}, \quad V = \begin{bmatrix} 1 & 0 & 0 & 1 \\ 0 & 1 & 1 & 0 \\ 1 & 1 & 0 & 0 \end{bmatrix}

  1. 计算点积 S=QKTS = QK^TS=[112110201]S = \begin{bmatrix} 1 & 1 & 2 \\ 1 & 1 & 0 \\ 2 & 0 & 1 \end{bmatrix}

  2. 缩放(除以 dk=4=2\sqrt{d_k} = \sqrt{4} = 2): S=[0.50.51.00.50.501.000.5]S' = \begin{bmatrix} 0.5 & 0.5 & 1.0 \\ 0.5 & 0.5 & 0 \\ 1.0 & 0 & 0.5 \end{bmatrix}

  3. Softmax 归一化(按行操作,得到注意力权重): A=softmax(S)[0.2740.2740.4520.3840.3840.2330.5060.1860.307]A = \text{softmax}(S') \approx \begin{bmatrix} 0.274 & 0.274 & 0.452 \\ 0.384 & 0.384 & 0.233 \\ 0.506 & 0.186 & 0.307 \end{bmatrix}

  4. 加权求和计算输出 Output=AV\text{Output} = AVOutput[0.7260.7260.2740.2740.6160.6160.3840.3840.8140.4940.1860.506]\text{Output} \approx \begin{bmatrix} 0.726 & 0.726 & 0.274 & 0.274 \\ 0.616 & 0.616 & 0.384 & 0.384 \\ 0.814 & 0.494 & 0.186 & 0.506 \end{bmatrix}

每个输出位置的向量是所有位置值向量的加权组合,权重由查询和键的匹配度决定。在一般形式中,注意力输出维度是 dvd_v;在标准 Transformer 实现里,多头拼接后再通过输出投影回到 dmodeld_{\text{model}},因此可以直接传递给下一层。

下面用 PyTorch 代码完整实现上述计算过程,使用与正文相同的 Q、K 矩阵,读者可以直接运行验证每一步的结果:

运行上述代码后,可以得到如下结果。

把这个权重矩阵绘制为热力图(heatmap),从颜色的深浅可以直观感受每个查询对各键的关注程度。

缩放点积注意力权重热力图

图 2-1:缩放点积注意力权重热力图(生成脚本

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

最后更新于