用 Blackwell 跑 Flash Attention 4 居然能把 MXFP8 的前向吞吐顶到 2.85 PF/s

产品经理大鹏 初级 48分钟前 565 浏览 13 点赞 约 2 分钟

这次 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 跑 Flash Attention 4 居然能把 MXFP8 的前向吞吐顶到 2.85 PF/s

其实在 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 能满速跑。

用 Blackwell 跑 Flash Attention 4 居然能把 MXFP8 的前向吞吐顶到 2.85 PF/s

其次是量化操作的融合。他们写了融合的 RMSNorm+QuantizeGEMM+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。

用 Blackwell 跑 Flash Attention 4 居然能把 MXFP8 的前向吞吐顶到 2.85 PF/s

如果你想尝试,代码已经开源了,路径在这里:

https://github.com/facebookresearch/ads_model_kernel_library/tree/main/lp_fa4

总的来说,这次更新把 MXFP8 的端到端链路跑通了,尤其是解决了 TMEM 争抢和量化开销的问题。对于追求极致吞吐的训练任务来说,这种 1.5 倍以上的提升非常可观,但前提是你得有 Blackwell 的卡,而且得能搞定这套复杂的 SF 管理逻辑。

AI编程MetaBlackwellFlash Attention 4MXFP8

全部回复 (3)

夜猫子创业者 专家 45分钟前

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

0 回复
小柯爱学习 专家 39分钟前

这种性能提升简直救命,我上个月跑那个 70B 简直慢到想砸机器,快冲去试下这个 FP8 到底稳不稳!

0 回复
数据分析师大山 中级 39分钟前

这吞吐量顶得我心慌,要是精度崩了得多少钱打水漂?没说在哪个规模的 model 上测的?

0 回复

发表回复

支持 Markdown 格式