
做NLP的朋友应该都被Transformer里的attention那一套公式“洗礼”过。我第一次照着论文手写Scaled Dot-Product Attention的时候最让我不解的就是那个scale——为什么QK^T后面非要除以一个√dk这个除法到底在干什么。后来调试模型踩了几次坑又回头啃了一遍数学推导才彻底把它吃透。如果你也在手写Transformer、复现论文或者调注意力机制这篇文章就是来解决这个问题的我会从统计、梯度、数值稳定性三个层面拆解scale的作用再带你用代码做个小实验亲自看看缩放前后的差异有多大。这背后不止是“防止数值溢出”这么简单它其实直接关系到attention能否顺利训练。理解了它你对Transformer的理解深度会上一个台阶以后看各类注意力变体也会轻松很多。1. 先定位scale在attention公式里的精确位置1.1 从Scaled Dot-Product Attention公式说起标准Transformer里注意力计算的公式长这样attention(Q,K,V) softmax(QK^T / √dk) V其中Q、K、V分别是query、key、value矩阵dk是每个注意力头的维度。Q和K做点积得到相似度分数然后经过softmax转成概率权重最后和V加权求和。公式里那个“除以√dk”的操作就是我们说的scale。很多代码实现里它写成scores torch.bmm(q, k.transpose(1, 2)) / math.sqrt(dk)也有的写法是提前算好一个缩放系数乘上去scale 1.0 / math.sqrt(dk) scores torch.bmm(q, k.transpose(1, 2)) * scale两种写法本质一样只是把除法变成了乘法运算。重点在于这个缩放必须发生在softmax之前因为它的目标是“预处理”那些还没变成概率的logits分数。有人可能会问为什么是dk而不是其他维度因为每个注意力头里Q和K的最后一维是dk点积是在这一维上求和。所以这个维度才是点积分数的“膨胀源”跟序列长度、batch大小、模型总隐藏维度都没有直接关系。这也是为什么head_dim通常被设为64或者128这类值的时候scale特别重要。1.2 新手实现时最常见的两种错误我见过不少初次手写attention的人在这两步容易出错。第一种把除法整个漏掉了。直接用QK^T进softmax结果训练时loss下降特别慢甚至根本不下降。一开始你可能以为是学习率调得不对但实际根源在注意力分布已经“坏掉”了。第二种把除法放在了softmax之后。比如先softmax(QK^T)再除以√dk最后乘V。这完全错了。softmax之后向量已经是概率分布整体再除一个常数只会把所有注意力权重同比例缩小最后乘V会让输出整体被缩小破坏整个注意力机制的意义。scale应该也只应该作用在softmax之前的logits上。2. 问题根源为什么点积会“随维度膨胀”2.1 一个容易被忽视的统计现象要理解scale的必要性先看一个非常基础的统计事实。假设Q和K向量里的每个分量都是从均值为0、方差为1的分布里独立采样出来的那么对于单独一个维度q_i乘以k_i的期望是0方差是1。为什么因为E[q_i k_i]E[q_i]E[k_i]0而Var(q_i k_i)E[q_i²]E[k_i²]1。现在把dk个这样的乘积加起来s Σ_{i1}^{dk} q_i k_i每一个乘积的方差都是1加起来之后方差的线性叠加让整体方差直接变成了dk。也就是说Var(QK^T) dk标准差 √dk这个结果非常关键。在Transformer常见的配置里dk64的话点积分数的标准差就是8dk128时标准差是11.3。这还不是极端情况只是“典型的波动范围”就已经达到个位数到十几了。换句话说QK^T得到的相似度分数天然就会随着dk的增大而越来越大这个膨胀是统计规律躲不掉的。你可以把每个维度想象成一次独立的“评分”每个评分都是小范围的但dk个评分加在一起总分范围就会被放大√dk倍。就像你测量一个物体测量dk次每次误差都很小但误差会累积总误差会随着测量次数增长。2.2 几何视角高维空间中向量内积的表现从几何上看点积本身反映的是两个向量的“方向是否一致”和“长度是否够大”。Q和K向量在初始化阶段并不是单位向量它们的模长大约是√dk这个量级。两个独立随机向量的夹角在高维空间中会接近90度余弦值很小但是模长在增长所以点积整体既不会稳定在一个很小的范围也不会是0附近。常见的误解是既然两个随机向量在高维空间里接近正交那点积应该接近0才对。实际上“接近正交”指的是夹角的余弦接近0但我们算的是余弦乘以两个向量模长的结果。模长本身就随√dk增长所以就算方向几乎垂直点积的标准差依然保持在√dk的量级。这个几何直觉和统计推导是完全一致的。2.3 原论文怎么说Transformer原文其实已经点出了这个问题只是当时作者们用了比较谨慎的表述。论文里是这样写的We suspect that for large values of dk, the dot products grow large in magnitude, pushing the softmax function into regions where it has extremely small gradients.翻译过来就是我们怀疑当dk比较大时点积的数值会变得很大从而把softmax函数推入一个梯度非常小的区域。“suspect”这个词说明当时作者更多是凭直觉判断但后来的研究和大量实验都验证了这个猜测。所谓“梯度非常小的区域”其实是softmax最危险的地方也是下一章要展开的内容。3. 不缩放会怎样陷入softmax的“死亡区域”3.1 softmax对“大输入”的敏感性逐步消失softmax的本质是先把logits取指数再归一化成概率。假设有两个logits一个是0另一个是0.2那softmax输出大约是[0.45, 0.55]两者区别不大模型能感受到微小的差异。可如果一个是0另一个是10呢exp(10)大约22026而exp(0)1归一化之后第二个位置的概率几乎就是1第一个位置几乎就是0。这就是softmax的性质输入之间的差距越大输出就越接近one-hot分布。当QK^T不缩放时scores的标准差是√dk级别意味着同一行里不同位置的分值差个几分甚至十几分是很常见的事。一旦出现这种差距softmax直接就“锁定胜者”了其他位置的注意力权重约等于0。3.2 梯度饱和的数学证据光说“梯度变小”可能不够有说服力来看一下softmax的梯度到底是怎么算的。假设softmax的输入是z输出是p那么雅可比矩阵的每个元素是J_ij ∂p_i / ∂z_j p_i (δ_ij - p_j)这里δ_ij是克罗内克函数。当softmax输出接近one-hot时假设p_k≈1其他p_j≈0那对角线上的项p_k(1-p_k)≈0非对角线项-p_i p_j也≈0。换句话说整个雅可比矩阵的每一项都在迅速变小梯度范数趋近于0。这意味着什么反向传播的时候从损失函数传回来的梯度经过softmax这一层就被“截断”了几乎传不到Q和K上面。attention分布已经是“赢家通吃”的状态你说这模型还怎么学它根本得不到有效的反馈来调整Q和K注意力机制就退化成了一个固定的、极端的筛选器。实际训练中如果去掉scale最常见的现象是模型在最初的几步loss几乎不动或者下降得极其缓慢勉强训练一段时间后效果也明显不如带scale的版本。而且这种情况不是调大学习率就能解决的因为问题出在梯度信号本身已经消失了放大学习率只会让其他层震荡并不会让softmax区域的梯度起死回生。3.3 数值上更容易溢出还有一个非常现实的问题数值溢出。exp的增长速度是指数级的exp(20)已经超过4.8亿exp(30)接近10的13次方exp(88)在float32里直接就变成inf了。之前说过dk64时scores标准差是8也就是说logits超过20甚至30并不稀奇一旦进入softmax计算exp中间结果随时可能变得极大。float32的情况还好一些至少在exp(30)这个量级还没溢出但float16就完全不行了。float16的最大值是65504exp(20)4.8亿早就爆了。这就是为什么在混合精度训练或者FP16推理时你会看到各种奇怪的问题比如loss变成nan。归根结底很可能是某个attention层在softmax前没有正确应用scale。所以你要明白scale不是锦上添花的优化技巧它直接决定了attention能不能稳定训练以及能不能在低精度环境下正常工作。4. 除以√dk之后到底发生了什么4.1 方差“保值”从dk回到1既然不缩放的根源是方差膨胀那除以√dk正好把这个膨胀抵消掉。回到第二章的推导如果s的方差是dk那么Var(s / √dk) dk / (√dk)² dk / dk 1也就是说缩放之后的分数标准差从√dk降到了1左右。scores的典型范围变成了一个以0为中心、波动幅度只有几个单位的分布。在这个范围里不同位置之间的logits差异通常是3到4个标准差以内也就是几分的差距。可别小看这个变化。在几分的差距下softmax的输出还能保留足够的“竞争空间”不会一发入魂直接变成one-hot。模型依然能从softmax的梯度中学习到“分数高一点还是低一点”的反馈梯度信号生存下来了。4.2 softmax落到了“信号敏感区”这里可以借助一个温度系数的概念来理解。softmax函数其实可以写成一个带温度参数T的形式softmax(z / T)T越大分布越平滑T越小分布越尖锐。在标准Transformer里QK^T除以√dk本质上就是在用softmax(z / √dk)。因为√dk通常大于1这个操作相当于给原来的logits“降温”其实是加大分母让分布更平滑让attention在一开始不会那么极端。你可能会想attention难道不应该越尖锐越好吗注意力集中不是更好吗问题是尖锐应该是训练过程中学出来的而不是初始化时就固化的。一开始模型还没学会哪些信息重要如果注意力权重直接进入one-hot状态那模型就只能从少数几个token里获取信息其他位置的梯度全断了。这就像你刚入职还没弄明白哪些事重要就只盯着一件事干很大的概率会漏掉关键信息。合适的scale让初始化阶段的attention分布保持一定熵既保留区分度又给训练留出调整空间。4.3 数值稳定性的“额外红利”除了优化层面的好处把logits控制在合理范围还给数值计算带来了极大便利。一方面exp的中间结果不会动不动就爆掉FP16训练成为可能另一方面像FlashAttention这类高效实现它们在做online softmax时需要在累加过程中不断更新最大值和归一化指数如果logits一开始就很大数值误差也会被迅速放大。scale在这里相当于给整个计算过程做了一个“数值预归一化”让后续运算都处在一个安全区间。后面你会看到几乎所有主流的高效注意力实现无论是FlashAttention还是各种kernel融合都保留了这个scale因子。它不是某个实现的小偏好而是整个注意力家族公认的标配。5. 为什么偏偏是√dk数字背后的推导5.1 严格推导在什么假设下成立很多人记住结论“要除以√dk”但不清楚这个√dk是怎么来的。我把完整推导写一遍你以后就不用在网上翻来翻去了。假设q_i和k_i是独立同分布的随机变量均值为0方差为1。那么对于单独的一项q_i k_iE[q_i k_i] E[q_i] E[k_i] 0E[(q_i k_i)²] E[q_i²] E[k_i²] 1 * 1 1所以 Var(q_i k_i) E[(q_i k_i)²] - (E[q_i k_i])² 1 - 0 1由于不同维度之间独立加总之后Var(Σ_{i1}^{dk} q_i k_i) Σ_{i1}^{dk} Var(q_i k_i) dk因此点积的标准差为√dk。要想让这个标准差变回1最自然的方法就是把所有分数同时除以√dk。这里有一个隐藏假设q_i和k_i是独立的而且每个分量的方差为1。实际训练中Q和K经过网络层和LayerNorm之后并不是严格的“独立零均值单位方差”但Transformer各种初始化方案在设计时都会尽量让输入到attention的向量保持单位方差量级所以这个推导在工程上仍然近似成立。这也是为什么标准实现里直接用1/√dk而不需要每次在模型里额外统计方差的根本原因。5.2 为什么不用dk、2√dk或者其他既然要让方差变回1那除以√dk看起来就是最自然的选择。你可能会问除以dk行不行当然也行但效果完全不同。如果除以dk那方差就变成1/dk。以dk64为例方差是1/64标准差只有0.125。这会导致所有scores都挤在0附近非常窄的区间里softmax输出的分布会接近均匀分布。不同token之间的注意力权重几乎一样相当于模型在做“平均池化”注意力机制的核心价值就丧失了。一个差不多的类比是你把所有候选人的得分都压缩到0.099和0.101之间最后大家的评分都差不多那排序还有什么意义那如果用2√dk或者更大的scale呢同理过度压缩会让区分度进一步下降attention的表达能力被限制。过犹不及。所以√dk不是随便挑的它是理论上最匹配“方差归一化”这个目标的值。5.3 和temperature的关系以及可学习scalescale和temperature本质上是一回事。你可以把attention的logits缩放理解成一个固定温度下的softmax采样temperature越小分布越尖锐temperature越大分布越平滑。标准Transformer的T√dk正好处在一个温和的位置。有的模型曾尝试把scale改成可学习的参数让网络自己决定什么温度最合适。理论上这听起来很灵活但在实际工作中标准Transformer固定用1/√dk的变体占绝对主流。原因很简单可学习的scale等于给网络增加了一个自由度但这个自由度没有带来多少增益反而可能引发训练不稳定。尤其在小数据集上网络容易把这个参数学到极端的值导致注意力分布过尖或者过平。固定scale简单、有效、可复现在绝大多数场景下都够用。6. 动手验证一个简单实验看scale的作用6.1 最小复现代码理论说了一大堆不如亲手做个实验直观。下面这段PyTorch代码不需要训练任何模型只需要随机生成Q和K就能观察缩放前后的差异。我用随机初始化来模拟attention刚刚开始训练时的状态。import torch import torch.nn.functional as F torch.manual_seed(42) batch_size 8 seq_len 128 dk 64 q torch.randn(batch_size, seq_len, dk) k torch.randn(batch_size, seq_len, dk) # 未缩放 scores_no_scale torch.bmm(q, k.transpose(1, 2)) # 缩放后 scale 1.0 / (dk ** 0.5) scores_scale scores_no_scale * scale print(未缩放方差: {:.4f}.format(scores_no_scale.var().item())) print(缩放后方差: {:.4f}.format(scores_scale.var().item())) probs_no_scale F.softmax(scores_no_scale, dim-1) probs_scale F.softmax(scores_scale, dim-1) def avg_entropy(probs): eps 1e-12 return -(probs * (probs eps).log()).sum(-1).mean().item() print(未缩放平均熵: {:.4f}.format(avg_entropy(probs_no_scale))) print(缩放后平均熵: {:.4f}.format(avg_entropy(probs_scale))) # 用一个随机权重模拟下游传来的梯度观察梯度范数 w torch.randn(batch_size, seq_len, seq_len) loss_no (probs_no_scale * w).sum() grad_no torch.autograd.grad(loss_no, scores_no_scale)[0] loss_scale (probs_scale * w).sum() grad_scale torch.autograd.grad(loss_scale, scores_scale)[0] print(未缩放梯度范数: {:.4f}.format(grad_no.norm().item())) print(缩放后梯度范数: {:.4f}.format(grad_scale.norm().item()))随机种子固定为42时你会得到一个非常清晰的结果大致如下指标未缩放缩放后scores方差约64约1softmax平均熵很低接近0.1较高约3.6平均梯度范数明显更小明显更大6.2 结果解读未缩放时scores的方差基本就等于64和理论推导完全吻合。更关键的是softmax的平均熵未缩放时非常低说明softmax已经产生了极端尖锐的分布大量token的注意力权重归零。缩放之后的熵明显更高注意力分布还保留着充分的“不确定性”。梯度范数的差异更直接。未缩放情况下梯度信号的范数只有缩放后的一小部分这意味着反传到QK的更新信号会被削弱。当你真的去训练一个Transformer这种削弱在每一层都会叠加最后就是loss卡住不动训练效率极低。这个实验已经把问题从“理论上可能”变成了“肉眼可见”。6.3 工程上容易踩的坑结合我自己的一些调试经历这里说几个真实项目里容易翻车的地方。第一个坑混合精度训练下的溢出。前面提到fp16最大只能表示到65504如果attention的attention层忘记scaleQK^T的logits很容易超过这个范围。最典型的现象是训练跑着跑着loss突然变成nan很多人第一反应是学习率太大但其实元凶是attention层数值溢出。遇到nan先别急着调参检查一下有没有漏掉scale。第二个坑用的是PyTorch自带的F.scaled_dot_product_attention但它通过参数显式传入scale。这个API默认情况下会自己处理缩放但如果你从别的框架移植代码需要确认一下当前版本是否自动帮你乘了1/√dk。如果框架版本差异导致行为不同在对比实验结果时很容易造成莫名其妙的差异。第三个坑手写attention时scale的位置必须放在bmm之后、softmax之前。很多时候把这行代码写错位置模型依然可以训练但效果会莫名其妙地差。调试这类问题最快的方式就是打印一下attention输出的分布看看有没有异常尖或者异常平的趋势。7. 常见疑问速查与高阶讨论7.1 Q/K的初始化为什么也跟√dk有关有一个常见问题是既然都要做scale那为什么Q和K的初始化还要关注√dk其实这两个问题是相关的。如果Q和K的每个分量方差不是1而是某个σ²那么点积的方差就会变成dk·σ⁴这时候理论上应该除以σ²√dk而不是√dk。好在标准的初始化方法无论是Xavier还是LeCun正态初始化都会把注意力层输入输出的方差控制在1附近这让固定使用1/√dk成为一个非常稳健的选择。实际上即使初始化导致q和k的经验方差略微偏离1固定scale的注意力层也具备一定的鲁棒性因为后续的残差连接和LayerNorm会兜底。但如果你自己设计了一种特殊的初始化最好先估算一下Q/K分量的方差必要时再调整scale避免一开始就让注意力分布进入极端状态。7.2 为什么要对QK^T做scale而不对V做同样处理很多人一开始会疑惑QK^T被缩放是因为点积很大那V为什么不缩放这要从softmax的特殊性说起。softmax会对整行做归一化所以无论QK^T的绝对数值有多大最终注意力权重之和都等于1。缩放QK^T影响的是“权重的分布形态”以及“梯度是否饱和”而不是权重的整体大小。V则完全不同它直接跟注意力权重做加权求和它的绝对尺度决定了输出向量的整体大小而这个大小交给后续的残差连接和LayerNorm去统一管理就够了。如果给V也加一个固定的scale那只会把输出向量整体变小没有任何正向作用。7.3 主流预训练模型和框架里到底怎么写的在主流实现中scale的写法基本一致。BERT各系列变体、GPT系列、T5等模型在多头注意力里用的都是1/√dk这个固定值。以BERT base为例每个注意力头head_dim64scale就是1/8这也就是为什么很多老代码里你会看到scores除以8而不是写一长串数学公式。GPT-2也遵循同样的设计。后续出现的开源框架例如HuggingFace Transformers、PyTorch的nn.Transformer内部实现里都保留了这个除法。后来的一些注意力变体比如cosine attention因为把相似度限制在[-1,1]区间天然不存在方差膨胀的问题可以去掉scale。Linear Attention、Performer这类去掉softmax的变体会重新设计相似度计算方式scale的处理也各不相同。这反过来证明了一件事只要用softmax就必须面对logits尺度的问题一旦绕开softmaxscale的设定逻辑就完全改变了。7.4 FlashAttention等高效实现中的scaleFlashAttention这类算子之所以能在长序列上跑得飞快主要是靠分块计算和online softmax来降低显存访问。但这个过程中scale依然是必不可少的一环。在online softmax里每个分块都需要维护自己的最大值和求和项如果logits一开始就很大累积过程中数值误差会被迅速放大。因此FlashAttention的实现会在计算QK^T之后立即乘以scale因子再进行后续所有操作。你看到的一些高效注意力库比如flash-attn在调用时可能会要求传入scale参数默认值就是1/√dk。如果你在迁移模型时忘记了这个参数可能只是推理结果偏离一点但训练时却会出现各种奇怪的精度问题。这提醒我们凡是跟attention相关的优化实现scale都是核心细节不能把它当成可有可无的配置项。7.5 attention分数缩放对长序列场景的影响最后再补充一个长度相关的现象。序列越长attention矩阵里每一行的元素就越多未缩放时出现极端logits的概率也会变大。也就是说长序列场景下不缩放带来的影响可能比短序列更明显。这也可以解释为什么一些早期工作在处理长文本时会额外关注scale或者使用更平滑的注意力变体。虽然标准Transformer里的scale是固定的但理解这一层关系对你在长文本场景下调优模型是有帮助的。在我自己的实践里现在每次写注意力代码都会把scale单独拎出来作为一个变量比如self.scale 1.0 / math.sqrt(self.head_dim)。调试的时候我会故意把它改成1.0或者更小的值观察loss曲线的变化用来判断注意力机制是否真的在发挥作用。大多数情况下去掉scale之后模型的收敛速度和效果都会明显变差如果你发现它完全没影响那多半意味着任务太简单或者模型压根没有好好利用注意力机制这时候反而值得警惕。这个简单的“变量替换法”是我觉得最实用也最容易上手的排查技巧。