@PyTorch: 来自 @Meta Engineering 的新发布:FlashAttention-4 扩展了对 @nvidia Blackwell 的 MXFP8 支持——从正向和反向……
摘要
Meta Engineering 扩展了 FlashAttention-4,增加了对 NVIDIA Blackwell 的 MXFP8 支持,实现了高达 2.85 PFLOP/s 的正向性能,并集成了如 GEM 这样的生产训练工作流。
查看缓存全文
缓存时间: 2026/09/16 20:10
来自 @Meta Engineering 的最新动态:FlashAttention-4 新增对 @nvidia Blackwell 架构 MXFP8 精度的支持——涵盖前向与反向内核、融合量化及参差交叉注意力机制。
该团队开发了一个融合量化的端到端参差模块,采用 FP8 激活值与计算精度,目前正被 Meta 内部用于 GEM 模型训练。在最新硬件上,LP FA4 内核在前向传播中达到 2.85 PFLOP/s,在反向传播中达到 2 PFLOP/s,实现最高 1.30 倍的端到端模块加速。
✍️ Devashish Shankar, Santosh Mohan, Jiaqi Xu, Darren Liu, Han Xu
在我们的最新博文中探索设计细节与开源实现方案:https://pytorch.org/blog/low-precision-flash-attention-4-end-to-end-block-scaled-attention-for-blackwell/…
面向 Blackwell 架构的端到端块缩放注意力机制——PyTorch
来源:https://pytorch.org/blog/low-precision-flash-attention-4-end-to-end-block-scaled-attention-for-blackwell/ 博客(https://pytorch.org/blog/category/blog/)
低精度 FlashAttention-4:面向 Blackwell 架构的端到端块缩放注意力机制
Dev (Devashish) Shankar, Darren Liu, Chunzhi Yang, Jackie (Jiaqi) Xu, Markus Hoehnerbach, Jason Xie, Santosh Mohan, Han Xu, Rich Zhu, Josh Fromm, Hongtao Yu, Max Leung (https://pytorch.org/blog/low-precision-flash-attention-4-end-to-end-block-scaled-attention-for-blackwell/#) | 2026年9月16日 | 无评论 (https://pytorch.org/blog/low-precision-flash-attention-4-end-to-end-block-scaled-attention-for-blackwell/#respond)
核心摘要
我们扩展了 FlashAttention-4 [1] 以支持 MXFP8 前向与反向传播,在 LLM 典型架构下实现前向 2.85 PF/s 和反向 2 PF/s 的计算性能。在内部测试架构中,FA4 MX8 前向达 2.54 PF/s,反向达 1.58 PF/s,相比 BF16 精度分别提升最高 1.6 倍和 1.52 倍。我们将量化操作融合至周边生成器,并开发了端到端零索引参差模块,使大多数激活值与计算保持在 FP8 精度。该模块目前正被 Meta 内部用于 GEM 模型训练 (https://engineering.fb.com/2026/08/03/ml-applications/training-gem-at-llm-scale-meta-ads-recommendation-foundation-model/) [2]。据我们所知,这是首批投入生产训练工作的 MXFP8 FA4 前向与反向传播的顶尖实现方案之一。我们已在以下地址开源代码:https://github.com/facebookresearch/ads_model_kernel_library/tree/main/lp_fa4
1. 引言
Blackwell 架构的张量核心引入了块缩放 MMA 指令(tcgen05.mma.block_scale),可原生操作微缩放格式——MXFP8、MXFP6、MXFP4 及 NVFP4——其吞吐量较 BF16 MMA 提升 2-4 倍 [4,5]。然而要在实际训练工作中利用此特性,仅替换注意力内核的数据类型是不够的。缩放因子必须在已饱和的 TMEM 中进行管理;必须为每个操作数(包括在线计算的中间量如 P 和 dS)处理沿 GEMM K 维度的量化;且精度转换的开销需被隐藏,以确保 MMA 能全速持续执行。
在本文中,我们扩展了 FA4 注意力内核,为其添加端到端的 MXFP8 支持(涵盖前向与反向传播),并将其集成到用于广告训练的交叉注意力模块中,实现了融合的生成器与输出尾声处理。主要贡献包括:(1)TMEM 分配策略,将缩放因子适配到利用率已达峰值的 512 列 TMEM 中,且仅需极少新增的同步屏障;(2)利用 Blackwell 的 redux.sync.max.abs.f32 线程束级归约操作,实现了针对 dS 的在线转置不变方块缩放量化;(3)融合的 RMSNorm+量化 与 GEMM+量化 内核,通过单次处理生成具有双缩放因子布局的 FP8 输出,从而消除量化开销;(4)零索引参差模块,其中 FP8 数据保持在非填充位置,仅将规模小得多的缩放因子进行分散、填充和重排,以符合张量核心友好的 128 字节对齐地址,便于 TMA 使用。
2. 实现细节
2.1 注意力前向传播
注意力前向传播包含以下主要操作:
S = Q @ K.T P = Softmax(S) O = P @ V
为实现块缩放 MMA,我们遵循 Quack GEMM 内核与 CUTLASS C++ 示例中的现有 CuTe DSL 范例 [5,6]。我们使用 TMA 加载操作从全局内存(GMEM)获取缩放因子到共享内存(SMEM),并在触发 UMMA 前将缩放因子从 SMEM 复制到 TMEM。这里的主要挑战是 TMEM 争用,下文将详细解释。
目前,在 softmax 线程束中,softmax 计算在 FP32 精度下进行,然后在进行 PV 乘法前将结果转换为 BF16。我们则在将 P 转换为 MXFP8 的同时计算缩放因子。我们将在后文深入探讨为高效实现此目的所进行的 PTX 优化。
需要注意的一个微妙之处是,为使 P.V 块缩放 MMA 正常工作,缩放因子需沿 MMA K 维度计算。对于 Q 和 K,这是注意力的嵌入维度(D);但对于 V,缩放因子和量化需沿序列维度(N)计算。
2.1.1 TMEM 分配与屏障同步
Blackwell 架构的 TMEM 固定大小为 512 列,在现有的 Blackwell FA 内核中已完全用于 MMA 的操作数与累加器。这给添加块缩放 MMA 带来了挑战,因为缩放因子也需要占用 TMEM 空间。
FA4 前向传播使用两个 Q 分块进行乒乓计算。我们加载两个大小为 [128, 128] 的 Q 分块 Q0 和 Q1,并循环处理 K/V 分块(N 维度)。GEMM 的执行顺序如下:
GEMM 序言 S0 = Q0 @ K0 S1 = Q1 @ K0 主循环 (for n in 0 .. N-1) O0 = P0_n * V_n S0 = Q0 * K_{n+1} O1 = P1_n * V_n S1 = Q1 * K_{n+1} 尾声 O0 = P0_N * V_N O1 = P1_N * V_N
TMEM 的分配如下:如图所示,TMEM 已被完全利用,没有剩余空间容纳缩放因子(SF)。注意,我们不能将输入缩放因子与累加器 TMEM 重叠。为解决此问题,我们采用以下重叠方式:
- 对于序言 S(i) GEMM,我们可以使用 O(i) 区域存放 S(i) 的缩放因子,因为此时 O(i) 尚未开始计算。
- O(i) 的缩放因子可以与 S(i) 重叠存放——这与常规 FA 中将 P(i) 与 S(i) 重叠的策略相同。因此,在我们复制 O(i) 的缩放因子之前,已经存在一个屏障确保 S(i) 的 TMEM 内容已被消费。我们只需选择 S(i) 中与 P(i) 不同的区域。注意,对于 FP8,P(i) 占用 32 列,而 S(i)(FP32)占用 128 列。
- S(i) 的缩放因子与 S(1-i) 重叠存放。这需要在 MMA 线程束和 Softmax 线程束之间增加一个额外的同步屏障,因为 MMA 是异步执行的,S(1-i) 的累加器与 S(i) 的缩放因子之间可能存在写写冲突。因此,我们在 MMA 线程束中添加一个屏障,在复制 S(i) 的缩放因子前等待,当 softmax 线程束读取 S(1-i) 的累加器时该屏障被释放。注意,这个屏障通常不会增加额外开销,因为存在一个 O(i) GEMM 来重叠屏障前必须的 TMEM->寄存器读取操作。通常,GEMM 的耗时长于 TMEM->寄存器的读取时间。
具体来说,带有缩放因子放置位置的 GEMM 执行顺序如下所示:
GEMM 缩放因子TMEM区域 序言 S0 = Q0 @ K0 O0 (空闲,未开始) S1 = Q1 @ K0 O1 (空闲,未开始) 主循环 (for n in 0 .. N-1) O0 = P0_n * V_n S0 (S已消耗,P区域不同) S0 = Q0 * K_{n+1} S1 (新屏障!) O1 = P1_n * V_n O0 (现有屏障) S1 = Q1 * K_{n+1} S0 (新屏障!) 尾声 O0 = P0_N * V_N O1 = P1_N * V_N
2.1.2 改进的 unroll-KV
随着精度降低,MMA 吞吐量提升(MXFP8 相比 BF16 提升 2 倍,MXFP4 提升 4 倍),而 softmax 线程束中受 SFU 限制的计算量保持不变。这改变了瓶颈所在:原本隐藏在慢速 BF16 MMA 后面的 softmax 停顿现在暴露出来。
在持久化内核中,分块边界问题尤为突出。若不使用 unroll-KV,一个分块的最后两个 GEMM 都是 PV,接着是下一个分块的两个 QK。下一个分块的 Q0 的 softmax 必须等到 QK0 完成才能启动——但 QK0 必须等待当前分块(tile(n))的所有 PV GEMM 完成。这导致了每个分块边界处都出现停顿。
unroll-KV(灵感来自 GDPA (https://pytorch.org/blog/generalized-dot-product-attention-tackling-real-world-challenges-in-gpu-training-kernels/) [3])将当前分块的最后一个 PV 与下一个分块的第一个 QK 交织执行:
现在,Q0 的 softmax 提前一个 GEMM 触发,从而将分块边界延迟隐藏在 PV 流水线之后。
然而,在 BF16 下启用此特性会导致性能下降:校正线程束(按顺序处理两个阶段)因为被延迟的 PV1 而被延迟,这种延迟通过 softmax_corr_empty 屏障级联影响到下一个 softmax 操作。我们通过将屏障等待移至行求和计算之后(行求和不依赖该屏障)来解决此问题,从而使校正流水线与关键路径解耦。
2.1.3 优化的在线 MXFP8 转换
由于我们希望使用块缩放注意力进行 QK 和 PV 的 GEMM 计算,这需要将 softmax 后的输出(P)从 FP32 在线转换为 MXFP8,而非 BF16。应用块缩放是一个 3 步过程。对于一个包含 32 个元素的块 x:
这里,a 是该块的 amax(绝对最大值),sigma 是缩放因子。为在 Blackwell 上高效实现此过程,我们利用了 3 指令最大值计算和 fmul2 指令(在 CuteDSL nvvm 中暴露)用于步骤(1)和(3)——这些操作对每个元素执行。对于计算缩放因子(步骤 2)——我们使用优化的 PTX 序列,通过提取 FP32 的指数位和尾数位来避免 log2 和除法运算。
此外,我们注意到 softmax 的计算已经涉及在指数化之前计算 128 个元素的行最大值。由于 exp 是单调递增函数,我们可以复用 softmax 计算的最大值,从而在增加每 32 个元素 1 个 exp 操作的同时,避免了额外的最大值计算操作。我们发现这种方法能带来更好的性能。
我们近期的实验表明,通过为 P 使用常量缩放可以进一步提升性能;鉴于 softmax 算子自然地将输出限制在 [0, 1] 区间内,这在数值上是稳健的。
2.1.4 使用 TMA 处理变长序列张量
处理参差(jagged)数据在广告模型中很重要,但在 MXFP8 下尤其具有挑战性。Blackwell 块缩放 MMA 使用一个经过 swizzle 处理的 512 字节缩放因子原子 (https://docs.nvidia.com/cutlass/latest/media/docs/cpp/blackwell_functionality.html#scale-factor-layouts),对应于 128×128 数据分块的缩放因子 [5]。由于参差序列长度是任意的且通常不是 128 的倍数,它们的自然偏移不满足此布局要求。填充完整数据张量是可行的,但会增加昂贵的内存访问和存储开销。相反,我们采用分离寻址:仅将规模小得多的缩放因子张量填充到 128 对齐的位置,而 FP8 数据保持紧凑存储。在消费端 GEMM 或注意力内核中,TMA 加载使用填充后的偏移获取缩放因子,使用原始的参差偏移获取数据。我们使用轻量级的分散和布局转换内核,将缩放因子从紧凑的参差布局移动到下游 TMA 加载所期望的填充、swizzle 后的布局。更一般地说,这些内核在不改变数值的情况下重排缩放因子的字节,使得每个后续的 GEMM 或注意力内核可以直接在其首选布局中消费缩放因子。
2.2 注意力反向传播
注意力反向传播包含以下关键操作:
GEMMs: S = K @ Q.T dP = V @ dO.T dV = P.T @ dO dK = dS.T @ Q dQ = dS @ K
计算: P = softmax(S) dS = dsoftmax(dP, P)
注意在反向传播中,多个 GEMM 是转置的。例如: dP = V @ dO.T dV = P.T @ dO
对于 dP MMA,dO 的量化需要沿 D(嵌入维度)进行,然而对于 dV MMA,量化需要沿 M(序列维度)进行。这与我们之前在前向传播 P.V MMA 中看到的情况类似,其中 V 需要沿 M 维度量化。但在反向传播中,我们有 GEMM 需要同时沿两个轴进行量化。注意,这是使用块缩放 MMA 的一个局限,因为量化必须沿 GEMM K 维度进行。
为解决此问题,我们使用 [32,32] 方块对 Q、K 和 dO 进行量化,使 E4M3 载荷具有转置不变性。因此,两个 GEMM 可以复用同一量化表示,无需存储不同版本。规模更小的 E8M0 缩放因子仍然为每个 GEMM 所需的张量核心缩放布局分别布局和加载。dS 也面临类似的转置挑战。方块量化方案产生的单一表示可供 dK 和 dQ 共同使用。
2.2.1 TMEM 分配
[注:此方案描述的是 1-CTA 路径。我们目前对 MX8 使用 1-CTA 路径,因其当前表现更优]
前向传播有 2 个 GEMM(S 和 O),使用 4 个累加器进行 2 阶段 Q 流水线;反向传播则有 5 个 GEMM,它们都需要 TMEM 空间来存放累加器和缩放因子。TMEM 布局将全部 512 列打包:
核心约束:dK 和 dV 是持久化累加器——它们在整个 M 循环中持续累加,因此其 TMEM 区域在整个内核生命周期内都被占用。这意味着缩放因子放置只能使用 S 和 dP 区域。我们对前 2 个 GEMM(S 和 dK)使用 dP 区域,对后 3 个 GEMM 使用 S 区域。我们只需在 S GEMM 之前添加一个额外的屏障,这应该不会增加额外开销。TMEM 分配方案详述如下:
名称 GEMM 缩放因子 缩放因子TMEM区域 安全原因? 序言 SFK, SFQ, SFV, SFDO dK dK 在序言阶段空闲 S ([email protected]) SFK, SFQ dP pipeline_dP_drain.empty.wait 确保之前的 dP 值已被排空。在单集群中,这与现有 pipeline_dP 屏障复用。 dK (dS.T@Q) SFDS, SFQ_dK dP 由于 dS 是 dK GEMM 的输入,现有屏障确保该区域空闲可用。注意 dS 仅占 32 列(FP8),因此我们有 96 列空闲。 dQ (dS@K) SFDS_dQ, SFK_dQ S pipeline_S_drain.empty.wait 确保 S 在其 TMEM 区域被复用前已被读取。 dP ([email protected]) SFV, SFDO S 隐含顺序。 dV (P.T@dO) SFP, SFDO_dV S pipeline_S_P.empty.wait(P 就绪所需)
2.2.2 在线方块 dS 量化
在反向传播中,dS 以 FP32 精度计算,并被两个 GEMM 消费: dK = dS.T @ Q, dQ = dS @ K。
我们对 dS 的每个 32×32 块进行一次性量化。每个线程束车道拥有 32 个值,并计算一个线程局部的绝对最大值。Blackwell 的 redux.sync.max.abs.f32 指令随后在线程束内归约这 32 个部分最大值,以获得整个块的 AMAX。由于缩放区域是方块,当 dS 被转置时该区域仍然有效。因此,相同的缩放因子和载荷可以同时供给两个 GEMM。
一次性量化 dS 避免了额外的 E8M0 转换、反缩放乘法、E4M3 转换和载荷分块。dK 和 dQ MMA 仍然需要将共享的缩放因子复制到各自的硬件布局中。我们发现 [32, 32] 方块量化在我们的训练工作中在数值上是可接受的。
2.2.3 FP16 dQ 归约
通过消融实验,我们发现 dQ 归约是一个关键的吞吐量瓶颈。因为 dQ 需要在每次内部循环迭代中将一个 128×128 的分块写入全局内存(GMEM),这大幅增加了全局内存带宽消耗。为缓解此问题,我们评估了 BF16 和 FP16 dQ 归约策略,发现采用静态缩放的 FP16 dQ 能实现最优性能。鉴于在生产工作负载中 dQ 值始终较小,这种 FP16 方法配合较大的缩放因子能保持可接受的数值精度。我们将在结果部分提供 FP32 与 FP16 dQ 归约的详细性能对比。
相似文章
@PyTorch: PyTorch 成员 Meta 刚刚开源了一个 GPU 内核,使注意力在 NVIDIA Blackwell 上加速 2.3 倍。TLX Block Atte…
Meta 开源了 TLX Block Attention,这是一个 warp 特化的 Triton 内核,在 NVIDIA Blackwell GPU 上为块对角自注意力实现了 2.3 倍的加速,与旋转嵌入融合时加速可达 3.5 倍。
@Hass_Abdallah11: 我们有两篇新的 Colfax 文章,关于优化 FlashAttention-4,其中包括我(解码)和 Jack Carlisle(反向传播)的工作…
文章详细介绍了 FlashAttention-4 的优化技术,包括用于解码的 S/P 乒乓方法以重叠操作,在 NVIDIA Blackwell GPU 上实现了高达 16% 的性能提升。
FlashMLA sm_120 内核构建,性能比 SDPA 提升 2-3 倍
作者为消费级 Blackwell sm_120 构建了 FlashMLA,在长上下文训练和稀疏预填充等注意力密集型工作负载中,性能比 PyTorch SDPA 提升了 2-3 倍。
Flash-MSA: 利用稀疏注意力内核加速百万token训练
介绍Flash-MSA,首个针对MiniMax稀疏注意力在Hopper和Blackwell GPU上的高性能开源训练内核,实现高效的百万token训练。
@PyTorch: AMD has been upstreaming optimizations for improved FP8 training support in PyTorch/TorchTitan and PyTorch/TorchAO, mak…
AMD upstreamed optimizations to PyTorch/TorchTitan and TorchAO for FP8 training on AMD Instinct GPUs, achieving up to 13.4% throughput gains on Llama3-8B and recovering 89% of FP8 quantization overhead on DeepSeek-V3 via fused Triton kernels.