尧图精选

Transformer中Attention为何除以√d_k?缩放因子作用深度解析

🕒 发布时间:2026/10/2 15:53:01 📁 来源:尧图网络
1. 先看现象把√d_k去掉attention直接“过曝”我第一次看到Transformer里的attention公式时注意力全在那一堆矩阵乘法上根本没在意最后那个除以√d_k的操作。心想这不就是个缩放吗除以一个常数有什么好说的直到有一次自己动手从头实现一个简化版Transformer为了省事直接把scale项扔了结果训练loss曲线疯狂震荡根本压不下去。当时我还以为是学习率调得不对折腾了半天才反应过来问题就出在这个不起眼的除法上。要理解为什么必须做scale得先搞清楚一个事实attention的计算本质上是让每个Query去和所有Key做点积然后把点积分数送进softmax变成权重。也就是说softmax的输入是点积结果。如果用不用scale直接算当模型维度d_k稍微大一点比如128、512甚至更高时点积分数的取值范围会变得非常宽。这里的关键在于向量点积的大小会随着维度增长而线性增长。假设q和k的每个分量都是均值为0、方差为1的随机变量那么两个d_k维向量的点积均值是0方差是d_k。换句话说点积结果的标准差是√d_k。当d_k512时√d_k≈22.6也就是说点积结果经常能跑到几十甚至上百。你可能会觉得分数大点怎么了softmax不就是把大数变概率吗问题恰恰出在这里。softmax对输入幅度极其敏感它是一个“温度敏感”的函数。输入logits的幅度一旦偏大softmax的输出分布会迅速从“温和的加权平均”退化成“近乎one-hot的硬选择”。分布越尖锐梯度越小训练就越容易卡住。这就好比你用放大镜看东西倍数太高反而什么都看不清。一个随手就能验证的实验证明了这个现象。我写过一段简单的PyTorch代码分别算一下d_k64和d_k512时不scale和scale后的attention分布长什么样import torch import torch.nn.functional as F torch.manual_seed(42) batch_size, seq_len 2, 32 d_k 64 q torch.randn(batch_size, seq_len, d_k) k torch.randn(batch_size, seq_len, d_k) scores torch.bmm(q, k.transpose(1, 2)) print(d_k , d_k) print(scores 均值:, scores.mean().item()) print(scores 标准差:, scores.std().item()) # 不scale weights_no_scale F.softmax(scores, dim-1) # scale weights_scale F.softmax(scores / (d_k ** 0.5), dim-1) print(不scale的attention最大权重:, weights_no_scale.max().item()) print(不scale的attention最小权重:, weights_no_scale.min().item()) print(不scale的attention熵:, -(weights_no_scale * torch.log(weights_no_scale 1e-8)).sum(-1).mean().item()) print(做scale的attention最大权重:, weights_scale.max().item()) print(做scale的attention最小权重:, weights_scale.min().item()) print(做scale的attention熵:, -(weights_scale * torch.log(weights_scale 1e-8)).sum(-1).mean().item())我实际跑出来的结果大致是这样的d_k是否scale分数标准差attention最大权重熵64否8.0左右接近1.0低于1.564是约1.00.2左右2.8以上512否22.6左右几乎就是1.0极低512是约1.00.2左右2.8以上看到这个对比我相信你已经能理解为什么说“不scale会过曝”了。d_k一大attention矩阵就变成了一组one-hot向量模型只会机械地盯住某一个token其他位置的信息几乎全部丢失。这对训练来说是毁灭性的。1.1 “过曝”后的注意力矩阵长什么样把“过曝”这两个字落到具体的数据上你会发现情况比想象中更糟。当点积分数没有被缩放标准差达到22以上时softmax输出的最大概率往往直接超过0.99剩下的所有位置加起来的概率不到0.01。这意味着在反向传播的时候除了最大值对应的位置其他位置拿到的梯度都是无穷小模型几乎学不到任何有效信息。最直接的表现就是训练初期loss居高不下或者loss曲线呈现出一种非常有节奏感的锯齿形震荡。很多初学者会把这个问题归结为学习率没调好、初始化不对或者干脆怀疑是数据集有bug。但其实只要把attention里的scale加上这些表象立马就能改善大半。那年我在做一个文本分类的小项目模型结构是从网上找的简易Transformer代码里面很多细节都被简化掉了。首次训练时我盯着loss曲线十分钟看到它像心电图一样上下乱跳一度怀疑是batch size太小。排查到最后才发现那段代码里的attention实现把除以√d_k这一步给漏了。补上之后同一个模型、同一份数据loss曲线马上变得平滑收敛速度至少快了一倍。这段经历给我的教训是理解每个组件为什么存在比会写模型结构重要得多。有些细节看似无关紧要实际上是整个系统的承重墙。1.2 罪魁祸首是softmax的“温度敏感体质”softmax函数本身并不复杂输入一个向量输出一个和为1的概率分布。但它的行为对输入向量的尺度极其敏感。为了看清楚这件事我习惯把softmax和“温度参数”放在一起看。想象一个场景你是一个篮球教练要从10个候选人里挑一个首发。第一轮评审打分后大家分数普遍在70到80之间这时候选谁不选谁差别没那么大每个人都有机会。但如果打分规则变成满分1000分有人得了980有人得了750那这个差距直接把你逼到了悬崖边上你只能选那个980的。softmax的输入尺度就扮演了“打分规则”的角色。用数学语言说softmax的性能完全取决于logits的“锐度”。当logits整体偏小时输出接近均匀分布模型失去了区分能力当logits整体偏大时输出接近one-hot模型又失去了泛化能力。Transformer的attention在中间找一个平衡点而除以√d_k正是那个把点积分数的标准差拉回到1左右的调温器。这个调温器的意义并不仅仅是让softmax的输出分布更好看更关键的是它直接影响反向传播的梯度大小。softmax在极端区间的梯度会趋近于0如果注意力权重过于尖锐梯度回传就会被“截断”。Transformer是深层网络每一层的小梯度损失经过层层累加到浅层时早就消失殆尽了。所以scale不只是数值稳定性问题它直接决定了一个Transformer能不能被正常地训练起来。2. 从方差推公式为什么恰好是√d_k而不是d_k很多人知道要除以√d_k但问到为什么是这个数答案往往是“论文里写的”。在这件事上我觉得值得多花一点时间把数学地基夯实因为理解了推导过程你在设计变体模型时才知道什么时候可以动这个scale什么时候不能动。先从最简化的假设开始。假设q和k是独立的随机向量它们的每个分量都服从均值为0、方差为1的分布比如标准正态分布。两个向量做点积score q_1*k_1 q_2*k_2 ... q_{d_k}*k_{d_k}这个结果是一个随机变量。根据独立随机变量的性质它的均值是0方差是每一项方差之和。因为每一项q_i*k_i的方差是1两个标准正态分布乘积的方差为1一共有d_k项所以点积的方差就是d_k。Var(q·k) d_k std(q·k) √d_k也就是说点积分数的标准差与√d_k成正比。维度越高点积的“天然尺度”就越大。为了让这个尺度不随着d_k的变化而剧烈波动最自然的做法就是把点积除以它的标准差也就是√d_k。做完这一步点积的方差被拉回到1三个softmax的输入保持在“单位尺度”附近不管d_k是64还是128还是512attention都工作在同一个数值区间。为了让你对“方差为d_k”这件事建立更直观的体感我再跑一个快速实验import torch for d_k in [16, 64, 256, 1024]: q torch.randn(100000, d_k) k torch.randn(100000, d_k) dots (q * k).sum(dim-1) print(fd_k {d_k:5d} | 点积std {dots.std().item():.2f} | √d_k {d_k ** 0.5:.2f})输出的结果非常规律点积标准差几乎就等于√d_k。这串数字比任何公式都直观——如果你不除以√d_k那么模型结构的维度只要变一变attention的数值范围就会完全变一个量级。而Transformer往往是多层的每层都用不同维度的head如果不做scale让这些不同尺度的logits都过softmax整个模型的数值稳定性根本没法保证。2.1 余弦相似度和“模长缩放”的另一种理解再换一个角度理解这个除法。点积本身由两部分组成两个向量的模长乘积以及它们夹角的余弦值。写成公式就是q·k |q| * |k| * cos(θ)除以√d_k相当于给这个点积乘上了一个1/√d_k的因子。这个操作表面上是在缩小模长的影响实际上是在说在attention里我不希望两个向量之间的“绝对长度”来决定注意力权重而是希望由“相对方向”决定。你可能马上会想到一个替代方案既然要消掉模长的影响为什么不直接用余弦相似度也就是把点积换成(q·k)/(|q|·|k|)这个问题我专门验证过。直接使用余弦相似度确实可以把logits严格控制在[-1, 1]范围内从数值稳定性上说非常漂亮但它有一个致命的缺陷损失了向量模长的信息。q和k的模长并不是无意义的它们可能编码了token本身的“重要度”或“置信度”信息。比如某个词在句子里特别关键它的向量可能会自然学出一个比较大的模长在做attention的时候理应获得更大的权重这个信息被余弦相似度一刀切掉后模型表达能力会受到影响。除以√d_k采取的是一个折中方案不彻底丢掉模长信息只是把模长的“尺度优势”按维度进行压缩。从某种程度上说这就是一个“软化的余弦相似度”既保留了区分方向的能力又保留了模长带来的语义信息还不至于让数值爆炸。2.2 为什么除以d_k也不对有了上面的推导你可能会想那除以d_k行不行如果是想让除以之后的结果更稳定不如干脆除个狠的。我的经验是除以d_k会让attention变成“什么也看不见”的状态。假设d_k512除以d_k之后点积分数大约在0.04这个量级。所有logits都趋近于0softmax在这个区间会退化成均匀分布。也就是说每个位置对所有其他位置的注意力权重几乎相同没有任何区分度。这种情况下模型倒是能稳定训练但学出来的attention几乎是白学的因为每个token对上下文一视同仁完全失去了“注意力”的意义。从信息论的角度看均匀分布的熵是最高的它保留了最多的“可能性”但也意味着模型没有从数据中获取任何偏好。attention机制的价值恰恰在于“有选择地关注”过度的scale会让这种选择性消失。所以这个除法是“恰到好处”的除以√d_k保持方差为1既不偏向尖锐也不偏向均匀。这种居中状态让softmax的梯度处于一个比较灵敏的工作区间既能学到不同位置之间的相对差异又能把梯度顺畅地回传下去。3. scale的本质它是softmax温度参数的固定版本如果对softmax函数有更深的了解你可能会注意到一个有趣的联系除以√d_k实际上是在调整softmax的“温度”。softmax有一个常见的推广形式——带温度参数的版本softmax(z_i / τ)这里的τ就是温度参数。当τ1时输出分布变得更平滑接近均匀分布当τ1时输出分布变得更尖锐接近one-hot。知识蒸馏里就经常用高温softmax来软化标签分布。Transformer的attention公式是softmax(QK^T / √d_k)对照一下你就会发现这里的√d_k本质上就扮演了温度τ的角色。它不随训练过程变化是一个固定的温度。这个发现挺有意思原版Transformer没有把温度设置成一个可学习参数而是根据模型的维度d_k直接推导出一个固定值。这背后隐含的逻辑是只要初始化合理、数据的分布保持稳定1/√d_k这个温度在大多数情况下已经足够好用不需要让模型在训练中去额外调这个参数。3.1 softmax温度τ与注意力“锐度”的关系为了更好地说明温度对attention的影响我列过这样一张对比表在d_k64的设定下输入完全相同的Q、K温度控制方式等价操作attention分布特征适用场景不加scale相当于τ1过于尖锐接近one-hot基本不适用除以√d_k相当于τ8温和有区分度Transformer默认设置除以d_k相当于τ64接近均匀分布无区分度不可用除以√(d_k/2)相当于τ≈5.66稍微尖锐一点某些注意力“聚焦”变体这张表特别能说明问题温度的选择直接决定了attention在“探索”和“利用”之间的倾向。温度高意味着一开始每个位置都被均匀地关注训练后期模型才慢慢学会聚焦温度低意味着模型一开始就非常自信只关注少数几个位置如果这些位置恰好是错的那就是“自负害死人”。Transformer原论文选择1/√d_k并不是拍脑袋。当时作者也对比过additive attention加性注意力那种注意力实现不涉及点积缩放但计算复杂度更高。基于点积的attention在效率上有天然优势但必须配上这个除法才能在“过于自信”和“过于模糊”之间找到舒服的位置。3.2 固定1/√d_k vs 可学习温度谁更好了解了温度的含义之后一个顺理成章的问题出现了为什么不把1/√d_k设成可学习的参数让模型自己去找最佳温度这件事我在不同模型上都试过。结论是可学习温度在理论上有吸引力但实际收益并不明显而且会引入额外的训练不稳定性。我的理解是attention的logits在训练过程中本身会不断变化。早期训练时模型还没学会有效的表示q和k接近随机初始化点积分数的方差大致符合√d_k的理论值所以固定的1/√d_k恰好踩在合理的位置。到了训练后期如果q和k的分布发生了变化——比如模长整体变大或变小——这时候固定温度确实不是最优的。但实际上Transformer为每个token还配备了LayerNorm之类的归一化结构这些结构会把q和k的分布拉回到一个比较可控的范围因此固定温度的劣势被很大程度上抹平了。可学习温度的典型实现是给scale加一个可训练参数class AttentionWithLearnableTemp(nn.Module): def __init__(self, d_k): super().__init__() self.d_k d_k self.log_t nn.Parameter(torch.zeros(1)) def forward(self, q, k): scores torch.bmm(q, k.transpose(1, 2)) temperature torch.exp(self.log_t) * (self.d_k ** 0.5) return F.softmax(scores / temperature, dim-1)实际跑下来这个可学习温度经常会出现一个问题在训练的早期它的值会剧烈振荡因为模型还没学到稳定的表示时梯度无法给出一个合理的温度更新方向。与其让模型在训练初期分心去调温度不如让它集中精力去调Q、K、V矩阵。固定scale省心、稳定、效果也不差所以原版Transformer的选择至今仍是绝大多数模型的主流配置。3.3 从梯度回传看scale对训练的直接影响抛开数学和直觉再看一个最实际的层面梯度。softmax函数的梯度有一个非常优雅的性质。对输入向量z求梯度结果可以写成∂softmax(z_i)/∂z_j softmax(z_i) * (δ_ij - softmax(z_j))也就是说梯度等于softmax输出自身的函数。当某个位置的softmax输出接近1其他位置接近0时这个位置的梯度会非常小因为式子中的(δ_ij - softmax(z_j))项在其他位置上被压到了几乎为0。这就是“softmax饱和区”。在一个不scale的attention里logits动辄几十上百softmax几乎总是处于饱和区。如果你去打印attention层的梯度你会发现大量位置的梯度值小到浮点数精度都难以表达。前向传播时模型“只看得到”一个token反向传播时梯度也“只走得了”一个token。这种单向死胡同式的信息流动对一个需要捕捉长距离依赖的模型来说是致命的。加了1/√d_k之后logits被压在单位尺度附近softmax离开饱和区每个位置都能拿到有意义的梯度信息才能在长距离上顺畅流动。这也是为什么Transformer能够有效建模长序列的底层原因之一——它不仅结构上允许任意两个位置通信数值上也保证了这种通信的梯度信号不会被截断。4. 不只是原论文scale在真实工程里的那些细节理论聊透了回到工程实践。scale这个东西看着只有一行代码但它在真实框架和优化库里的处理方式包含了很多值得注意的细节。我在阅读和魔改各种Transformer实现时踩过不少和scale相关的坑这些坑值得单独拿出来说说。4.1 FlashAttention内核里的scale参数如果你用过FlashAttention应该知道它有一个专门的scale参数。这个参数的存在本身就说明scale不是attention外侧的一个附加操作它应该被嵌入到attention计算的核心流程中。FlashAttention的基本思路是分块计算不让完整的QK^T矩阵驻留显存而是通过online softmax的算法在块级别更新注意力统计量。在这个过程中每个块的点积结果都会除以scale。在标准实现里这个scale通常就是1/√d_k。一个严肃的工程陷阱是如果你在图解里先在外部把Q除以√d_k或者预先对QK^T做了缩放再把结果传给FlashAttention的同时又传入scale参数就会出现“双重缩放”。这属于那种“看起来没错、跑起来也没报错、但模型效果莫名变差”的隐蔽bug。我见过一个实际案例某同学自己实现了一个简洁版attention为了省事直接在Q上除以√d_k后来想升级成FlashAttention加速又传了一个默认的scale参数进去。模型训练出来之后BLEU分数比原来低了好几个点他一度怀疑是flash版本的数值精度问题。后来逐行对比才发现这个重复缩放把attention分布压得太均匀模型基本又退回到了“没有注意力”的状态。所以在工程上有一个值得养成的习惯检查你的attention库的scale参数是内部处理的还是需要外部传入的两者只能选其一。用HuggingFace或者PyTorch官方SDPA的话优先使用它们提供的scale参数而不要自己在外部手工缩放。PyTorch 2.0之后的SDPA接口长这样attn_output torch.nn.functional.scaled_dot_product_attention( query, key, value, attn_maskNone, dropout_p0.0, scaled_k ** -0.5, # 显式传入scale代替在q/k上提前缩放 is_causalFalse )这种写法的好处不仅是避免了双重缩放的问题更重要的是F.scaled_dot_product_attention会自动选择最合适的kernel实现比如内存高效的flash kernel、cudnn kernel或者math fallback。你只需要把scale传进去性能和安全都能得到保证。4.2 QK Norm、余弦相似度与scale的关系近两年很多模型在进入attention之前会对Q和K做额外的归一化也就是业内常说的QK Norm。这个设计和scale其实是相辅相成的但如果理解不透彻也会造成混合使用时的混乱。所谓QK Norm就是对Q和K分别做一次LayerNorm或者RMSNorm。它的作用是稳定Q和K的行向量模长让它们在训练过程中保持在一个比较可控的尺度。如果没有这一步Q和K的模长在训练中可能学到很大导致点积分数的绝对值整体变大即便除以√d_k也无济于事。加上QK Norm之后q和k的模长被限制在一个标准范围此时再除以√d_k才能保证logits始终处于理想的分布区间。但这里有一个关系需要理清QK Norm和scale不是二选一而是互补的。QK Norm控制的是单个向量的模长scale控制的是点积结果的整体方差。两者关注的对象不同叠加使用不会冲突反而能让数值更稳定。有些人尝试用QK Norm彻底替代scale也就是做完QK归一化之后不再除以√d_k。这种做法在部分实现里确实可以工作但并没有理论上的必然优势。因为即便单个q和k的模长被归一化了两个高维向量的点积仍然会随着维度增加而逐渐偏离0。余弦相似度本质上就是俩单位向量的点积它对维度仍然有一种偏向性——维度越高两个随机单位向量更容易落在彼此正交的方向上点积的标准差会变得很小。这个现象在多维几何里有一个很直观的解释高维空间中随机向量的夹角普遍趋向于90度。所以在实际建模中比较稳妥的组合是对Q和K做RMSNorm或LayerNorm保持特征尺度稳定保留1/√d_k的scale维持点积分数的标准差在1附近必要时再加上可学习温度或位置偏置这个组合在多个开源模型里已经成了默认配置比如某些ViT变体和多模态模型都采用了类似的思路。4.3 实际实现中scale、mask、bias的顺序问题最后一个工程细节是顺序。在注意力计算中scale、maskpadding mask或因果mask和位置偏置如ALiBi、RoPE生成的bias三者之间的先后顺序不是一个可以随意的选择。我推荐的标准顺序是# 1. 先算点积 scores torch.bmm(q, k.transpose(-2, -1)) # 2. 再做scale scores scores / math.sqrt(d_k) # 3. 再加位置偏置如果有 scores scores bias # 4. 最后加maskmask区域填负无穷 scores scores.masked_fill(mask 0, float(-inf)) # 5. softmax attn_weights F.softmax(scores, dim-1)为什么mask要放在最后因为mask需要把不需要的位置填成负无穷如果先做softmax再mask那些位置仍然会被赋予非零概率语义就错了。而scale和bias放在mask之前是为了保证所有位置都先被“调温”然后被mask硬性屏蔽的才彻底屏蔽顺序错乱不会直接导致报错但会引入不易察觉的数值偏差。另外一个更隐蔽的问题是如果bias的值非常大直接加到scores上即使做了scalesoftmax仍然可能进入饱和区。ALiBi里的距离惩罚项在某些长序列上会累积出较大的值如果不加限制attention依然会变得过分尖锐。这种场景下需要的往往不是调整scale而是对bias本身做截断或缩放不要指望一个固定的1/√d_k能解决所有问题。5. 常见误区与踩坑记录这一节我打算专门整理一下和scale相关的常见误区都是我在自己写代码或帮别人debug时真实遇到过的。这几个坑隐蔽性很强通常不会让代码直接报错但会让模型性能莫名下降。5.1 误区把d_model和d_k搞混这是最常见的bug没有之一。很多人在实现Multi-Head Attention时会错误地使用d_model作为缩放因子而不是单头维度d_k。一个标准的多头注意力设置下假设d_model768num_heads12那么每个head的维度d_k64。正确的scale应该是1/√641/8。如果你错用了1/√768相当于多除了一个√12≈3.46attention分布会被过度压平模型的学习能力会显著下降。这个bug很难通过报错发现因为代码完全能跑loss也可能在慢慢下降只是最终的效果始终达不到预期。排查时很多人会去调学习率、改层数很少有人想到scale用错了维度。从这点也能看出理解公式里每个变量的含义比抄代码重要得多。5.2 误区FP16下先乘再除导致上溢在混合精度训练的时候如果把attention的前向计算放在FP16下进行要格外小心数值范围。一个很常见的不规范做法是# 不规范 scores torch.bmm(q, k.transpose(-2, -1)) scores scores / math.sqrt(d_k)这看起来没什么问题但如果你是在AMP自动混合精度环境下torch.bmm产生的结果会进入FP16自动放大的范围。当d_k比较大时点积分数的绝对值可能超过FP16的最大表示能力65504一旦上溢scores里会出现NaN或者Inf。更稳妥的做法是先对q做scale再做矩阵乘# 更稳妥 q_scaled q / math.sqrt(d_k) scores torch.bmm(q_scaled, k.transpose(-2, -1))这样点积计算过程中的数值就始终被约束在安全范围内。这个例子也说明scale放在哪一步做不是纯粹的数学等价问题在实际数值精度下不同顺序可能带来完全不同的结果。5.3 误区把scale设成可学习参数后loss爆炸前面提到可学习温度在理论上可行但工程上需要非常小心。我曾经在一次实验里把scale设置成可学习参数初始化为1/√d_k然后允许模型微调这个值。结果训练到十几步之后loss直接爆掉了数值变成NaN。后来看了梯度才发现scale的梯度在训练早期非常大导致它迅速偏离合理区间把整体logits放大到softmax饱和区失去了梯度信号。如果确实想要可学习温度一个更安全的做法是把它约束在log空间并且加一个正则项限制它不要离初始值太远log_t torch.nn.Parameter(torch.zeros(1)) # 在loss里加入正则 reg 0.01 * (log_t ** 2)但说实话基于我自己的实验经验绝大多数场景下固定scale就够了。与其让模型在训练早期分心去学温度不如把这份建模能力放在Q、K、V矩阵上。可学习温度更适合那些特殊任务比如某些对比学习场景或检索场景这些场景需要显式地调节注意力的锐度。结尾前的一段实务总结如果你也要从零实现或者魔改Transformer我的建议是保持最简单的固定scale也就是1/√d_k不要擅自去掉也不要在没有充分验证的情况下把它改成可学习参数。这个看似普通的除法背后是整个attention机制数值稳定性的地基之一。如果你在调试一个既有模型时发现效果不对可以第一时间去检查attention实现的这些细节scale是否正确、是否被重复施加、mask和bias的顺序是否正确、FP16环境中是否存在上溢风险。这几个点排查完往往能解决很多玄学级别的性能问题。我自己在接触Transformer的早期曾因为忽略scale导致模型无法收敛也因此花了不少时间去理解背后的数学原理。那段debug经历虽然绕了远路但让我真正对这个模型组件建立了直觉。希望你不用再踩一遍我踩过的坑。
上一篇/下一篇内容由系统自动关联 返回资讯列表 →