大模型面试必备13-Softmax 前为什么要除以根号 d?
面试硬核解析: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(dkQKT)
今天,我们就来彻底拆解这道硬核面试题。
一、 直观理解: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 的数值范围始终处于合理区间,我们必须对 QQQ 和 KKK 相乘的结果进行缩放 。
二、 核心症结:维度 ddd 带来的“方差爆炸”
为什么 QQQ 和 KKK 相乘后,数值会变大呢?因为维度叠加导致了方差爆炸。
假设输入向量 QQQ 和 KKK 的维度为 ddd,并且它们的每个元素都是独立随机抽取的,服从均值为 0、方差为 1的分布 。
当 QQQ 和 KKK 进行点积后:
-
结果的期望值依然为 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_iqi 与 kjk_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[Q⋅K]=E[i=1∑dqiki]=i=1∑dE[qiki]
因为 qiq_iqi 和 kik_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[Q⋅K]=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(Q⋅K)=E[(Q⋅K)2]−(E[Q⋅K])2
因为期望为 0,后一项消去,只需计算平方的期望 E[(Q⋅K)2]\mathbb{E}[(Q \cdot K)^2]E[(Q⋅K)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=1∑dqiki)2=i=1∑d(qiki)2+i=j∑qikiqjkj
然后分别求期望 :
-
平方项: ∑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(Q⋅K)=d+0−0=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}}d1。
💡 面试终极一句话总结 (Cheat Sheet):
“点积操作会将输入特征的方差放大 ddd 倍 ;如果不进行缩放,数值过大会导致 Softmax 陷入饱和区,引发梯度消失 。除以 d\sqrt{d}d 是为了将点积结果的方差强行拉回 1,保证 Softmax 输入处于合理区间,从而保障反向传播时梯度的稳定 。”

print('hello world')
更多推荐



所有评论(0)