Transformer 中缩放点积注意力机制的数学原理
为什么注意力分数要除以 dk\sqrt{d_k}dk
1. 引言
Transformer 的核心结构之一是注意力机制。注意力机制允许模型根据当前任务,动态判断输入序列中的哪些位置更加重要。
在 Transformer 中,最常用的注意力计算方式是缩放点积注意力:
softmax(QKTdk)V \operatorname{softmax} \left( \frac{QK^{\mathrm T}}{\sqrt{d_k}} \right)V softmax(dk QKT)V
其中:
- QQQ 表示 Query,即查询矩阵;
- KKK 表示 Key,即键矩阵;
- VVV 表示 Value,即值矩阵;
- dkd_kdk 表示单个 Query 或 Key 向量的维度;
- QKTQK^{\mathrm T}QKT 表示 Query 与 Key 之间的点积相似度;
- Softmax 将相似度分数转换为总和为 1 的注意力权重。
这个公式中,一个看似简单但非常重要的操作是:
QKTdk \frac{QK^{\mathrm T}}{\sqrt{d_k}} dk QKT
也就是将 Query 和 Key 的点积结果除以 dk\sqrt{d_k}dk 。
这个缩放操作不是经验性的随意设计,而是由点积的统计性质和 Softmax 的指数放大特性共同决定的。
它的主要作用是:
控制注意力分数的数值尺度,避免随着向量维度增大,Softmax 过早进入过于尖锐的区域,从而提高模型训练的稳定性。
2. Query、Key 和 Value 的作用
假设输入序列经过嵌入层后得到矩阵:
X∈Rn×dmodel X\in\mathbb R^{n\times d_{\text{model}}} X∈Rn×dmodel
其中:
- nnn 表示序列长度;
- dmodeld_{\text{model}}dmodel 表示每个输入向量的维度。
Transformer 使用三个可学习的线性变换矩阵:
WQ,WK,WV W_Q,\qquad W_K,\qquad W_V WQ,WK,WV
将输入分别映射为 Query、Key 和 Value:
Q=XWQ Q=XW_Q Q=XWQ
K=XWK K=XW_K K=XWK
V=XWV V=XW_V V=XWV
可以将三者理解为:
- Query 表示"当前正在寻找什么";
- Key 表示"每个位置能够提供什么特征";
- Value 表示"每个位置真正携带的信息"。
对于某个 Query 向量 qqq,模型会将它与所有 Key 向量进行比较:
qTk1,qTk2,...,qTkn q^{\mathrm T}k_1,\quad q^{\mathrm T}k_2,\quad \ldots,\quad q^{\mathrm T}k_n qTk1,qTk2,...,qTkn
点积越大,通常表示当前 Query 与对应 Key 越匹配。
经过 Softmax 后,模型得到每个位置的注意力权重:
a1,a2,...,an a_1,a_2,\ldots,a_n a1,a2,...,an
最后使用这些权重对 Value 进行加权求和:
o=∑j=1najvj o=\sum_{j=1}^{n}a_jv_j o=j=1∑najvj
因此,注意力机制的基本流程可以概括为:
计算匹配程度→转换为注意力权重→加权汇总信息 \text{计算匹配程度} \rightarrow \text{转换为注意力权重} \rightarrow \text{加权汇总信息} 计算匹配程度→转换为注意力权重→加权汇总信息
3. Query 与 Key 的点积计算
设一个 Query 向量和一个 Key 向量分别为:
q=q1 q2 ⋮ qdk q= \begin{bmatrix} q_1\ q_2\ \vdots\ q_{d_k} \end{bmatrix} q=q1 q2 ⋮ qdk
k=k1 k2 ⋮ kdk k= \begin{bmatrix} k_1\ k_2\ \vdots\ k_{d_k} \end{bmatrix} k=k1 k2 ⋮ kdk
它们的点积为:
qTk=∑i=1dkqiki q^{\mathrm T}k= \sum_{i=1}^{d_k}q_ik_i qTk=i=1∑dkqiki
展开后为:
qTk=q1k1+q2k2+⋯+qdkkdk q^{\mathrm T}k=q_1k_1+ q_2k_2+ \cdots+ q_{d_k}k_{d_k} qTk=q1k1+q2k2+⋯+qdkkdk
这个点积可以看成将 dkd_kdk 个维度上的局部匹配结果累加起来。
当 dkd_kdk 较小时,累加项较少,点积结果的波动范围通常较小。
当 dkd_kdk 较大时,需要累加大量乘积项,点积结果的波动范围会随之增大。
4. 为什么维度增大会导致点积方差增大
为了分析点积的统计特性,假设 Query 和 Key 的每个元素满足:
qi∼N(0,1) q_i\sim N(0,1) qi∼N(0,1)
ki∼N(0,1) k_i\sim N(0,1) ki∼N(0,1)
并且不同维度之间相互独立。
定义:
Xi=qiki X_i=q_ik_i Xi=qiki
那么点积可以写成:
qTk=∑i=1dkXi q^{\mathrm T}k=\sum_{i=1}^{d_k}X_i qTk=i=1∑dkXi
4.1 单个乘积项的均值
由于 qiq_iqi 和 kik_iki 的均值都为 0,并且二者相互独立,因此:
EXi=Eqiki EX_i=Eq_ik_i EXi=Eqiki
根据独立随机变量期望可分离的性质:
Eqiki=EqiEki Eq_ik_i=Eq_iEk_i Eqiki=EqiEki
因为:
Eqi=0 Eq_i=0 Eqi=0
Eki=0 Ek_i=0 Eki=0
所以:
EXi=0 EX_i=0 EXi=0
因此,每一个乘积项的期望都为 0。
4.2 单个乘积项的方差
随机变量 XiX_iXi 的方差为:
Var(Xi)=EXi2−EXi2 \operatorname{Var}(X_i)=EX_i\^2-EX_i^2 Var(Xi)=EXi2−EXi2
因为:
EXi=0 EX_i=0 EXi=0
所以:
Var(Xi)=EXi2 \operatorname{Var}(X_i)=EX_i\^2 Var(Xi)=EXi2
又因为:
Xi=qiki X_i=q_ik_i Xi=qiki
所以:
Xi2=qi2ki2 X_i^2=q_i^2k_i^2 Xi2=qi2ki2
因此:
Var(Xi)=Eqi2ki2 \operatorname{Var}(X_i)=Eq_i\^2k_i\^2 Var(Xi)=Eqi2ki2
由于 qiq_iqi 和 kik_iki 相互独立:
Eqi2ki2=Eqi2Eki2 Eq_i\^2k_i\^2=Eq_i\^2Ek_i\^2 Eqi2ki2=Eqi2Eki2
标准正态分布的方差为 1,因此:
Eqi2=1 Eq_i\^2=1 Eqi2=1
Eki2=1 Ek_i\^2=1 Eki2=1
所以:
Var(Xi)=1 \operatorname{Var}(X_i)=1 Var(Xi)=1
4.3 点积结果的方差
点积由 dkd_kdk 个独立随机变量相加得到:
qTk=X1+X2+⋯+Xdk q^{\mathrm T}k=X_1+X_2+\cdots+X_{d_k} qTk=X1+X2+⋯+Xdk
独立随机变量相加时,方差也相加:
Var(qTk)=∑i=1dkVar(Xi) \operatorname{Var}(q^{\mathrm T}k)=\sum_{i=1}^{d_k}\operatorname{Var}(X_i) Var(qTk)=i=1∑dkVar(Xi)
由于每个 XiX_iXi 的方差都是 1:
Var(qTk)=∑i=1dk1 \operatorname{Var}(q^{\mathrm T}k)=\sum_{i=1}^{d_k}1 Var(qTk)=i=1∑dk1
因此:
Var(qTk)=dk \boxed{ \operatorname{Var}(q^{\mathrm T}k)=d_k } Var(qTk)=dk
对应的标准差为:
Std(qTk)=dk \boxed{ \operatorname{Std}(q^{\mathrm T}k)=\sqrt{d_k} } Std(qTk)=dk
这说明,随着向量维度增加,点积结果的典型波动幅度会按照 dk\sqrt{d_k}dk 增长。
例如:
| dkd_kdk | 点积方差 | 点积标准差 |
|---|---|---|
| 111 | 111 | 111 |
| 444 | 444 | 222 |
| 161616 | 161616 | 444 |
| 646464 | 646464 | 888 |
| 256256256 | 256256256 | 161616 |
当 dk=64d_k=64dk=64 时:
\\operatorname{Std}(q\^{\\mathrm T}k)=# \\sqrt{64} 8
这意味着,未经缩放的点积分数通常会在比单个输入维度大约 8 倍的尺度上波动。
5. 方差与标准差的区别
这里必须区分方差和标准差。
点积的方差为:
Var(qTk)=dk \operatorname{Var}(q^{\mathrm T}k)=d_k Var(qTk)=dk
点积的标准差为:
Std(qTk)=dk \operatorname{Std}(q^{\mathrm T}k)=\sqrt{d_k} Std(qTk)=dk
方差描述的是随机变量偏离均值程度的平方尺度,而标准差描述的是随机变量实际的典型波动幅度。
因此,当 dk=64d_k=64dk=64 时:
- 点积的方差为 646464;
- 点积的标准差为 888。
不能简单地说点积数值扩大了 64 倍。更加准确的说法是:
点积分布的方差扩大为 dkd_kdk,而点积结果的典型波动幅度扩大为 dk\sqrt{d_k}dk 。
6. 真正重要的是分数之间的差值
注意力分数绝对值变大,并不是最直接的问题。
Softmax 对所有分数同时加上同一个常数是不敏感的。
例如:
softmax(2,0,−1) \operatorname{softmax}(2,0,-1) softmax(2,0,−1)
与:
softmax(102,100,99) \operatorname{softmax}(102,100,99) softmax(102,100,99)
计算结果完全相同。
证明如下。设所有输入同时增加常数 ccc:
pi=ezi+c∑jezj+c p_i=\frac{e^{z_i+c}} {\sum_j e^{z_j+c}} pi=∑jezj+cezi+c
由于:
ezi+c=ecezi e^{z_i+c}=e^ce^{z_i} ezi+c=ecezi
所以:
pi=eceziec∑jezj p_i=\frac{e^ce^{z_i}} {e^c\sum_j e^{z_j}} pi=ec∑jezjecezi
约去 ece^cec 后:
pi=ezi∑jezj p_i=\frac{e^{z_i}} {\sum_j e^{z_j}} pi=∑jezjezi
因此,Softmax 真正关心的是不同分数之间的差值,而不是它们的绝对大小。
7. Softmax 如何放大分数差异
Softmax 的定义为:
pi=ezi∑jezj p_i=\frac{e^{z_i}} {\sum_j e^{z_j}} pi=∑jezjezi
对于任意两个位置 iii 和 jjj,它们的概率比为:
pipj=eziezj \frac{p_i}{p_j}=\frac{e^{z_i}}{e^{z_j}} pjpi=ezjezi
因此:
pipj=ezi−zj \boxed{ \frac{p_i}{p_j}=e^{z_i-z_j} } pjpi=ezi−zj
这条公式说明:
Softmax 会将线性的分数差转换为指数级的概率差。
例如,当两个分数相差 2 时:
e2≈7.39 e^2\approx7.39 e2≈7.39
当两个分数相差 8 时:
e8≈2980.96 e^8\approx2980.96 e8≈2980.96
当两个分数相差 16 时:
e16≈8.89×106 e^{16}\approx8.89\times10^6 e16≈8.89×106
因此,高维点积使分数差距增大后,Softmax 会通过指数函数进一步放大这种差距。
8. 低方差状态下的 Softmax
假设当前 Query 与三个 Key 的点积分数为:
z=2,0,−1 z=2,0,-1 z=2,0,−1
Softmax 定义为:
pi=ezi∑jezj p_i=\frac{e^{z_i}} {\sum_j e^{z_j}} pi=∑jezjezi
先计算指数:
e2≈7.389 e^2\approx7.389 e2≈7.389
e0=1 e^0=1 e0=1
e−1≈0.368 e^{-1}\approx0.368 e−1≈0.368
总和为:
7.389+1+0.368=8.757 7.389+1+0.368=8.757 7.389+1+0.368=8.757
因此,第一个位置的概率为:
p1=7.3898.757≈0.844 p_1=\frac{7.389}{8.757} \approx0.844 p1=8.7577.389≈0.844
第二个位置的概率为:
p2=18.757≈0.114 p_2=\frac{1}{8.757} \approx0.114 p2=8.7571≈0.114
第三个位置的概率为:
p3=0.3688.757≈0.042 p_3=\frac{0.368}{8.757} \approx0.042 p3=8.7570.368≈0.042
最终结果为:
softmax(2,0,−1)≈0.844,0.114,0.042 \operatorname{softmax}(2,0,-1) \approx 0.844,0.114,0.042 softmax(2,0,−1)≈0.844,0.114,0.042
这个注意力分布虽然主要关注第一个位置,但第二个和第三个位置仍然保留了一定权重。
因此,这种分布仍然具有较好的可调整性。
9. 未缩放高维状态下的 Softmax
假设 dk=64d_k=64dk=64。
点积分数的标准差约为:
64=8 \sqrt{64}=8 64 =8
为了直观展示高维点积的影响,可以将前面的分数放大 8 倍:
2,0,−1\]×8=\[16,0,−8\] \[2,0,-1\]\\times8=\[16,0,-8\] \[2,0,−1\]×8=\[16,0,−8
现在计算:
softmax(16,0,−8) \operatorname{softmax}(16,0,-8) softmax(16,0,−8)
首先:
e16≈8,886,111 e^{16}\approx8,886,111 e16≈8,886,111
e0=1 e^0=1 e0=1
e−8≈0.000335 e^{-8}\approx0.000335 e−8≈0.000335
因此:
p1≈8,886,1118,886,112.000335 p_1 \approx \frac{8,886,111} {8,886,112.000335} p1≈8,886,112.0003358,886,111
p1≈0.999999887 p_1 \approx 0.999999887 p1≈0.999999887
第二个位置的概率约为:
p2≈1.125×10−7 p_2 \approx 1.125\times10^{-7} p2≈1.125×10−7
第三个位置的概率约为:
p3≈3.775×10−11 p_3 \approx 3.775\times10^{-11} p3≈3.775×10−11
最终:
softmax(16,0,−8)≈0.999999887,1.125×10−7,3.775×10−11 \operatorname{softmax}(16,0,-8) \approx 0.999999887, 1.125\\times10\^{-7}, 3.775\\times10\^{-11} softmax(16,0,−8)≈0.999999887,1.125×10−7,3.775×10−11
这个结果已经非常接近:
1,0,0\] \[1,0,0\] \[1,0,0
也就是说,注意力分布几乎变成了 one-hot 分布。
第一个位置获得了几乎全部的注意力,而其他位置基本被忽略。
10. 注意力尖锐本身是否一定有问题
需要注意,注意力集中在少数位置本身不一定是错误的。
当模型已经训练成熟后,某些任务确实可能要求模型集中关注某个特定位置。
真正的问题在于:
如果训练初期注意力就因为数值尺度问题过早接近 one-hot,模型将很难重新调整 Query 与 Key 之间的匹配关系。
例如,训练初期由于随机初始化,模型可能错误地认为第一个位置最重要。
如果注意力权重为:
0.40,0.35,0.25\] \[0.40,0.35,0.25\] \[0.40,0.35,0.25
模型仍然可以比较容易地重新调整不同位置的权重。
但是,如果注意力权重已经变成:
0.9999999,0.0000001,0\] \[0.9999999,0.0000001,0\] \[0.9999999,0.0000001,0
其他位置对输出的影响就几乎消失了。
此时模型想要把注意力从第一个位置移动到第二个位置,会更加困难。
11. Softmax 的完整导数
很多简化解释会写:
Softmax 导数=pi(1−pi) \text{Softmax 导数}=p_i(1-p_i) Softmax 导数=pi(1−pi)
这个说法并不完整。
Softmax 是一个多输入、多输出函数。一个输入分数的变化,会影响所有输出概率。
设:
pi=ezi∑jezj p_i=\frac{e^{z_i}} {\sum_j e^{z_j}} pi=∑jezjezi
Softmax 的完整导数为:
∂pi∂zj=pi(δij−pj) \boxed{ \frac{\partial p_i}{\partial z_j}=p_i(\delta_{ij}-p_j) } ∂zj∂pi=pi(δij−pj)
其中,δij\delta_{ij}δij 是克罗内克符号:
δij={1,i=j 0,i≠j \delta_{ij}=\begin{cases} 1,&i=j\ 0,&i\neq j \end{cases} δij={1,i=j 0,i=j
11.1 当 i=ji=ji=j 时
当输出位置和输入位置相同时:
∂pi∂zi=pi(1−pi) \frac{\partial p_i}{\partial z_i}=p_i(1-p_i) ∂zi∂pi=pi(1−pi)
11.2 当 i≠ji\neq ji=j 时
当输出位置和输入位置不同时:
∂pi∂zj=−pipj \frac{\partial p_i}{\partial z_j}=-p_ip_j ∂zj∂pi=−pipj
因此,Softmax 的导数不仅包括自身项,还包括不同位置之间的交叉影响。
12. Softmax 的雅可比矩阵
对于三个输出概率:
p=p1,p2,p3T p= p_1,p_2,p_3^{\mathrm T} p=p1,p2,p3T
Softmax 的雅可比矩阵为:
J=p1(1−p1)−p1p2−p1p3 −p2p1p2(1−p2)−p2p3 −p3p1−p3p2p3(1−p3) J= \begin{bmatrix} p_1(1-p_1)&-p_1p_2&-p_1p_3\ -p_2p_1&p_2(1-p_2)&-p_2p_3\ -p_3p_1&-p_3p_2&p_3(1-p_3) \end{bmatrix} J=p1(1−p1)−p1p2−p1p3 −p2p1p2(1−p2)−p2p3 −p3p1−p3p2p3(1−p3)
也可以写成更加紧凑的形式:
J=diag(p)−ppT\boxed{J =\operatorname{diag}(p)-pp^{\mathrm T} } J=diag(p)−ppT
其中:
diag(p) \operatorname{diag}(p) diag(p)
表示以概率向量 ppp 为对角元素的对角矩阵。
13. 为什么接近 one-hot 时梯度变小
假设 Softmax 输出为:
p=0.9999999,0.0000001,0 p= 0.9999999,0.0000001,0 p=0.9999999,0.0000001,0
第一个位置的对角导数为:
p1(1−p1)=0.9999999×0.0000001 p_1(1-p_1)=0.9999999\times0.0000001 p1(1−p1)=0.9999999×0.0000001
因此:
p1(1−p1)≈10−7 p_1(1-p_1) \approx10^{-7} p1(1−p1)≈10−7
第二个位置的对角导数为:
p2(1−p2)=0.0000001×0.9999999 p_2(1-p_2)=0.0000001\times0.9999999 p2(1−p2)=0.0000001×0.9999999
所以:
p2(1−p2)≈10−7 p_2(1-p_2) \approx10^{-7} p2(1−p2)≈10−7
交叉导数为:
−p1p2=−0.9999999×0.0000001 -p_1p_2=-0.9999999\times0.0000001 −p1p2=−0.9999999×0.0000001
因此:
−p1p2≈−10−7 -p_1p_2 \approx-10^{-7} −p1p2≈−10−7
可以看出,Softmax 雅可比矩阵中的大部分元素都非常小。
反向传播时,损失对分数 zzz 的梯度为:
∂L∂z=JT∂L∂p \frac{\partial L}{\partial z}=J^{\mathrm T} \frac{\partial L}{\partial p} ∂z∂L=JT∂p∂L
即使后续网络传回的梯度:
∂L∂p \frac{\partial L}{\partial p} ∂p∂L
并不小,乘上接近零的雅可比矩阵后,传递到注意力分数的梯度仍然可能变得很弱。
因此,Softmax 饱和更准确的含义是:
注意力概率对分数变化变得不敏感,使模型难以通过调整注意力分数改变当前的注意力分布。
14. 梯度如何传递到 Query 和 Key
缩放后的注意力分数为:
sj=qTkjdk s_j=\frac{q^{\mathrm T}k_j}{\sqrt{d_k}} sj=dk qTkj
对于 Query 向量 qqq,有:
∂sj∂q=kjdk \frac{\partial s_j}{\partial q}=\frac{k_j}{\sqrt{d_k}} ∂q∂sj=dk kj
因此,根据链式法则:
∂L∂q=∑j∂L∂sj∂sj∂q \frac{\partial L}{\partial q}=\sum_j \frac{\partial L}{\partial s_j} \frac{\partial s_j}{\partial q} ∂q∂L=j∑∂sj∂L∂q∂sj
代入后得到:
∂L∂q=1dk∑j∂L∂sjkj \boxed{ \frac{\partial L}{\partial q}=\frac{1}{\sqrt{d_k}} \sum_j \frac{\partial L}{\partial s_j}k_j } ∂q∂L=dk 1j∑∂sj∂Lkj
对于第 jjj 个 Key,有:
∂sj∂kj=qdk \frac{\partial s_j}{\partial k_j}=\frac{q}{\sqrt{d_k}} ∂kj∂sj=dk q
因此:
∂L∂kj=1dk∂L∂sjq \boxed{ \frac{\partial L}{\partial k_j}=\frac{1}{\sqrt{d_k}} \frac{\partial L}{\partial s_j}q } ∂kj∂L=dk 1∂sj∂Lq
如果 Softmax 饱和导致:
∂L∂sj≈0 \frac{\partial L}{\partial s_j} \approx0 ∂sj∂L≈0
那么:
∂L∂q≈0 \frac{\partial L}{\partial q} \approx0 ∂q∂L≈0
并且:
∂L∂kj≈0 \frac{\partial L}{\partial k_j} \approx0 ∂kj∂L≈0
由于:
Q=XWQ Q=XW_Q Q=XWQ
K=XWK K=XW_K K=XWK
所以生成 Query 和 Key 的参数矩阵 WQW_QWQ 和 WKW_KWK 得到的训练信号也会减弱。
这会使模型难以学习:
- Query 应该关注哪些特征;
- 哪些 Key 应该与当前 Query 匹配;
- 当前注意力应该从哪个位置转移到另一个位置;
- 不同注意力头应该形成怎样的分工。
15. 为什么不能说整个模型一定停止训练
虽然 Softmax 饱和会削弱 Query 和 Key 路径上的梯度,但不能简单地认为整个 Transformer 会彻底停止训练。
主要有以下几个原因。
15.1 梯度通常只是很小,而不是严格为零
从数学上来说,只要 Softmax 输出不是严格的 0 或 1,其导数就不是严格的 0。
实际训练中,更常见的情况是梯度非常小,而不是完全消失。
15.2 Value 路径仍然可以传播梯度
注意力输出为:
O=AV O=AV O=AV
其中:
A=softmax(S) A=\operatorname{softmax}(S) A=softmax(S)
损失对 Value 的梯度为:
∂L∂V=AT∂L∂O \frac{\partial L}{\partial V}=A^{\mathrm T} \frac{\partial L}{\partial O} ∂V∂L=AT∂O∂L
即使注意力矩阵 AAA 非常尖锐,被选中的 Value 位置仍然可以获得梯度。
因此,Value 路径不一定完全失去训练信号。
15.3 Transformer 具有残差连接
Transformer 中通常采用残差结构:
Y=X+Attention(X)Y =X+\operatorname{Attention}(X) Y=X+Attention(X)
对输入 XXX 求导:
∂Y∂X=I+∂Attention(X)∂X \frac{\partial Y}{\partial X}=I+ \frac{\partial\operatorname{Attention}(X)} {\partial X} ∂X∂Y=I+∂X∂Attention(X)
其中,III 表示恒等映射。
即使注意力分支的梯度较弱,残差连接仍然可以提供一条直接的梯度传播路径。
因此,更准确的说法是:
未缩放点积不会必然让整个模型停止训练,但会使注意力分布过早尖锐,削弱 Query 和 Key 路径上的有效训练信号,从而降低模型的训练稳定性和学习效率。
16. 为什么除以的是 dk\sqrt{d_k}dk
已知原始点积为:
s=qTk s=q^{\mathrm T}k s=qTk
并且:
Var(s)=dk \operatorname{Var}(s)=d_k Var(s)=dk
现在使用一个常数 ccc 对点积进行缩放:
s~=sc \tilde{s}=\frac{s}{c} s~=cs
随机变量除以常数后,方差会除以该常数的平方:
Var(s~)=Var(s)c2 \operatorname{Var}(\tilde{s})=\frac{\operatorname{Var}(s)}{c^2} Var(s~)=c2Var(s)
代入:
Var(s~)=dkc2 \operatorname{Var}(\tilde{s})=\frac{d_k}{c^2} Var(s~)=c2dk
为了让缩放后的方差保持在大约 1 的尺度,希望:
dkc2=1 \frac{d_k}{c^2}=1 c2dk=1
于是:
c2=dk c^2=d_k c2=dk
所以:
c=dk c=\sqrt{d_k} c=dk
最终得到:
s~=qTkdk \boxed{ \tilde{s}=\frac{q^{\mathrm T}k}{\sqrt{d_k}} } s~=dk qTk
此时:
Var(qTkdk)=Var(qTk)dk \operatorname{Var} \left( \frac{q^{\mathrm T}k}{\sqrt{d_k}} \right)=\frac{\operatorname{Var}(q^{\mathrm T}k)} {d_k} Var(dk qTk)=dkVar(qTk)
因为:
Var(qTk)=dk \operatorname{Var}(q^{\mathrm T}k)=d_k Var(qTk)=dk
所以:
Var(qTkdk)=1 \boxed{ \operatorname{Var} \left( \frac{q^{\mathrm T}k}{\sqrt{d_k}} \right)=1 } Var(dk qTk)=1
这就是缩放因子 dk\sqrt{d_k}dk 的数学来源。
17. 缩放前后的数值对比
假设:
dk=64 d_k=64 dk=64
那么缩放因子为:
\\sqrt{d_k}=# \\sqrt{64} 8
未缩放的注意力分数为:
16,0,−8\] \[16,0,-8\] \[16,0,−8
缩放后:
16,0,−8\]8=\[2,0,−1\] \\frac{\[16,0,-8\]}{8}=\[2,0,-1\] 8\[16,0,−8\]=\[2,0,−1
未缩放时:
softmax(16,0,−8)≈0.999999887,0.000000113,0 \operatorname{softmax}(16,0,-8) \approx 0.999999887, 0.000000113, 0 softmax(16,0,−8)≈0.999999887,0.000000113,0
缩放后:
softmax(2,0,−1)≈0.844,0.114,0.042 \operatorname{softmax}(2,0,-1) \approx 0.844,0.114,0.042 softmax(2,0,−1)≈0.844,0.114,0.042
缩放前,模型几乎只关注第一个位置。
缩放后,第一个位置仍然最重要,但其他位置仍然保留有效权重。
因此,缩放并不是消除注意力的区分能力,而是防止注意力过早变成绝对选择。
18. 为什么不直接除以 dkd_kdk
如果使用:
qTkdk \frac{q^{\mathrm T}k}{d_k} dkqTk
那么缩放后的方差为:
Var(qTkdk)=Var(qTk)dk2 \operatorname{Var} \left( \frac{q^{\mathrm T}k}{d_k} \right)=\frac{\operatorname{Var}(q^{\mathrm T}k)} {d_k^2} Var(dkqTk)=dk2Var(qTk)
由于:
Var(qTk)=dk \operatorname{Var}(q^{\mathrm T}k)=d_k Var(qTk)=dk
所以:
Var(qTkdk)=dkdk2 \operatorname{Var} \left( \frac{q^{\mathrm T}k}{d_k} \right)=\frac{d_k}{d_k^2} Var(dkqTk)=dk2dk
即:
Var(qTkdk)=1dk \boxed{ \operatorname{Var} \left( \frac{q^{\mathrm T}k}{d_k} \right)=\frac{1}{d_k} } Var(dkqTk)=dk1
当 dk=64d_k=64dk=64 时:
Var=164 \operatorname{Var}=\frac{1}{64} Var=641
对应的标准差为:
Std=18 \operatorname{Std}=\frac{1}{8} Std=81
这时分数会被压缩得过小。
例如:
16,0,−8\]64=\[0.25,0,−0.125\] \\frac{\[16,0,-8\]}{64}=\[0.25,0,-0.125\] 64\[16,0,−8\]=\[0.25,0,−0.125
Softmax 结果约为:
softmax(0.25,0,−0.125)≈0.406,0.316,0.279 \operatorname{softmax}(0.25,0,-0.125) \approx 0.406,0.316,0.279 softmax(0.25,0,−0.125)≈0.406,0.316,0.279
这个分布变得比较平坦。
如果分数被过度压缩,模型就难以突出真正重要的位置。
因此:
- 不缩放时,注意力可能过于尖锐;
- 除以 dkd_kdk 时,注意力可能过于平坦;
- 除以 dk\sqrt{d_k}dk 时,注意力分数保持在更加合理的尺度。
19. 缩放与温度 Softmax 的关系
温度 Softmax 的一般形式为:
pi=fracezi/T∑jezj/T p_i=frac{e^{z_i/T}} {\sum_j e^{z_j/T}} pi=fracezi/Tj∑ezj/T
其中,TTT 称为温度参数。
当:
T<1 T<1 T<1
时,分数会被放大,Softmax 分布更加尖锐。
当:
T>1 T>1 T>1
时,分数会被压缩,Softmax 分布更加平缓。
缩放点积注意力为:
softmax(QKTdk) \operatorname{softmax} \left( \frac{QK^{\mathrm T}}{\sqrt{d_k}} \right) softmax(dk QKT)
从形式上看,相当于使用:
T=dk T=\sqrt{d_k} T=dk
但是,在 Transformer 中,这个温度并不是随意选择的超参数,而是根据点积方差随维度增长的规律推导出来的。
因此,可以将其理解为:
向量维度越大,点积越容易产生较大的分数差;除以 dk\sqrt{d_k}dk 相当于提高 Softmax 的有效温度,将注意力分布恢复到较稳定的尺度。
20. 数值溢出与 Softmax 饱和的区别
数值溢出和 Softmax 饱和是两个不同的问题。
20.1 数值溢出
如果直接计算:
e1000 e^{1000} e1000
计算机可能无法表示这个数,从而得到无穷大。
实际实现 Softmax 时,通常先减去输入中的最大值:
softmax(z)=softmax(z−max(z)) \operatorname{softmax}(z)=\operatorname{softmax}(z-\max(z)) softmax(z)=softmax(z−max(z))
例如:
z=1000,990,980 z=1000,990,980 z=1000,990,980
减去最大值后:
z−max(z)=0,−10,−20 z-\max(z)=0,-10,-20 z−max(z)=0,−10,−20
于是:
softmax(1000,990,980)=softmax(0,−10,−20) \operatorname{softmax}(1000,990,980)=\operatorname{softmax}(0,-10,-20) softmax(1000,990,980)=softmax(0,−10,−20)
这样可以避免直接计算 e1000e^{1000}e1000。
20.2 Softmax 饱和
虽然减去最大值可以避免数值溢出,但它不能解决分布过于尖锐的问题。
因为:
softmax(0,−10,−20) \operatorname{softmax}(0,-10,-20) softmax(0,−10,−20)
仍然会非常接近:
1,0,0\] \[1,0,0\] \[1,0,0
因此:
- 减去最大值解决的是指数计算的数值稳定性;
- 除以 dk\sqrt{d_k}dk 解决的是注意力分数尺度和优化稳定性。
两者解决的不是同一个问题。
21. 缩放能否保证 Softmax 永远不饱和
答案是否定的。
除以 dk\sqrt{d_k}dk 主要保证:
注意力分数不会仅仅因为向量维度增加而自动扩大。
但是,Query 和 Key 是模型学习得到的:
Q=XWQ Q=XW_Q Q=XWQ
K=XWK K=XW_K K=XWK
如果训练过程中 Query 和 Key 的范数变得很大,那么即使经过缩放,注意力分数仍然可能很大。
例如:
qTk=800 q^{\mathrm T}k=800 qTk=800
当:
dk=64 d_k=64 dk=64
时:
\\frac{q\^{\\mathrm T}k}{\\sqrt{d_k}}=# \\frac{800}{8} 100
分数 100 进入 Softmax 后,仍然会产生极度尖锐的分布。
因此,Transformer 还需要结合其他稳定化机制,例如:
- 合理的参数初始化;
- Layer Normalization;
- 残差连接;
- 多头注意力;
- 合适的学习率;
- 梯度裁剪;
- 权重衰减;
- 数值稳定的 Softmax 实现。
所以,更准确的说法是:
除以 dk\sqrt{d_k}dk 可以降低 Softmax 因维度增大而过早饱和的风险,但不能保证训练过程中永远不会出现尖锐注意力。
22. 注意力 Softmax 与分类 Softmax 的区别
还需要注意,Softmax 饱和并不意味着所有使用 Softmax 的场景都会出现相同的梯度问题。
在分类任务中,Softmax 通常直接与交叉熵损失结合。
交叉熵为:
L=−∑iyilogpiL =-\sum_i y_i\log p_i L=−i∑yilogpi
Softmax 与交叉熵组合后的梯度为:
∂L∂zi=pi−yi \frac{\partial L}{\partial z_i}=p_i-y_i ∂zi∂L=pi−yi
假设模型预测为:
p=0.9999,0.0001 p=0.9999,0.0001 p=0.9999,0.0001
而真实标签为:
y=0,1 y=0,1 y=0,1
那么:
p−y=0.9999,−0.9999 p-y=0.9999,-0.9999 p−y=0.9999,−0.9999
此时梯度并不小。
这是因为模型虽然非常自信,但它预测错误,交叉熵仍然会产生较强的纠正信号。
而注意力中的 Softmax 不是最终分类输出,它只是模型内部的加权系数:
O=softmax(S)VO =\operatorname{softmax}(S)V O=softmax(S)V
它没有直接对应的 one-hot 分类标签。
因此,在注意力机制内部,当 Softmax 雅可比矩阵接近零时,更容易削弱传递到 Query 和 Key 的梯度。
所以不能简单地认为:
p→0或p→1 p\rightarrow0 \quad\text{或}\quad p\rightarrow1 p→0或p→1
就必然导致任何模型中的梯度都消失。
必须结合 Softmax 所处的位置和后续损失函数进行分析。
23. 多头注意力中的 dkd_kdk
在多头注意力中,模型会将总特征维度拆分成多个注意力头。
假设:
dmodel=512 d_{\text{model}}=512 dmodel=512
注意力头数量为:
h=8 h=8 h=8
那么每个注意力头的维度通常为:
dk=fracdmodelh d_k=frac{d_{\text{model}}}{h} dk=fracdmodelh
代入后:
dk=5128=64 d_k=\frac{512}{8}=64 dk=8512=64
因此,每个注意力头的缩放因子为:
\\sqrt{d_k}=# \\sqrt{64} 8
每个注意力头分别计算:
headi=softmax(QiKiTdk)Vi \operatorname{head}_i=\operatorname{softmax} \left( \frac{Q_iK_i^{\mathrm T}} {\sqrt{d_k}} \right)V_i headi=softmax(dk QiKiT)Vi
随后将所有注意力头拼接:
MultiHead(Q,K,V)=Concat(head1,...,headh)WO \operatorname{MultiHead}(Q,K,V)=\operatorname{Concat} ( \operatorname{head}_1, \ldots, \operatorname{head}_h ) W_O MultiHead(Q,K,V)=Concat(head1,...,headh)WO
这里使用的是每个注意力头内部的 Key 维度 dkd_kdk,而不是总模型维度 dmodeld_{\text{model}}dmodel。
24. 缩放点积注意力的完整计算流程
假设输入矩阵为:
X∈Rn×dmodel X\in\mathbb R^{n\times d_{\text{model}}} X∈Rn×dmodel
24.1 生成 Query、Key 和 Value
Q=XWQ Q=XW_Q Q=XWQ
K=XWK K=XW_K K=XWK
V=XWV V=XW_V V=XWV
24.2 计算原始点积相似度
S=QKT S=QK^{\mathrm T} S=QKT
其中:
S∈Rn×n S\in\mathbb R^{n\times n} S∈Rn×n
矩阵中的元素:
Sij=qiTkj S_{ij}=q_i^{\mathrm T}k_j Sij=qiTkj
表示第 iii 个 Query 与第 jjj 个 Key 的匹配程度。
24.3 对点积分数进行缩放
S~=Sdk \tilde{S}=\frac{S}{\sqrt{d_k}} S~=dk S
即:
S~QKTdk \tilde{S}\frac{QK^{\mathrm T}}{\sqrt{d_k}} S~dk QKT
24.4 加入掩码
如果模型中存在填充位置或者因果约束,则需要加入掩码矩阵 MMM:
Smask=QKTdk+M S_{\text{mask}}=\frac{QK^{\mathrm T}}{\sqrt{d_k}} + M Smask=dk QKT+M
对于不允许关注的位置,掩码通常设置为一个非常大的负数。
经过 Softmax 后,这些位置的权重会接近 0。
24.5 计算注意力权重
A=softmax(Smask)A =\operatorname{softmax}(S_{\text{mask}}) A=softmax(Smask)
Softmax 通常沿矩阵的每一行进行,使每个 Query 对所有 Key 的注意力权重之和为 1:
∑jAij=1 \sum_j A_{ij}=1 j∑Aij=1
24.6 对 Value 进行加权求和
O=AV O=AV O=AV
因此,完整的缩放点积注意力公式为:
O=softmax(QKTdk+M)V \boxed{O =\operatorname{softmax} \left( \frac{QK^{\mathrm T}}{\sqrt{d_k}} + M \right)V } O=softmax(dk QKT+M)V
25. 一个小型矩阵计算示例
假设有两个 Query 和两个 Key:
Q=11 10 Q= \begin{bmatrix} 1&1\ 1&0 \end{bmatrix} Q=11 10
K=10 11 K= \begin{bmatrix} 1&0\ 1&1 \end{bmatrix} K=10 11
并且:
dk=2 d_k=2 dk=2
首先计算:
KT=11 01 K^{\mathrm T}=\begin{bmatrix} 1&1\ 0&1 \end{bmatrix} KT=11 01
因此:
QKT=11 1011 01 QK^{\mathrm T}=\begin{bmatrix} 1&1\ 1&0 \end{bmatrix} \begin{bmatrix} 1&1\ 0&1 \end{bmatrix} QKT=11 1011 01
得到:
QKT=12 11 QK^{\mathrm T}=\begin{bmatrix} 1&2\ 1&1 \end{bmatrix} QKT=12 11
缩放因子为:
dk=2 \sqrt{d_k}=\sqrt{2} dk =2
因此:
QKT2=1/22/2 1/21/2 \frac{QK^{\mathrm T}}{\sqrt{2}}=\begin{bmatrix} 1/\sqrt{2}&2/\sqrt{2}\ 1/\sqrt{2}&1/\sqrt{2} \end{bmatrix} 2 QKT=1/2 2/2 1/2 1/2
近似为:
QKT2≈0.7071.414 0.7070.707 \frac{QK^{\mathrm T}}{\sqrt{2}} \approx \begin{bmatrix} 0.707&1.414\ 0.707&0.707 \end{bmatrix} 2 QKT≈0.7071.414 0.7070.707
对每一行进行 Softmax。
第一行:
softmax(0.707,1.414)≈0.330,0.670 \operatorname{softmax}(0.707,1.414) \approx 0.330,0.670 softmax(0.707,1.414)≈0.330,0.670
第二行:
softmax(0.707,0.707)=0.5,0.5 \operatorname{softmax}(0.707,0.707)=0.5,0.5 softmax(0.707,0.707)=0.5,0.5
因此,注意力矩阵为:
A≈0.3300.670 0.50.5 A \approx \begin{bmatrix} 0.330&0.670\ 0.5&0.5 \end{bmatrix} A≈0.3300.670 0.50.5
这说明:
- 第一个 Query 更关注第二个 Key;
- 第二个 Query 对两个 Key 给予相同的注意力权重。
假设 Value 矩阵为:
V=100 020 V= \begin{bmatrix} 10&0\ 0&20 \end{bmatrix} V=100 020
注意力输出为:
O=AV O=AV O=AV
即:
O=0.3300.670 0.50.5100 020O =\begin{bmatrix} 0.330&0.670\ 0.5&0.5 \end{bmatrix} \begin{bmatrix} 10&0\ 0&20 \end{bmatrix} O=0.3300.670 0.50.5100 020
得到:
O=3.313.4 510O =\begin{bmatrix} 3.3&13.4\ 5&10 \end{bmatrix} O=3.313.4 510
因此,第一个 Query 的输出更多地吸收了第二个 Value 的信息,而第二个 Query 对两个 Value 进行了平均汇总。
26. 更一般的方差情况
前面的推导假设:
Var(qi)=1 \operatorname{Var}(q_i)=1 Var(qi)=1
Var(ki)=1 \operatorname{Var}(k_i)=1 Var(ki)=1
在更加一般的情况下,假设:
Eqi=0 Eq_i=0 Eqi=0
Eki=0 Ek_i=0 Eki=0
并且:
Var(qi)=σq2 \operatorname{Var}(q_i)=\sigma_q^2 Var(qi)=σq2
Var(ki)=σk2 \operatorname{Var}(k_i)=\sigma_k^2 Var(ki)=σk2
那么:
Var(qiki)=σq2σk2 \operatorname{Var}(q_ik_i)=\sigma_q^2\sigma_k^2 Var(qiki)=σq2σk2
点积方差为:
Var(qTk)=dkσq2σk2 \operatorname{Var}(q^{\mathrm T}k)=d_k\sigma_q^2\sigma_k^2 Var(qTk)=dkσq2σk2
缩放后:
Var(qTkdk)=σq2σk2 \operatorname{Var} \left( \frac{q^{\mathrm T}k}{\sqrt{d_k}} \right)=\sigma_q^2\sigma_k^2 Var(dk qTk)=σq2σk2
可以看到,除以 dk\sqrt{d_k}dk 消除了由维度 dkd_kdk 带来的线性增长,但不会自动消除 Query 和 Key 本身方差过大的问题。
这也是为什么 Transformer 还需要合理初始化和归一化机制。
27. 点积与向量范数的关系
点积也可以写成:
qTk=∣q∣∣k∣cosθ q^{\mathrm T}k=|q||k|\cos\theta qTk=∣q∣∣k∣cosθ
其中:
- ∣q∣|q|∣q∣ 表示 Query 的范数;
- ∣k∣|k|∣k∣ 表示 Key 的范数;
- θ\thetaθ 表示两个向量之间的夹角。
当向量各维度方差相近时,向量范数通常会随着维度增大而增大。
例如:
E∣q∣2=dk E\|q\|\^2=d_k E∣q∣2=dk
因此:
∣q∣≈dk |q| \approx \sqrt{d_k} ∣q∣≈dk
同理:
∣k∣≈dk |k| \approx \sqrt{d_k} ∣k∣≈dk
所以:
∣q∣∣k∣≈dk |q||k| \approx d_k ∣q∣∣k∣≈dk
虽然随机方向下的 cosθ\cos\thetacosθ 通常会缩小,但最终点积的标准差仍然呈现 dk\sqrt{d_k}dk 的增长。
这一视角说明:
高维向量的范数会自然增大,点积不仅包含方向相似性,也受到向量长度的影响,因此需要进行适当的尺度控制。
28. 与余弦相似度的区别
余弦相似度为:
cos(q,k)=qTk∣q∣∣k∣ \operatorname{cos}(q,k)=\frac{q^{\mathrm T}k} {|q||k|} cos(q,k)=∣q∣∣k∣qTk
余弦相似度只关注向量方向,不直接受到向量范数影响。
而标准 Transformer 注意力使用:
qTk q^{\mathrm T}k qTk
它同时受到:
- 向量方向;
- 向量范数;
两方面影响。
除以 dk\sqrt{d_k}dk 并不等于将点积转换为余弦相似度,因为它没有分别除以 ∣q∣|q|∣q∣ 和 ∣k∣|k|∣k∣。
缩放点积只是对维度带来的平均尺度进行校正:
qTkdk \frac{q^{\mathrm T}k}{\sqrt{d_k}} dk qTk
因此:
- 余弦相似度消除的是向量范数影响;
- 缩放点积消除的是维度增长带来的典型尺度变化。
29. 常见误解一:维度越高,点积一定越大
这个说法并不准确。
在均值为 0 的条件下:
EqTk=0 Eq\^{\\mathrm T}k=0 EqTk=0
无论 dkd_kdk 多大,点积的期望仍然为 0。
维度增加带来的不是点积一定变成正的大数,而是点积分布更加分散:
Var(qTk)=dk \operatorname{Var}(q^{\mathrm T}k)=d_k Var(qTk)=dk
因此,更准确的说法是:
维度越高,点积出现较大正值或较大负值的概率越高,点积结果的波动范围越大。
30. 常见误解二:Softmax 导数就是 p(1−p)p(1-p)p(1−p)
这个说法只描述了 Softmax 的对角导数:
∂pi∂zi=pi(1−pi) \frac{\partial p_i}{\partial z_i}=p_i(1-p_i) ∂zi∂pi=pi(1−pi)
完整导数还包括交叉项:
∂pi∂zj=−pipji≠j \frac{\partial p_i}{\partial z_j}=-p_ip_j \qquad i\neq j ∂zj∂pi=−pipji=j
因此,Softmax 的完整导数是一个雅可比矩阵:
J=diag(p)−ppTJ =\operatorname{diag}(p)-pp^{\mathrm T} J=diag(p)−ppT
31. 常见误解三:Softmax 饱和后整个模型梯度都会消失
这个说法过于绝对。
Softmax 饱和主要会削弱:
- 注意力分数的梯度;
- Query 路径的梯度;
- Key 路径的梯度;
- 注意力重新分配的能力。
但以下路径仍然可能继续传播梯度:
- Value 路径;
- 输出投影路径;
- 残差连接;
- Transformer 中的其他层和其他注意力头。
因此,更准确的说法是:
Softmax 饱和会降低注意力匹配关系的可学习性,但不一定让整个模型完全停止训练。
32. 常见误解四:除以 dk\sqrt{d_k}dk 可以完美避免饱和
这个说法也不准确。
缩放只能保证:
点积分数不会仅仅因为 dk 增大而扩大 \text{点积分数不会仅仅因为 }d_k\text{ 增大而扩大} 点积分数不会仅仅因为 dk 增大而扩大
如果 Query 和 Key 的范数在训练中变得很大,那么缩放后的分数仍然可能进入 Softmax 饱和区。
因此,除以 dk\sqrt{d_k}dk 是降低风险,而不是提供绝对保证。
33. 常见误解五:缩放是为了避免指数溢出
缩放确实有助于减小分数,但它的主要理论目的不是解决指数溢出。
指数溢出通常通过减去最大值来解决:
softmax(z)=softmax(z−max(z)) \operatorname{softmax}(z)=\operatorname{softmax}(z-\max(z)) softmax(z)=softmax(z−max(z))
缩放的主要目的则是:
控制注意力分数的方差和梯度尺度 \text{控制注意力分数的方差和梯度尺度} 控制注意力分数的方差和梯度尺度
因此,两者的功能不同。
34. 缩放点积注意力的完整因果链条
整个逻辑可以概括为以下过程。
首先,向量维度增大:
dk↑ d_k\uparrow dk↑
导致点积方差增大:
Var(qTk)=dk \operatorname{Var}(q^{\mathrm T}k)=d_k Var(qTk)=dk
点积标准差随之增大:
Std(qTk)=dk \operatorname{Std}(q^{\mathrm T}k)=\sqrt{d_k} Std(qTk)=dk
因此,不同 Key 之间的注意力分数差距更容易增大。
Softmax 的概率比满足:
pipj=ezi−zj \frac{p_i}{p_j}=e^{z_i-z_j} pjpi=ezi−zj
所以分数差距会被指数级放大。
注意力分布可能过早接近:
1,0,0,...\] \[1,0,0,\\ldots\] \[1,0,0,...
此时 Softmax 的雅可比矩阵:
J=diag(p)−ppTJ =\operatorname{diag}(p)-pp^{\mathrm T} J=diag(p)−ppT
中的大量元素非常小。
于是:
∂L∂S \frac{\partial L}{\partial S} ∂S∂L
减弱,进一步导致:
∂L∂Q \frac{\partial L}{\partial Q} ∂Q∂L
和:
∂L∂K \frac{\partial L}{\partial K} ∂K∂L
减弱。
最终,模型难以调整 Query 和 Key 之间的匹配关系。
为了消除维度造成的尺度增长,使用:
QKTdk \frac{QK^{\mathrm T}}{\sqrt{d_k}} dk QKT
使得:
Var(qTkdk)≈1 \operatorname{Var} \left( \frac{q^{\mathrm T}k}{\sqrt{d_k}} \right) \approx1 Var(dk qTk)≈1
从而让 Softmax 分布和梯度尺度保持相对稳定。
35. 最终总结
Transformer 在计算注意力时,将 Query 和 Key 的点积除以 dk\sqrt{d_k}dk ,其根本原因不是简单地防止数值过大,而是控制点积分数的统计尺度。
在各维度相互独立、均值为 0、方差为 1 的假设下:
Var(qTk)=dk \operatorname{Var}(q^{\mathrm T}k)=d_k Var(qTk)=dk
因此:
Std(qTk)=dk \operatorname{Std}(q^{\mathrm T}k)=\sqrt{d_k} Std(qTk)=dk
随着 dkd_kdk 增大,注意力分数之间的差距更容易增大。
Softmax 又会通过指数函数放大分数差异:
pipj=ezi−zj \frac{p_i}{p_j}=e^{z_i-z_j} pjpi=ezi−zj
因此,未经缩放的高维点积容易使注意力权重过早接近 one-hot 分布。
当注意力分布过度尖锐时,Softmax 对分数变化的敏感性降低,传递到 Query 和 Key 路径的梯度会减弱,模型难以修正错误的注意力关系。
通过缩放:
QKTdk \frac{QK^{\mathrm T}}{\sqrt{d_k}} dk QKT
可以将点积分数的方差恢复到大约与维度无关的尺度:
Var(qTkdk)≈1 \operatorname{Var} \left( \frac{q^{\mathrm T}k}{\sqrt{d_k}} \right) \approx1 Var(dk qTk)≈1
这样可以:
- 避免注意力分数仅仅因为维度增大而失控;
- 降低 Softmax 过早进入尖锐区域的风险;
- 保持注意力权重的可调整性;
- 改善 Query 和 Key 路径的梯度传播;
- 提高不同注意力头和不同模型规模下的训练稳定性。
因此,最准确的总结是:
除以 dk\sqrt{d_k}dk 是一种针对高维点积的方差归一化操作。它抵消了点积标准差随向量维度增长的问题,使 Softmax 的概率分布和梯度尺度不会因为 dkd_kdk 增大而失控,从而让 Transformer 更稳定地学习 Query 与 Key 之间的匹配关系。