FP8 混合精度训练的核心矛盾,是用更窄的数值表示换取更高的矩阵计算吞吐,同时避免溢出、下溢和量化误差破坏收敛。动态缩放正是连接“数值稳定性”与“硬件利用率”的关键机制:它通过持续调整张量进入 FP8 前的尺度,让有效数值尽量落入可表示区间。
为什么 FP8 训练离不开缩放
常见 FP8 格式包括 E4M3 和 E5M2。E4M3 尾数位更多,精度相对较好,但动态范围较小,通常适合权重与前向激活;E5M2 的指数位更多,动态范围更大,常用于梯度。即使按前向和反向分别选用格式,LLM 不同层、不同训练阶段的数值分布仍可能相差很大,因此不能简单地把 BF16 张量直接转换为 FP8。相关格式与混合使用方式可参考 Transformer Engine FP8 说明。
缩放通常先统计张量的绝对值最大值 amax,再根据 FP8 可表示上限计算比例。量化前将原始值乘以比例,反量化时再恢复尺度。比例过小会浪费 FP8 的有效区间,使大量小值被舍入为零;比例过大则可能导致异常值饱和。动态缩放的任务,就是在这两类风险之间寻找平衡。
当前缩放:响应快,但会增加关键路径开销
当前缩放使用本次迭代或当前张量的 amax 立即生成缩放因子。它能快速跟随激活值和梯度分布变化,对训练早期、学习率切换、长序列样本或稀疏专家路由引起的数值突变更敏感。其直接优势是减少因缩放信息过期而产生的溢出,调试逻辑也比较直观。
代价在于,必须先完成 amax 归约,才能确定缩放因子并执行后续量化。这会形成额外的数据依赖和同步点。如果张量较小、微批次较少,或者通信与计算重叠不足,统计与量化开销可能抵消一部分 FP8 GEMM 的吞吐收益。因此,评估当前缩放不能只看 Tensor Core 峰值,还要把归约、类型转换、内存访问和内核启动纳入端到端测量。其基本流程可参见 FP8 Current Scaling 文档。
延迟缩放:更利于流水化,但依赖历史质量
延迟缩放不等待当前张量统计完成,而是根据此前若干次迭代保存的 amax 历史计算本轮尺度。这样可以削弱当前归约对 GEMM 关键路径的阻塞,更容易让缩放统计、量化、矩阵计算和分布式通信形成流水线,从而提升大模型训练中的设备忙碌率。
不过,历史尺度本质上带有滞后性。当数据分布突然改变时,旧尺度可能低估当前峰值,从而产生饱和;如果历史窗口长期保留一次极端异常值,也可能造成尺度过于保守,使普通数值集中在 FP8 区间的一小部分,增加舍入误差。实践中可选择历史最大值、最近值或其他平滑策略,还可设置缩放余量,但余量越大并不代表越稳定,因为它也会牺牲有效精度。
缩放粒度影响稳定性,也影响系统成本
每张量缩放只为整个张量维护一个尺度,元数据少、实现简单,也容易与高性能 GEMM 融合。但 LLM 张量中常同时存在普通通道与离群通道,一个全局 amax 可能被少数极值主导。分块缩放则为局部数据分别配置尺度,使更多数值充分利用 FP8 范围,通常更能适应分布不均匀的权重、激活和梯度。
更细的粒度并非没有成本。尺度数量增加后,会带来额外存储、读取、布局转换和 kernel 约束,实际收益取决于硬件是否提供原生支持。例如 MXFP8 按连续元素块共享尺度,并依赖相应硬件加速;其机制和维度要求可参考 MXFP8 官方文档。因此,不能脱离 GPU 架构讨论哪种粒度“绝对更快”。
如何在稳定性与利用率之间做选择
- 先建立 BF16 基线:比较训练损失、梯度范数、验证指标和吞吐,避免把数据问题或超参数问题误判为 FP8 问题。
- 监控量化健康度:除 loss 外,还应记录 amax、缩放因子、饱和比例、零值比例以及非有限数出现次数。
- 区分不同张量:权重相对平稳,激活和梯度变化更快,不必强制所有张量使用相同的历史窗口、格式或缩放策略。
- 关注阶段切换:预热结束、学习率突变、上下文长度变化和 MoE 负载偏移时,可缩短历史窗口或暂时采用响应更快的尺度更新。
- 测量端到端吞吐:同时观察每步耗时、Tensor Core 活跃度、内存带宽、通信重叠和量化内核占比,而不是只比较单个 GEMM。
总结
动态缩放不会改变 FP8 格式本身,却决定了 FP8 的有限表示范围被如何使用。当前缩放适合强调即时适应性的场景,但可能增加关键路径开销;延迟缩放更容易提高流水化效率,却需要防范历史尺度滞后;分块缩放可改善局部量化精度,但必须结合硬件支持衡量元数据和布局成本。对 LLM 训练而言,合理策略通常不是追求单一算法,而是按张量类型、训练阶段和目标硬件组合缩放粒度与更新时间,并通过稳定性指标和端到端性能共同验证。