Transformer中Attention为何除以√dk?解析缩放因子的数学原理与训练价值

发布时间:2026/10/1 20:44:36
Transformer中Attention为何除以√dk?解析缩放因子的数学原理与训练价值 开篇先亮个观点attention里的这个除以√dk是整个Transformer架构里最容易被忽视、但又是最不该被忽视的细节之一。很多人看完代码觉得不就是多除了一个数嘛调参的时候甚至想把它去掉试试结果训练直接崩了或者收敛慢到怀疑人生。这不是玄学背后是一套非常明确的数学逻辑搞懂它你才算真正看懂了Transformer。我看了不少Transformer的源码和论文解读也自己动手复现过attention机制踩过不少坑。今天这篇就把“为什么做scale”这件事彻底讲透从softmax的数学性格、方差推导、实验对比到面试高频追问一次聊完。1. 先把attention的计算流程摆上桌1.1 从三个向量到加权求和要理解scale的作用先得看清楚attention到底在算什么。以最基础的自注意力Self-Attention为例输入是一串token的向量表示每个token会生成三个向量Query、Key、Value分别对应“我想找什么”、“我有什么标签”和“我实际的内容”。注意力权重的计算分为两步每个query和所有key做点积得到一组相似度分数。对这组分数做softmax归一化成权重。用权重去加权求和所有value。用公式表示就是Attention(Q, K, V) softmax(QK^T / √dk) V注意公式中间的那个“/ √dk”就是今天的主角。1.2 为什么点积能表示相似度先把点积这件事讲透。两个向量做点积数学上等于各维度相乘再求和几何上可以写成a · b |a||b|cosθ在attention的场景里如果两个向量的方向越接近cosθ越大点积结果越大说明两者越相似。如果方向相反点积为负说明不相似。这本质上是一种“打分”机制分数越高说明当前这个query越应该关注对应的key。当维度dk比较小的时候比如只有16、32维点积结果自然在个位数徘徊softmax还能正常处理。但在真实的Transformer模型里dk通常是64、128甚至更大问题就来了。这也是为什么原始论文里专门留了一段话解释scale的必要性很多初学者读到这里都会跳过去但它恰恰是关键。提一句源码实践在PyTorch的nn.MultiheadAttention实现里q k.transpose(-2, -1)之后会紧跟着乘以1 / sqrt(head_dim)这个head_dim就是dk。所以不是论文里的花架子是每一行代码都在用的东西。2. Softmax的“性格”决定了你必须scale2.1 Softmax有两个致命区饱和区和死区Softmax的公式大家都会背softmax(z_i) exp(z_i) / Σ_j exp(z_j)关键在exp函数。exp(x)的曲线是指数增长x稍微变大一点值就爆炸式增长。这就导致softmax对输入的绝对大小极其敏感。当输入是一组比较接近的小数值比如[1, 2, 3]softmax输出是[0.09, 0.24, 0.67]这还算有区分度梯度也能正常传播。但当输入是一组大数值比如[10, 20, 30]softmax输出会变成[2.1e-9, 4.5e-5, 0.99995]。最大的那个值几乎占满了所有权重其他值直接归零这就是“饱和区”。更麻烦的是“死区”效应。当输入数值很大时除了最大值那个位置其他位置在softmax之后的梯度会极度趋近于零也就是神经元“死掉”了。梯度过小意味着参数更新幅度微乎其微模型几乎学不动。我做一个直观的类比你交了一份全班同学的成绩单给老师平时成绩都分布在60~70分老师还能排出个先后。但如果全班都考到90多分放在Excel里看大家的差距反而被压缩了很难看出谁比谁强多少。2.2 点积结果的大小是怎么失控的现在回到attention的场景。Q和K的各个维度在初始化和训练过程中通常会被控制在近似均值为0、方差为1的分布里。两个dk维向量的点积本质上是在把dk个乘积项相加。由于每一项的值可正可负、均值为0这些项加在一起并不会变成0而是会像一个随机游走一样方差会累积。结果就是点积的方差等于dk标准差等于√dk。用数字说明dk64时点积标准差约8。dk128时点积标准差约11.3。dk512时点积标准差约22.6。也就是说如果dk64点积的结果大概率落在-8到8之间但如果你把维度提高到512点积结果的范围就扩大到正负22左右。输入softmax的数值大了好几个量级直接就把softmax推进饱和区。这个数学事实就是为什么原始论文里明确说“我们怀疑对于大的dk点积的幅度会变大把softmax函数推到梯度极小的区域”的原因。注意这里的“方差为1”是理想假设下的结论实际训练过程中每一层的分布都会偏移所以scale的数值并不是一个需要精调的超参数而是从理论出发算出来的“标准答案”。但理论只能保证初始分布合理真正的收敛效果还需要配合学习率、初始化等一起调。3. 除以√dk到底谁给出的最优解3.1 从方差角度看scaling的必要性假设q和k的每个维度都是独立同分布的随机变量均值为0方差为1。那么单维度乘积q_i * k_i的期望是0方差是1。点积q · k是dk个这样的乘积之和均值还是0方差变成dk。如果我们想把点积的方差稳定回1最简单的方式就是除以标准差√dk。除完之后新的随机变量的方差就是Var(q · k / √dk) Var(q · k) / (√dk)² dk / dk 1这一步操作的意义在于不管k的维度是64还是512除以√dk之后输入softmax的数值分布都保持在相对一致的尺度上。既不会因为维度大而饱和也不会因为维度小而丧失区分度。这个推导过程是理解Transformer最关键的数学环节之一面试官问这个问题本质上是在考验你对“训练稳定性”的理解程度。3.2 为什么是√dk而不是dk这是第二阶段面试官最爱追加的问题。有人会直觉想既然点积分摊那我直接除以dk不就行了吗反正也把方差缩小了。答案是不行因为除以dk会把数值压得太小。举个例子dk64时点积结果可能在-8到8之间除以√dk8后结果落在-1到1之间这对softmax来说是刚刚好的输入范围。但如果除以dk64结果落在-0.125到0.125之间所有数值都被压缩到一个极窄的区间里softmax的输出会变得过于平滑几乎接近均匀分布模型就丧失了“集中注意力”的能力。总结一下处理方式点积方差对softmax的影响不缩放dk数值过大进入饱和区梯度消失除以√dk1输入范围合理梯度稳定除以dk1/dk数值过小输出接近均匀分布丧失区分度这个表格基本就能回答案面试里的“为什么不用dk”了。不缩放的问题是梯度消失除以dk的问题是注意力分散除以√dk刚好卡在中间走的是中庸之道。3.3 换个视角看它就是个温度参数如果你对统计物理或者对比学习里的温度系数熟悉就会知道softmax里除以一个数其实是常规操作。在对比学习中那叫temperature温度公式是softmax(sim / τ)τ大分布平滑τ小分布尖锐。在Transformer里这个τ就等于√dk。所以你看attention里的scale因子本质上就是一个“固定的温度”只不过这个温度的取值不是调参调出来的而是通过方差分析算出来的理论最优值。这个视角很重要因为它把attention和对比学习、知识蒸馏等领域统一起来了。理解了这一点以后再看到任何softmax(x / T)形式的东西你都能立刻反应出它在调节分布的“尖锐程度”。4. 不scale会怎样一个暴力实验告诉你真相4.1 实验设置理论说了半天不如动手看一眼。我写了一个极简实验分别在有scale和没有scale的条件下对比attention输出的分布和梯度回传的状态。实验设定非常简单随机初始化一批Q、K、V。分别走完整attention流程。对比softmax前后的数值分布和梯度范数。我用PyTorch写个核心代码片段import torch import torch.nn.functional as F torch.manual_seed(42) batch_size, seq_len, dk 2, 10, 64 q torch.randn(batch_size, seq_len, dk) k torch.randn(batch_size, seq_len, dk) scores torch.bmm(q, k.transpose(1, 2)) # [batch, seq, seq] print(f点积结果均值: {scores.mean().item():.4f}) print(f点积结果标准差: {scores.std().item():.4f}) # 不缩放 weights_no_scale F.softmax(scores, dim-1) # 缩放 scores_scaled scores / (dk ** 0.5) weights_scaled F.softmax(scores_scaled, dim-1) print(f不缩放最大权重 {weights_no_scale.max().item():.4f}, f最小权重 {weights_no_scale.min().item():.6f}) print(f缩放后最大权重 {weights_scaled.max().item():.4f}, f最小权重 {weights_scaled.min().item():.6f})输出结果很直观dk64时不缩放的attention权重最大值会达到0.85以上最小值基本是0。缩放之后最大值降到0.25左右分布明显平滑了很多。4.2 训练中的连锁反应上面的实验只是静态观察放在真实训练里问题会更严重。我在曾经的一次尝试中把Transformer的scale去掉结果出现了两个特别典型的症状第一个是“注意力坍塌”。注意力权重迅速退化成one-hot分布每个token只关注某一个token其他位置的梯度全部消失。模型看起来学得挺快但关注到的信息非常单一泛化能力极差。第二个是“训练震荡”。因为softmax在饱和区梯度忽大忽小loss曲线来回跳无法稳定收敛。你必须把学习率调得非常小才能勉强维持训练但那样模型几乎学不动。实操心得试着在训练前把attention权重的分布打出来看一下。如果发现大部分权重都是个位数级别并且softmax之后几乎成了one-hot那大概率就是scale出了问题。这种情况在小模型上还不明显一旦模型变大、序列变长立刻原形毕露。4.3 为什么训练初期最容易踩这个坑还有一个很值得说的现象训练初期模型参数还在随机初始化状态Q和K的分布其实接近理想状态点积方差理论值和实际值差距不会太大。但如果初始化的尺度没有控制好或者学习率设置得偏大参数更新几步之后Q和K的分布就和初始状态完全不一样了。这个阶段如果不做scalesoftmax饱和的问题会被迅速放大。换句话说scale不仅仅是为了初始状态稳定更是为了在训练过程中给网络一个“安全垫”让梯度传播不至于因为局部数值波动突然失去信号。这也是为什么我在使用FlashAttention这类高效attention实现时发现它们在内部做块状计算之后依然会保留scale操作的原因。这不是性能能优化掉的是数学上的刚性需求。5. 高频追问关于scale的几个延伸问题5.1 为什么不直接对Q和K做归一化面试中常问的一个追问是既然点积可能大那我不如直接对Q和K做L2归一化让点积最大只能到1这样不是更稳吗这样确实能解决问题操作上的确可以但实际用的非常少。原因在于L2归一化会把Q和K的模长信息丢掉。在attention里向量的模长可能携带了“置信度”之类的信息强行归一化等于牺牲掉这部分信息。除以√dk是一种比较轻量的处理保留了模长信息只是从统计意义上把数值压到合理范围。另外L2归一化会限制attention的表示能力让它只能表示方向上的相似性而方向和尺度一起才构成完整的语义信息。实践里你用cosine attention也会发现效果常常不如标准scaled dot-product attention。5.2 MHA里的head_dim到底取多少多头注意力机制里每个head的维度head_dim通常是总维度除以头数。原始论文里8个头、512维每个head的维度是64那么scale就是√648。head_dim的选择直接影响scale值。head_dim越大scale越大head_dim越小scale越小。如果你用的是GQA或者MQA头的维度可能不同但scale始终跟随head_dim走而不是总维度这点千万别搞错。我见过有人写代码时不小心对全维度做了scale导致实际缩放系数远大于理论值最终模型收敛速度明显变慢。这种小坑排查起来特别费时间最好是写一个单元测试把中间层输出打印出来核对。5.3 和FlashAttention之间的关联FlashAttention把attention计算分块执行尽量减少显存读写但数学计算和标准attention完全等价。在分块过程中每个块的softmax需要做rescaling但它不是通过改scale实现的而是用running maximum和running sum做在线softmax修正。所以你在看FlashAttention源码时会看到它在内部仍然保留了softmax_scale这个参数这个参数的值依然是1/√dk。哪怕是FlashAttention 2、3这一点都没有变过。理解这个点对后续接触长文本推理、大模型加速、量化部署都有帮助。很多高性能优化的基础仍然要回到最基础的那几个公式。5.4 Softmax之前做还是之后做这是一个很容易被忽略的细节scale应该加在softmax之前而不是之后。因为scale的目的是把输入到softmax的数值压回合理范围如果加在softmax之后那只是放缩了权重向量完全起不到调节梯度的作用。# 正确写法 weights softmax(scores / math.sqrt(dk)) # 错误写法 weights softmax(scores) / math.sqrt(dk)第二种写法不但没用还会把权重向量的和破坏掉导致注意力权重不再归一。我在看一些初学者代码时确实见过这种错误特别容易发生在从别的框架迁移代码的时候。5.5 面试高频题总结这些问题的标准回答思路整理一张速查表问题核心回答为什么要scale防止点积随维度增大而过大把softmax推进饱和区导致梯度消失为什么除以√dk将点积方差归一为1让softmax输入落在合理范围为什么不除以dk会把分布压得太扁注意力权重失去区分度可以用其他方法吗可以用L2归一化或可学习的温度但通常没必要理论最优即可scale和温度有什么关系√dk等价于温度参数控制分布的尖锐程度6. 一些想跟你分享的实战心得最后说点真实的个人感受。我最初看Transformer论文时对scale这个操作基本就是一眼带过心里觉得这不就是个调节参数嘛随便设个值不都一样。直到自己复现模型时训练中期loss出现周期性抖动查了两天才定位到scale被我在代码重构时不小心弄丢了。当时打出来的attention分布几乎全是one-hot梯度范数小得离谱才真正意识到scale的重要性。之后我养成了一个习惯在实现attention的时候一定会加一行注释记录dk的来源并且在单元测试里断言attention权重的分布不会出现极端one-hot化。这个习惯帮我避免了很多后续的隐性bug。调试的时候可以做一个快速的“健康检查”拿一批真实输入跑一次前向把中间attention权重的均值、方差打出来看看有没有异常的集中。如果发现权重的信息熵特别低先怀疑scale的数值是否正确再考虑是不是模型结构的问题。如果还想再深挖建议动笔推导一下完整的梯度回传公式看看scale因子在每一路梯度上分别产生了什么样的影响。这个过程有点费时间但是做完之后你对attention的理解会有一个质的提升。