面试硬核解析:Transformer 中 Softmax 前为什么要除以根号 ddd

在深度学习和大模型的面试中,Transformer 架构的细节是绝对的必考区。其中,关于自注意力机制(Self-Attention)经常有一个极为经典的连环问:

“在计算 Attention Score 时,为什么要除以 dk\sqrt{d_k}dk ?”
“你能从数学期望和方差的角度推导一下吗?”

它的核心公式如下 :

Attention Scores=Softmax(QKTdk)\text{Attention Scores} = \text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)Attention Scores=Softmax(dk QKT)

今天,我们就来彻底拆解这道硬核面试题。


一、 直观理解:Softmax 的“敏感病”

首先,我们要明白 Softmax 并不是一个完美的函数,它对输入的绝对大小极其敏感

  • 输入过大: Softmax 的输出概率分布会变得极度尖锐(即某个值无限逼近 1,其他值逼近 0) 。这会导致一个致命问题——梯度消失在反向传播时梯度趋近于 0,模型几乎无法学习

  • 输入过小: 输出分布会趋于均匀(大家都是平庸的概率),注意力机制就失去了“区分度”和重点 。

灾难演示:
假设我们在未经缩放时,计算出的注意力分数差异稍微大一点,比如 [1000, 10, 5]
经过 Softmax 后,输出会变成:

Softmax([1000,10,5])≈[1,0,0]\text{Softmax}([1000, 10, 5]) \approx [1, 0, 0]Softmax([1000,10,5])[1,0,0]

此时,梯度几乎无法传播到非最大值的位置,整个网络在这个节点上就“卡死”了 。为了保证输入 Softmax 的数值范围始终处于合理区间,我们必须对 QQQKKK 相乘的结果进行缩放 。


二、 核心症结:维度 ddd 带来的“方差爆炸”

为什么 QQQKKK 相乘后,数值会变大呢?因为维度叠加导致了方差爆炸

假设输入向量 QQQKKK 的维度为 ddd,并且它们的每个元素都是独立随机抽取的,服从均值为 0、方差为 1的分布 。
QQQKKK 进行点积后:

  • 结果的期望值依然为 0

  • 但结果的方差变成了 ddd

这意味着,如果你的词向量维度 d=512d=512d=512,那么点积结果的方差直接飙升到了 512!数值波动的范围变得极大,极易触发 Softmax 的梯度消失。


三、 面试高光:手撕数学推导

如果面试官让你证明“方差为什么变成了 ddd”,请自信地写下以下推导过程:

Q=[q1,q2,…,qd]Q = [q_1, q_2, \dots, q_d]Q=[q1,q2,,qd]K=[k1,k2,…,kd]K = [k_1, k_2, \dots, k_d]K=[k1,k2,,kd] ,且满足:

  • E[qi]=E[kj]=0\mathbb{E}[q_i] = \mathbb{E}[k_j] = 0E[qi]=E[kj]=0

  • Var(qi)=Var(kj)=1Var(q_i) = Var(k_j) = 1Var(qi)=Var(kj)=1

  • qiq_iqikjk_jkj 相互独立

1. 证明期望为 0

点积的期望展开为 :

E[Q⋅K]=E[∑i=1dqiki]=∑i=1dE[qiki]\mathbb{E}[Q \cdot K] = \mathbb{E}\left[\sum_{i=1}^{d}q_i k_i\right] = \sum_{i=1}^{d}\mathbb{E}[q_i k_i]E[QK]=E[i=1dqiki]=i=1dE[qiki]

因为 qiq_iqikik_iki 独立且均值为 0,所以 E[qiki]=E[qi]⋅E[ki]=0×0=0\mathbb{E}[q_i k_i] = \mathbb{E}[q_i] \cdot \mathbb{E}[k_i] = 0 \times 0 = 0E[qiki]=E[qi]E[ki]=0×0=0
最终得出 E[Q⋅K]=0\mathbb{E}[Q \cdot K] = 0E[QK]=0

2. 证明方差为 ddd

根据方差的定义:

Var(Q⋅K)=E[(Q⋅K)2]−(E[Q⋅K])2Var(Q \cdot K) = \mathbb{E}[(Q \cdot K)^2] - (\mathbb{E}[Q \cdot K])^2Var(QK)=E[(QK)2](E[QK])2

因为期望为 0,后一项消去,只需计算平方的期望 E[(Q⋅K)2]\mathbb{E}[(Q \cdot K)^2]E[(QK)2] 。我们将平方项展开 :

(∑i=1dqiki)2=∑i=1d(qiki)2+∑i≠jqikiqjkj\left(\sum_{i=1}^{d}q_i k_i\right)^2 = \sum_{i=1}^{d}(q_i k_i)^2 + \sum_{i \neq j}q_i k_i q_j k_j(i=1dqiki)2=i=1d(qiki)2+i=jqikiqjkj

然后分别求期望 :

  • 平方项: ∑i=1dE[qi2ki2]=∑i=1dE[qi2]⋅E[ki2]\sum_{i=1}^{d}\mathbb{E}[q_i^2 k_i^2] = \sum_{i=1}^{d}\mathbb{E}[q_i^2] \cdot \mathbb{E}[k_i^2]i=1dE[qi2ki2]=i=1dE[qi2]E[ki2] 。因为方差为 1 且均值为 0,所以 E[qi2]=1\mathbb{E}[q_i^2] = 1E[qi2]=1。每一项都是 1×1=11 \times 1 = 11×1=1,求和后为 ddd

  • 交叉项: ∑i≠jE[qikiqjkj]=∑i≠jE[qiqj]⋅E[kikj]\sum_{i \neq j}\mathbb{E}[q_i k_i q_j k_j] = \sum_{i \neq j}\mathbb{E}[q_i q_j] \cdot \mathbb{E}[k_i k_j]i=jE[qikiqjkj]=i=jE[qiqj]E[kikj] 。由于变量相互独立且均值为 0,E[qiqj]=0×0=0\mathbb{E}[q_i q_j] = 0 \times 0 = 0E[qiqj]=0×0=0,所以所有交叉项的期望均为 0 。

因此,方差为:

Var(Q⋅K)=d+0−0=dVar(Q \cdot K) = d + 0 - 0 = dVar(QK)=d+00=d


四、 总结:为什么要除以 d\sqrt{d}d

因为点积后的方差变成了 ddd,为了让方差重新回到 1(保持数据的稳定分布),我们根据统计学中方差的性质 Var(cX)=c2Var(X)Var(cX) = c^2 Var(X)Var(cX)=c2Var(X),需要给原本的变量乘以 1d\frac{1}{\sqrt{d}}d 1

💡 面试终极一句话总结 (Cheat Sheet):
“点积操作会将输入特征的方差放大 ddd 倍 ;如果不进行缩放,数值过大会导致 Softmax 陷入饱和区,引发梯度消失 。除以 d\sqrt{d}d 是为了将点积结果的方差强行拉回 1,保证 Softmax 输入处于合理区间,从而保障反向传播时梯度的稳定 。”

在这里插入图片描述

print('hello world')
Logo

脑启社区是一个专注类脑智能领域的开发者社区。欢迎加入社区,共建类脑智能生态。社区为开发者提供了丰富的开源类脑工具软件、类脑算法模型及数据集、类脑知识库、类脑技术培训课程以及类脑应用案例等资源。

更多推荐