用 Blackwell 跑 Flash Attention 4 居然能把 MXFP8 的前向吞吐顶到 2.85 PF/s
这次 Meta 搞的这个 Low Precision Flash Attention 4 (LP-FA4) 确实有点意思,直接把 MXFP8 的前向和后向全部打通了。最核心的结论是,在 LLM 规模的 shape 下,前向能跑到 2.85 PF/s,后向 2 PF/s。对比 BF16,前向提升了 1.6 倍,后向提升了 1.52 倍。而且这玩意儿已经实装在 Meta 内部的 GEM 训练里了,不是那种跑个 benchmark 就扔的 Demo。
其实在 Blackwell 上用 MXFP8 没那么简单,不能直接把数据类型给换了就完事。最头疼的是 Scale Factor(缩放因子)的管理。Blackwell 的 Tensor Core 虽然支持块缩放 MMA 指令(tcgen05.mma.block_scale),但这些 SF 必须塞进已经快被挤爆的 TMEM 里。如果处理不好,转换精度的开销会直接把 MMA 的速度拖下来。
Meta 这次在实现上解决了几个关键痛点:
首先是 TMEM 的内存分配。Blackwell 的 TMEM 只有 512 列,之前的 FA kernel 基本上把空间占满了。为了塞进 SF,他们用了两个 Q tile ([128, 128]) 之间做 ping-pong 计算,通过精准控制加载顺序,在不增加过多 barrier 的情况下,把 SF 从 GMEM 搬到 SMEM 再到 TMEM,保证了 UMMA 能满速跑。
其次是量化操作的融合。他们写了融合的 RMSNorm+Quantize 和 GEMM+Quantize kernel,让输出直接就是带双重 SF 布局的 FP8,省去了单独做量化的开销。另外,针对 dS 的量化,利用了 Blackwell 的 redux.sync.max.abs.f32 指令做 warp-wide 归约,实现了在线转置不变的平方块缩放量化。
最硬核的是那个 zero-gather jagged module。在处理非对齐数据时,FP8 的激活值依然留在原位(unpadded),只有体积小得多的 SF 会被 scatter、padding 并排列到符合 TMA 要求的 128 对齐地址。这样就避免了大规模搬运 FP8 数据的浪费。
具体到前向计算流程,虽然 S = Q @ K.T 和 O = P @ V 看起来没变,但在 MXFP8 下,V 的 scale 和量化必须沿着序列维度 (N) 来算,而 Q 和 K 则是沿着嵌入维度 (D) 算,这个细节如果不处理好,结果直接就错了。而且在 softmax warp 里,计算依然用 FP32,但之后会转换成 MXFP8 并同步计算 scale,确保 P @ V 能够触发块缩放 MMA。
如果你想尝试,代码已经开源了,路径在这里:
https://github.com/facebookresearch/ads_model_kernel_library/tree/main/lp_fa4
总的来说,这次更新把 MXFP8 的端到端链路跑通了,尤其是解决了 TMEM 争抢和量化开销的问题。对于追求极致吞吐的训练任务来说,这种 1.5 倍以上的提升非常可观,但前提是你得有 Blackwell 的卡,而且得能搞定这套复杂的 SF 管理逻辑。

2.85 PF/s 看着是猛,但要是遇到那个离谱的精度掉点就全白搞了,没说用哪个量化工具对齐的?