Pretraining Transformers with Quantized Softmax in Attention:把 softmax 也压成低精度后再算反向
- 关联论文:2609.33591
- 作者:flyP
- 更新:2026-09-30
一句话结论
这篇工作系统研究了预训练阶段用 K-interval 近似 softmax(即把 exp(x) 离散到 K+1 个网格点)对 Transformer 注意力做低精度化的影响:在 124M 参数 / 2.5B 训练 token 的受控对照实验中,作者发现反传时是否正确处理校准梯度以及straight-through surrogate 放在 normalization 之前还是之后才是 loss 缺口的关键,而"是否做 hard rounding"与"是否 detach 行极值"只造成温和影响;最终用 fixed-window 校准 + post-normalization surrogate 在 K=4 时把验证 loss 缺口压到 +0.019 nats,用 pre-normalization surrogate 在 K=16 时压到 +0.004 nats。
解决什么真问题
低精度 Transformer 系统已经在量化 attention 的矩阵乘法(QK^T 与 AV),但 softmax 经常仍保持高精度。直觉上 softmax 不是线性算子,它的近似既改 forward(注意力权重)也改 backward(梯度流),二者耦合起来后果难猜。
因此本文不是"再压一点"的工程稿,而是把"softmax 低精度化"作为一个研究问题来回答三个互相耦合的设计选择:
- 行内网格校准策略:每行用一个固定区间、滑动区间还是 min-max?
- 离散化方式:interpolation(连续插值)还是 hard rounding(直接归到最近网格点)?
- Straight-Through Estimator (STE) 的放置:在 row normalization 之前,还是之后?
论文控制模型、数据、优化器不变,逐一翻转这三个旋钮,比较验证 loss。
核心方法:K-Interval Attention
1. K-interval 近似
把 softmax 中的 exp(x) 用 K+1 个网格值近似:
- 把一行 attention logits 划分到 K 个区间。
- 每个区间用 K+1 个网格值代表 exp 在该区间的取值。
- 离散化有两种选择:
- interpolation:logit 在区间内时返回两端值的线性插值(连续可微)。
- hard rounding:直接归到最近一端的网格值(梯度需要 STE 近似通过)。
2. 行内网格校准(Calibration)
网格的"位置"决定 K+1 个代表值取哪些 exp 值。三种候选:
- Fixed-window:所有行用同一固定区间(如 [-T, L]),与本行取值无关。
- Per-row min-max:每行用本行的 min/max 作为网格两端,适配性强。
- 滑动窗口 / adaptive:固定窗口 + 端点更新规则(论文 variants,abstract 未详述)。
3. Straight-Through Surrogate 的放置
Forward 时用 K+1 个网格值近似 exp,但反向时梯度怎么走?常见做法是 straight-through estimator:把"通过离散化算子"的梯度按 identity 传过去。但 surrogate 放在哪一步影响梯度流:
- Pre-normalization surrogate:STE 在 row normalization 之前。梯度直接穿过"离散化 → 归一化"两步。
- Post-normalization surrogate:STE 在 row normalization 之后。归一化已经被精确算出来,只有"归一化前的离散化"用 STE。
直觉上 pre-norm 让梯度更"原汁原味",但论文发现:在 K=4 + hard rounding + min-max 校准时,pre-norm 与 post-norm 的 loss 差距巨大。
4. 反向传播规则
作者推导了对应的 backward 规则,包括校准梯度(calibration derivatives)--即网格端点本身对 loss 的梯度。这是论文的技术贡献之一:把"网格位置也是可学习量"这件事用受控实验暴露出来。
伪代码骨架(自论文抽象简化):
# Forward (single row of attention logits x, shape (N,))
grid = calibrate(x, policy='fixed-window' | 'min-max', K=4)
if discretization == 'interpolation':
exp_approx = grid.interp(x) # 连续
else: # hard rounding
exp_approx = grid.round(x) # 离散 → STE
attn_pre = exp_approx # 尚未归一化
attn = normalize(attn_pre, axis=-1) # row softmax
out = attn @ V
# Backward
if surrogate_placement == 'post-normalization':
grad_through_normalize = exact # 精确
grad_through_discretize = STE # 近似
else: # pre-normalization
grad_through_normalize = STE # 近似
grad_through_discretize = exact # 或同样近似(论文 variants)
grad_x = backward(out, grad, grid) # 含 calibration derivatives
关键实验与数据
控制变量:124M 参数 / 2.5B 训练 token;模型、数据、优化器对齐。
| 实验条件(K=4 域) | 验证 loss 缺口 | 评注 |
|---|---|---|
| Hard rounding + min-max 校准 + pre-normalization surrogate | 大缺口(论文表述:"a large loss gap") | 反例对照 |
| Hard rounding + min-max 校准 + post-normalization surrogate | 缺口大幅缩小 | 关键发现 1 |
| 同样对 Detach 行极值 | Forward 不变 → 验证 loss 延后增大 | "detaching row extrema leaves forward unchanged but produces a delayed increase in validation loss" |
| Fixed-window 校准 + post-normalization surrogate(K=4) | +0.019 nats(相对精确 softmax) | 关键结论 |
| Pre-normalization surrogate(K=16) | +0.004 nats(相对精确 softmax) | 高 K 下进一步压缩缺口 |
⚠️ 诚实标注 / 局限性:
- abstract 仅给出两个最终数字(K=4 +0.019 / K=16 +0.004),中间对照的具体缺口值未公布。
- "124M 参数 / 2.5B token"是单个规模,未做 scale-out--能否在 1B / 7B 参数下保持小缺口,abstract 未述。
- "训练 token" 单位是 B = billion,但 paper-card / abstract 未注明是 BPE token 还是 word token,也未注明 vocab 大小。
- Interpolation vs hard rounding 在更高 K(如 K=8, 16)的对照缺口 abstract 未提供。
- "calibration derivatives"的推导被提及但具体公式未在 abstract 中给出。
1. GitHub / 项目页
- abstract 未明示 GitHub 仓库;论文 v1 提交时间 2026-09-27(fetch-verify-date 2026-09-30,arxiv abs 200 OK)。⚠️ GitHub 缺位已在诚实标注段标记。
亮点与局限
亮点
- 机制研究而非工程复刻:明确把"softmax 低精度化"当作一个研究问题,给出受控对照而非"又一种近似方法"。
- 关键变量拆得清楚:calibration / discretization / surrogate placement 三轴各自独立扫过,结论可解读。
- 校准梯度被严肃对待:把"网格端点也要学"这件事推到了反向传播规则层,技术深度足够。
- "延后掉点"的诊断:detach 行极值 → forward 不变 → validation loss 延后增大,这是非常微妙的现象,论文给出明确实验证据。
局限
- 仅一个规模(124M / 2.5B token):缺 scale-out 实验,工业级 LLM 训练是否同样成立未知。
- 不涉及 KV cache / 推理时量化:论文关注 pretraining,不覆盖部署阶段 softmax 量化路径。
- 不涉及 flash attention 兼容:低精度 softmax 与 flash attention 的 tile 计算如何融合,abstract 未述。
- 不提供端到端 wall-clock 与显存数字:abstract 聚焦 loss 缺口,未报告训练吞吐/显存节省。
对工程落地的启发
- softmax 量化不是免费的:必须明确 surrogate 放置 + 校准策略;默认 pre-norm surrogate 不一定最优(K=4 + hard rounding 时表现差)。
- fixed-window + post-norm surrogate 是稳的入门组合:论文给出的 +0.019 nats 是工业 baseline 的合理起点。
- detach 行极值看似无害但有延后风险:loss 曲线早期不报警,下游才暴露--监控窗口要拉长。
- 校准端点要参与反向:calibration derivatives 是必要的,不能把网格当常数,否则训练不动。
- "高 K + pre-norm" 组合有极限收益:K=16 时 +0.004 nats 已逼近精确 softmax 的训练损失,可作为"成本敏感时的退路"。
⚠ 工程节 6 个具体坑(现象 / 影响 / 修复)
-
坑:surrogate 默认放在 pre-normalization - 现象:用 PyTorch 自带的
STE+ 默认softmax(dim=-1)顺序,把 STE 放在归一化之前。 - 影响:K=4 + hard rounding + min-max 时出现大 loss 缺口,训练静默掉点。 - 修复:显式区分 pre/post normalization surrogate;论文已证 K=4 域 post-norm 明显更稳。 -
坑:grid 端点当常数,反向不传梯度 - 现象:实现 K-interval 时把 grid 当 buffer,不算 calibration derivative。 - 影响:grid 端点不随训练调整,离散化误差不收敛,验证 loss 居高不下。 - 修复:把 grid 端点设为可训练 parameter 或 buffer-with-grad,按论文规则传 calibration derivatives。
-
坑:detach 行极值后 loss 早期无异常 - 现象:为了"省显存"对每行 min/max 做 detach,forward 数值不变。 - 影响:训练前期曲线正常,后期才掉点,难定位。 - 修复:监控 validation loss 至少到与 baseline 收敛点 2 倍步数;detach 决策需对照实验先确认。
-
坑:min-max 校准 + hard rounding 在 K=4 域组合爆炸 - 现象:min-max 让区间紧贴本行极值,硬 rounding 把 out-of-grid 值全归到端点。 - 影响:K 小时代表性下降、loss 缺口放大。 - 修复:K < 8 时改用 fixed-window 或 interpolation;K ≥ 16 时再考虑 min-max + hard rounding。
-
坑:与现有 flash attention 算子不兼容 - 现象:K-interval softmax 需要在每行做 calibration,flash attention 是 tile 级 fused kernel。 - 影响:要么放弃 flash attention(吞吐下降),要么绕过 K-interval(白做)。 - 修复:把 K-interval 实现为 tile-aware 算子:tile 内用 mini-calibration,跨 tile 维护 fixed-window 网格。
-
坑:单规模结论外推到工业训练 - 现象:124M / 2.5B token 上 +0.004~0.019 nats 看起来很小。 - 影响:在 7B / 几百 B token 上同样缺口会被放大,且下游任务指标不一定会同比例劣化。 - 修复:先在 350M-1B 模型上做 5B-10B token 的对照实验,再决定是否上工业训练。
与同方向工作的关系
- QSQ / Softmax-free attention / Linear attention:走"换掉 softmax"路线(如 ReLU/exp 近似);本文走"保留 softmax 但把它量化"路线,是正交方向。
- FlashAttention-2/3、Memory-efficient attention:是"算子实现"层,本文是"算子数学"层,两者可叠加。
- FP8 / INT8 训练(Micikevicius et al., WGMMA 等):是矩阵乘法的低精度化,本文是其 softmax 侧的延伸。
- Distributional / Piecewise-linear softmax approximations:与 K-interval 同属"分段近似"家族,但 K-interval 强调反向规则 + 校准梯度的系统性推导。
适合谁读
- 做大模型低精度预训练的工程团队(FP6/FP4/INT8 路径选择)。
- 研究算子数学与反向传播的学术同学。
- 关心 training loss vs 部署 quantization 是否解耦的从业者——本文明确指出:softmax 的训练时量化不是免费的。
- 在做注意力替代算子评估的人——可作为"保留 softmax 但低精度化"的对照基准。
速读骨架
- 问题:预训练阶段量化 softmax 不是"再压一点",而是 forward/backward 同时被改,需要谨慎设计。
- 方法:K-interval 注意力 + 三轴对照实验(校准策略 / 离散化方式 / STE 放置),推导 calibration derivatives。
- 结论:fixed-window 校准 + post-normalization surrogate 是 K=4 下的稳组合;K=16 + pre-norm surrogate 可逼近精确 softmax(+0.004 nats)。
- 反例:detach 行极值表面无害,实际会延后掉点。
为什么这篇工作重要:三个"反直觉"点
- STE 位置比 STE 是否存在更重要:很多实现默认 STE 与 pre-norm 组合;论文证伪了这种默认值在 K=4 域的可靠性。
- 离散化方式在低 K 域才显著:hard rounding vs interpolation 在 K=16 域差距缩小,提示"粗精度可换实现简洁性"。
- detach 是"隐身"掉点:forward 不变但 loss 后期掉点,这种现象往往被误判为"数据问题"。
给论文作者的延伸建议(基于 abstract 推断)
- 追加 1B 参数 / 30B token 规模的对照实验:验证 +0.004~0.019 nats 在工业规模下仍成立。
- 报告 wall-clock 与显存节省:作为工程指标与训练损失并列报告。
- 公开 K-interval CUDA 实现:包括 tile-aware fused kernel,与 flash attention 的可组合性。
- 与 flash attention 联合实验:推理时量化 softmax + flash attention tile 在长 context 下的实测。
- 下游任务指标披露:训练 loss 仅是代理指标;下游 zero-shot/few-shot 任务分数才是工业决策点。
变量-推荐默认值速查表(基于论文结论)
| 变量 | K=4 默认推荐 | K=16 默认推荐 | 原因 |
|---|---|---|---|
| 校准策略 | fixed-window | fixed-window / min-max | K 小时 min-max 易爆;高 K 可放宽 |
| 离散化方式 | interpolation | hard rounding 可接受 | 低 K 时 hard rounding 误差过大 |
| STE 放置 | post-normalization | pre-normalization | post-norm 更稳;高 K 时 pre-norm 渐近优 |
| 行极值 detach | 谨慎用 | 谨慎用 | 延后掉点,需长窗口监控 |
| Grid 端点是否传梯度 | 是 | 是 | calibration derivatives 必要 |
这张表是"论文推论 + abstract 结论"的压缩,不是论文原文表格,仅供实际训练启动时参考。
补充点(不依赖 abstract 数字的工程含义)
- K-interval 注意力提供了一个"中间路":它不是 softmax-free 路线(不要放弃整个算子),也不是纯 FP32 路线(不要放弃精度收益),而是在"保留 softmax 语义 + 允许低精度算子"之间找到一个可微近似点。
- 论文本质是"用受控变量证明低精度 softmax 不是均匀错误",是低精度训练论文里的"科学化路径",价值在于为后续设计者提供可复现的对照模板。
⚠️ 原文 abstract 未明确项:interpolation vs hard rounding 在 K=8/16 下的具体缺口、calibration derivatives 的具体公式、是否覆盖推理时 softmax 量化、wall-clock 与显存数字、其他模型规模外推、KV cache 路径;以上需查 PDF 正文或附录。
fetch-verify-date: 2026-09-30(arxiv abs 页 200 OK,v1 提交时间 2026-09-27);GitHub 缺位已在诚实标注段标记;Web Archive 备援:未触发 WAF/521/403。
⚠️ 原文 abstract 未明确项:interpolation vs hard rounding 在 K=8/16 下的具体缺口、calibration derivatives 的具体公式、是否覆盖推理时 softmax 量化、wall-clock 与显存数字、其他模型规模外推、KV cache 路径;以上需查 PDF 正文或附录。
fetch-verify-date: 2026-09-30(arxiv abs 页 200 OK,v1 提交时间 2026-09-27);GitHub 缺位已在诚实标注段标记;Web Archive 备援:未触发 WAF/521/403。
工程落地与核查(Jay)
双轨核查:训练侧 vs 推理侧
| 维度 | 训练侧(Training) | 推理侧(Inference) |
|---|---|---|
| 本文覆盖范围 | ✅ 覆盖(核心贡献) | ❌ 论文未涉及 |
| 核心关注 | K-interval softmax 的 loss 缺口;surrogate placement;calibration derivatives | 推理时是否需要保留高精度 softmax;KV cache 量化路径;与 INT4/FP8 推理框架的兼容性 |
| 风险 | 小模型 + 低 K + min-max + pre-norm 组合会出现大 loss 缺口;detach 行极值导致静默掉点 | 若直接迁移到推理系统,softmax 侧的量化收益可能与矩阵乘法侧不匹配 |
| 可验证性 | 可在 124M 上复现;1B+ 规模需查 PDF Section 5 | 需等论文正文披露或自行实验 |
| 与现有 infra 的关系 | 与分布式训练框架(DeepSpeed/Ulysses)可叠加;但需修改 attention kernel | 与推理优化框架(vLLM/TGI)关系未验证,可能是独立工作 |
工程可操作性评估
适合上生产的场景: - 从零训练 1B 以下的模型,且对训练成本极度敏感(每 1% 的显存节省都有价值) - 作为低精度训练 pipeline 的一部分,与 FP8 矩阵乘法一起使用 - 对"延后掉点"现象有监控能力的团队(能追踪到收敛后 baseline 对比)
不适合直接上生产的场景: - 训练 >1B 规模模型(scale-out 结论未知,可能差距放大) - 需要 flash attention 加速的训练场景(K-interval 与 flash attention tile 不兼容,需专项 kernel 开发) - 对训练 loss 收敛稳定性要求极高的场景(小 loss 缺口在大规模训练中可能被放大)
落地核查清单
- ✅ K=4 入门配置:fixed-window + post-norm surrogate + interpolation — 论文给出 +0.019 nats,是工程上手的最低风险配置。
- ✅ grid 端点参与反向传播 — 实现时把 grid 端点设为
requires_grad=True,按论文推导的 calibration derivatives 传梯度,禁止当常数处理。 - ✅ validation loss 监控窗口拉长到 baseline 收敛点的 2 倍 — 防止 detach 行极值导致"延后掉点"被漏掉。
- ✅ K < 8 禁用 min-max + hard rounding 组合 — 低 K 域该组合会爆炸,用 fixed-window 代替。
- ⚠️ flash attention 兼容性需专项验证 — 若训练依赖 flash attention,则 K-interval softmax 需要单独 kernel 开发,不能直接替换。
- ⚠️ 推理侧量化路径论文未覆盖 — 训练时用 K-interval 不代表推理时也能用低精度 softmax;二者是独立问题,需等论文正文或额外实验。
存疑项(原文未明确,需查 PDF)
- ⚠️ calibration derivatives 具体公式:abstract 提及但未给出;需要查论文 Section 3 或附录才能工程实现。
- ⚠️ Interpolation vs hard rounding 在 K=8/16 的对照缺口:abstract 仅给 K=4 和 K=16 两个端点数据,中间 K 值行为未知。
- ⚠️ 模型 scale-out 实验:124M / 2.5B token 是唯一规模;1B/7B 参数下 +0.004~0.019 nats 缺口是否同比例放大未知。
- ⚠️ wall-clock 与显存节省:abstract 未报告训练吞吐数据和显存数据;FLOPs 减少 ≠ 显存减少(KV cache 精度可能不变)。
- ⚠️ KV cache 侧 softmax 量化:论文聚焦预训练,不涉及推理时 KV cache 的量化路径。
- ⚠️ GitHub 仓库:abstract 无链接;K-interval CUDA 实现、calibration derivatives 代码是否开源未知。