在最初的 transformer 架构中,完整的 attention 公式为:
Attention(Q,K,V)=softmax(dkQKT)V
这篇文章推导公式中的 dk 是如何得到的。
值得说明一点,这里的 dk 在代码中常被称为 norm_factor,它的取值可能根据模型架构、attention 变体以及长度外推方法不同而有所改变。
推导过程
dk 表示每个 attention head 中 query/key 向量的维度,也就是通常所说的 head_dim。
假设某个 query 向量和 key 向量分别为:
q=[q1,q2,…,qdk],k=[k1,k2,…,kdk]
它们的点积为:
q⋅k=i=1∑dkqiki
也就是:
q⋅k=q1k1+q2k2+⋯+qdkkdk
为了分析这个点积的数值尺度,假设 qi 和 ki 是相互独立的随机变量,并且满足:
E[qi]=0,E[ki]=0,Var(qi)=1,Var(ki)=1
令:
Xi=qiki
则:
q⋅k=i=1∑dkXi
首先计算单项 Xi 的期望:
E[Xi]=E[qiki]
由于 qi 和 ki 独立:
E[qiki]=E[qi]E[ki]
又因为 E[qi]=0,E[ki]=0,所以:
E[Xi]=0
接着计算 Xi 的方差:
Var(Xi)=E[Xi2]−E[Xi]2
因为 E[Xi]=0,所以:
Var(Xi)=E[Xi2]
代入 Xi=qiki:
Var(Xi)=E[(qiki)2]=E[qi2ki2]
由于 qi 和 ki 独立:
E[qi2ki2]=E[qi2]E[ki2]
又因为 Var(qi)=E[qi2]−E[qi]2=1,且 E[qi]=0,所以 E[qi2]=1。同理 E[ki2]=1。
因此:
Var(Xi)=E[qi2]E[ki2]=1
也就是说,每一项乘积 qiki 的方差约为 1。
由于点积 q⋅k 是 dk 个这样的项相加:
q⋅k=X1+X2+⋯+Xdk
在各项近似独立的情况下,协方差为0,故和的方差等于方差之和:
Var(q⋅k)=Var(i=1∑dkXi)=i=1∑dkVar(Xi)=dk
因此:
Var(q⋅k)=dk
对应的标准差为:
Std(q⋅k)=Var(q⋅k)=dk
这说明,当 dk 变大时,未缩放的点积 q⋅k 的典型数值幅度会随着 dk 增大。
而 attention score 后面会进入 softmax:
softmax(QKT)
如果 QKT 的数值过大,softmax 会变得非常尖锐。例如:
softmax([20,1,−3,0])≈[1,0,0,0]
这会导致 attention 分布过早饱和,梯度变小,训练变得不稳定。
因此,为了让进入 softmax 的 logits 保持在相对稳定的尺度,需要对点积结果做归一化缩放:
dkq⋅k
缩放之后:
Var(dkq⋅k)=dkVar(q⋅k)=dkdk=1
对应标准差为:
Std(dkq⋅k)=1
也就是说,除以 dk 后,attention logits 的尺度大致回到 O(1) 的量级,不会随着 head_dim 增大而持续变大。
因此,在原始 Transformer 中:
norm_factor=dk
其作用是控制 QKT 的数值尺度,使 softmax 不至于过度饱和,从而提升训练稳定性。
评论