把 softmax 量化进注意力后,训练能否从零跑通取决于反向规则:去掉行极值梯度会延迟发散,MinMax 配 Weight-STE 在 2.5B token 落后 0.89 nats
核心概要
该工作研究在预训练中用量化 softmax(K-interval attention,用 K+1 个网格值近似指数)替代精确 softmax,推导了包含校准导数的反向规则,并在模型、数据、优化器匹配的预训练实验中比较逐行网格校准(MinMax 与固定窗口 FWM)、插值(LERP)与硬取整(Nearest)、以及直通代理放在归一化前(Weight-STE)还是后(Prob-STE);结果显示前向不变但把行极值 detach 会导致 25–30M token 后延迟发散并最终高出 0.65–3.07 nats,硬取整下 MinMax 配 Weight-STE 在 250M token 时高出 0.89 nats、2.5B 时高出 0.89 nat
深度剖析
同一个前向计算可以仅因反向规则不完整而训练失败:把行极值 detach 后前向数值不变,但丢掉两个雅可比项并破坏平移不变性所要求的零和恒等式,训练先与 softmax 同步 25–30M token,随后发散,最终比带完整校准梯度的匹配运行高出 0.65–3.07 nats,并在另外两个随机种子上复现。 此前把减去行最大值当作数值便利、其梯度精确抵消,这一习惯被沿用到带校准的网格上;该工作把校准统计量放回自动微分图,并给出闭式校准导数与零和恒等式,指出 detach 会引入无法改变损失方向的伪共模分量。 LERP 的完整梯度运行与 softmax 无可检测差异(95% 置信区间包含零),同一前向的 detach 运行最终高出 0.65–3.07 nats;失败在 25M token 后出现、约 30M 时分离,并在种子 7 和 42 上复现;Nearest 在每个 K 上都相对其匹配运行失败。
把 detach 后的梯度投影到零和子空间并不能修复结果:投影精确去掉了共模分量,但该运行最终比朴素 detach 更差;而只保留依赖行最大值的校准梯度可恢复几乎全部完整–detach 差距,只保留依赖最小值的梯度几乎恢复不了。 这区分了“零和恒等式被破坏”与“行最大值通道缺失”两种解释,说明零和本身不足以解释结果,并把主导通道定位到行最大值一侧。 P1 投影运行前向与 detach 逐位相同,最终更差;最大-only 反向在两个种子上都落在完整梯度 0.015 nats 以内,最小-only 几乎无恢复;20M 与 40M 的固定上游分解结论一致。
在硬取整(Nearest)下,校准策略与代理位置强烈交互:MinMax 配 Weight-STE 在 K=4 时 250M token 高出 0.89 nats、2.5B 时高出 0.89 nats;把代理移到概率层(Prob-STE)或改用固定窗口 FWM 各自消除大部分差距,两者同时改则增益很小,说明交互大且为正。 重建后归一化的算子产生了“权重层代理/概率层代理”这一新选择,两者共享同一硬前向但归一化雅可比求值点不同;此前直通估计器在 Gumbel-softmax 中只有概率层一个位置。 250M token 的平衡因子设计(校准×代理)给出主效应与交互;方向在 1B 参数和全部五个种子上保持;五种子平均下 FWM 相对 MinMax 降低 0.067 nats、Prob-STE 相对 Weight-STE 降低 0.021 nats,条件排序 Kendall's τ 为 0.91。
粗粒度确定性重建未必妨碍接近基线的性能:LERP 在测试的 K 上 2.5B token 时与 softmax 相差约 0.01 nats 以内,FWM–Weight 在 K=16 时同样接近;但小 NLL 差距不等于下游等价,大差距条件在所有基准配置上都低于 softmax,而 0.01 nats 以内的条件在基准上出现双向的小幅任务相关差异。 该工作把“近似 softmax 能否用于预训练”从单点误差预算问题转为校准、重建与反向规则三者联合评估的问题,并给出训练预算会改变硬取整与插值相对优劣的观察。 2.5B token 的 15 个运行、1B 参数 100M token 的 12 个运行、五种子 130 个运行;下游在 WikiText-103、PTB、C4 上排序保持(Spearman 0.952–0.996),六个基准七种配置中 35 个大差距比较全部为负、27 个达到显著,近基线组 35 个比较中 8 个达到显著(6 负 2 正)。
启示与展望
该结果面向在预训练中替换 softmax 归一化算子的系统设计者,适用设定是 GPT-2 风格 124M 与 1B 参数模型、FineWebEdu-3B 与 WikiText-103 语料、最多 2.5B token 的训练预算,注意力算子以 fp32 运行、训练时启用 TF32 矩阵乘。它使后续工作可以在同一前向下比较校准策略与代理位置,并把行最大值梯度作为必须保留的通道;LERP 与 FWM 的若干配置可作为接近基线的候选,用于进一步的下游与硬件评估。
套件 A–C 各只有一个种子,套件 D 的种子离散度只适用于其自身设定;尾部策略比较为单种子且非等成本,通道消融为两个种子,严格前向重训与降精度基线各为 100M token 单种子。两个代理在数学上共享同一硬前向,但历史实现并非逐位相同,冻结检查点上平均验证 NLL 差异最大约 0.001 nats,单块最大可达 0.0125 nats,因此代理对比带有前向舍入混淆。窗口宽度 τ 由冻结权重上的推理期筛选确定,其最终阶段使用的 WikiText-103 验证集后来也用于套件 C 和 D,未证明对训练最优。范围–分辨率反馈假设未被任何干预识别,几何量在 detach 失败中先于损失变化、在正式 MinMax–Weight 运行中后于损失变化,因此不是通用预警。下游近基线组出现双向小差异,单种子条件下无法归因于算子本身。
