SheepNav
新上线今天0 投票

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 训练的最终效果并非单纯由损失曲线决定,而是缩放契约、操作数和执行路径三者共同作用的结果。

这意味着什么

对于正在探索低精度预训练的团队来说,这篇论文传递了两个信号:

  1. FP4 的瓶颈往往不在计算本身,而在围绕它的数据搬运和格式转换。格式感知融合提供了一种系统级优化思路。
  2. 评估 FP4 方案不能只看训练损失,下游表现可能给出完全不同的答案。

当然,目前这仍是单篇论文的实验结论,实际落地效果还需在更多模型规模和硬件配置上验证。但至少它说明了一点:FP4 预训练要真正跑快,光靠 Tensor Core 不够,还得把整条数据通路一起设计。

延伸阅读

  1. 亚马逊回应数据中心争议:不再使用保密协议,驳斥四大迷思
  2. 英伟达 Shield TV 上市七年逆势涨价 100 美元,AI 狂潮推高内存成本
  3. 卡普空布局AI:从辅助开发到“与AI共同创造游戏”的未来
查看原文