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 低精度化"作为一个研究问题来回答三个互相耦合的设计选择:

  1. 行内网格校准策略:每行用一个固定区间、滑动区间还是 min-max?
  2. 离散化方式:interpolation(连续插值)还是 hard rounding(直接归到最近网格点)?
  3. 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 缺位已在诚实标注段标记。

亮点与局限

亮点

  1. 机制研究而非工程复刻:明确把"softmax 低精度化"当作一个研究问题,给出受控对照而非"又一种近似方法"。
  2. 关键变量拆得清楚:calibration / discretization / surrogate placement 三轴各自独立扫过,结论可解读。
  3. 校准梯度被严肃对待:把"网格端点也要学"这件事推到了反向传播规则层,技术深度足够。
  4. "延后掉点"的诊断:detach 行极值 → forward 不变 → validation loss 延后增大,这是非常微妙的现象,论文给出明确实验证据。

局限

  1. 仅一个规模(124M / 2.5B token):缺 scale-out 实验,工业级 LLM 训练是否同样成立未知。
  2. 不涉及 KV cache / 推理时量化:论文关注 pretraining,不覆盖部署阶段 softmax 量化路径。
  3. 不涉及 flash attention 兼容:低精度 softmax 与 flash attention 的 tile 计算如何融合,abstract 未述。
  4. 不提供端到端 wall-clock 与显存数字:abstract 聚焦 loss 缺口,未报告训练吞吐/显存节省。

对工程落地的启发

  1. softmax 量化不是免费的:必须明确 surrogate 放置 + 校准策略;默认 pre-norm surrogate 不一定最优(K=4 + hard rounding 时表现差)。
  2. fixed-window + post-norm surrogate 是稳的入门组合:论文给出的 +0.019 nats 是工业 baseline 的合理起点。
  3. detach 行极值看似无害但有延后风险:loss 曲线早期不报警,下游才暴露--监控窗口要拉长。
  4. 校准端点要参与反向:calibration derivatives 是必要的,不能把网格当常数,否则训练不动。
  5. "高 K + pre-norm" 组合有极限收益:K=16 时 +0.004 nats 已逼近精确 softmax 的训练损失,可作为"成本敏感时的退路"。

⚠ 工程节 6 个具体坑(现象 / 影响 / 修复)

  1. 坑:surrogate 默认放在 pre-normalization - 现象:用 PyTorch 自带的 STE + 默认 softmax(dim=-1) 顺序,把 STE 放在归一化之前。 - 影响:K=4 + hard rounding + min-max 时出现大 loss 缺口,训练静默掉点。 - 修复:显式区分 pre/post normalization surrogate;论文已证 K=4 域 post-norm 明显更稳。

  2. 坑:grid 端点当常数,反向不传梯度 - 现象:实现 K-interval 时把 grid 当 buffer,不算 calibration derivative。 - 影响:grid 端点不随训练调整,离散化误差不收敛,验证 loss 居高不下。 - 修复:把 grid 端点设为可训练 parameter 或 buffer-with-grad,按论文规则传 calibration derivatives。

  3. 坑:detach 行极值后 loss 早期无异常 - 现象:为了"省显存"对每行 min/max 做 detach,forward 数值不变。 - 影响:训练前期曲线正常,后期才掉点,难定位。 - 修复:监控 validation loss 至少到与 baseline 收敛点 2 倍步数;detach 决策需对照实验先确认。

  4. 坑: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。

  5. 坑:与现有 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 网格。

  6. 坑:单规模结论外推到工业训练 - 现象: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 行极值表面无害,实际会延后掉点。

为什么这篇工作重要:三个"反直觉"点

  1. STE 位置比 STE 是否存在更重要:很多实现默认 STE 与 pre-norm 组合;论文证伪了这种默认值在 K=4 域的可靠性。
  2. 离散化方式在低 K 域才显著:hard rounding vs interpolation 在 K=16 域差距缩小,提示"粗精度可换实现简洁性"。
  3. detach 是"隐身"掉点:forward 不变但 loss 后期掉点,这种现象往往被误判为"数据问题"。

给论文作者的延伸建议(基于 abstract 推断)

  1. 追加 1B 参数 / 30B token 规模的对照实验:验证 +0.004~0.019 nats 在工业规模下仍成立。
  2. 报告 wall-clock 与显存节省:作为工程指标与训练损失并列报告。
  3. 公开 K-interval CUDA 实现:包括 tile-aware fused kernel,与 flash attention 的可组合性。
  4. 与 flash attention 联合实验:推理时量化 softmax + flash attention tile 在长 context 下的实测。
  5. 下游任务指标披露:训练 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 缺口在大规模训练中可能被放大)

落地核查清单

  1. ✅ K=4 入门配置:fixed-window + post-norm surrogate + interpolation — 论文给出 +0.019 nats,是工程上手的最低风险配置。
  2. ✅ grid 端点参与反向传播 — 实现时把 grid 端点设为 requires_grad=True,按论文推导的 calibration derivatives 传梯度,禁止当常数处理。
  3. ✅ validation loss 监控窗口拉长到 baseline 收敛点的 2 倍 — 防止 detach 行极值导致"延后掉点"被漏掉。
  4. ✅ K < 8 禁用 min-max + hard rounding 组合 — 低 K 域该组合会爆炸,用 fixed-window 代替。
  5. ⚠️ flash attention 兼容性需专项验证 — 若训练依赖 flash attention,则 K-interval softmax 需要单独 kernel 开发,不能直接替换。
  6. ⚠️ 推理侧量化路径论文未覆盖 — 训练时用 K-interval 不代表推理时也能用低精度 softmax;二者是独立问题,需等论文正文或额外实验。

存疑项(原文未明确,需查 PDF)

  1. ⚠️ calibration derivatives 具体公式:abstract 提及但未给出;需要查论文 Section 3 或附录才能工程实现。
  2. ⚠️ Interpolation vs hard rounding 在 K=8/16 的对照缺口:abstract 仅给 K=4 和 K=16 两个端点数据,中间 K 值行为未知。
  3. ⚠️ 模型 scale-out 实验:124M / 2.5B token 是唯一规模;1B/7B 参数下 +0.004~0.019 nats 缺口是否同比例放大未知。
  4. ⚠️ wall-clock 与显存节省:abstract 未报告训练吞吐数据和显存数据;FLOPs 减少 ≠ 显存减少(KV cache 精度可能不变)。
  5. ⚠️ KV cache 侧 softmax 量化:论文聚焦预训练,不涉及推理时 KV cache 的量化路径。
  6. ⚠️ GitHub 仓库:abstract 无链接;K-interval CUDA 实现、calibration derivatives 代码是否开源未知。