FP4 预训练提速新思路:格式感知融合让 8B 模型吞吐提升近一倍
当 FP4 快不起来的时候
四比特浮点(FP4)Tensor Core 本应是加速矩阵乘法的利器,但现实往往没那么美好。量化过程中的缩放计算、操作数打包、布局构建,以及反向传播中需要保存的状态,这些额外开销很容易把 Tensor Core 带来的收益抵消掉。换句话说,算得是快了,但准备工作太慢。
arXiv 上最新的一篇论文《Format-Aware Fusion for Fast FP4 Pretraining》(arXiv:2610.00053)提出了一个名为 format-aware fusion(格式感知融合) 的方案,试图从系统层面解决这个问题。
核心思路:让量化生产者与消费者“对齐”
这篇论文的关键洞察在于:量化不是一个孤立步骤,而是贯穿数据流的过程。作者将每一个**量化生产者(producer)与其缩放域(scale domain)和消费者布局(consumer layout)**协同设计,分别针对三种 FP4 格式做了定制化处理:
- native mxfp
- global nvfp
- cooperative-thread-array-local nvfp
这种“生产者-缩放域-消费者”三者联动的方式,本质上是在消除格式转换和布局重排带来的隐性成本。
实测数据:从 18.8K 到 37.9K tokens/s/GPU
作者用 Llama-3 系列 8B 模型做了预训练实验,训练量达到 1600 亿 token,输出投影使用 bfloat16,交叉熵部分做了编译优化。
在相同加速器的对比测试中:
- bfloat16 基线:18.8K tokens/s/GPU
- Transformer Engine nvfp:27.6K tokens/s/GPU
- 本文最快定制路径:37.9K tokens/s/GPU
另一条路径——mxfp 配合行梯度随机舍入(row-gradient stochastic rounding)和固定符号 32 值 Hadamard 权重梯度预条件——达到 37.2K tokens/s/GPU,对应 86.3% 的 bfloat16 模型 FLOP 利用率。
精度代价:训练损失终点略高,但下游表现另有玄机
速度提升并非没有代价。上述 mxfp 路径最终训练损失比原始 bfloat16 高 2.11%。而另一种 Transformer Engine 方案(最后四个 block 保留 bfloat16)在 27.1K tokens/s/GPU 下,损失仅高出 0.87%。
值得注意的是,论文指出下游任务的排名与训练损失排名并不一致。这意味着,FP4 训练的最终效果并非单纯由损失曲线决定,而是缩放契约、操作数和执行路径三者共同作用的结果。
这意味着什么
对于正在探索低精度预训练的团队来说,这篇论文传递了两个信号:
- FP4 的瓶颈往往不在计算本身,而在围绕它的数据搬运和格式转换。格式感知融合提供了一种系统级优化思路。
- 评估 FP4 方案不能只看训练损失,下游表现可能给出完全不同的答案。
当然,目前这仍是单篇论文的实验结论,实际落地效果还需在更多模型规模和硬件配置上验证。但至少它说明了一点:FP4 预训练要真正跑快,光靠 Tensor Core 不够,还得把整条数据通路一起设计。
