PyTorch 发布 Flash Attention 4 MXFP8 版本:Blackwell 架构训练性能再创新高

ADK PyTorch官方 / ADK编译 2026-08-03 5 分钟 107 次浏览
速览导读 / Summary

Meta 团队在 PyTorch 中开源了针对 NVIDIA Blackwell 架构优化的 Flash Attention 4 (FA4) 低精度版本,全面支持 MXFP8 前向与反向传播。该实现通过端到端的零散列模块(Zero-gather jagged module)和优化的 TMEM 分配策略,在 LLM 训练形状上实现了高达 2.85 PF/s 的前向吞吐,相比 BF16 模式提升 1.6 倍。这是首个在生产级 GEM 训练中大规模应用 MXFP8 FA4 的前沿方案,显著降低了大模型训练的计算成本与显存压力。

前向吞吐 (MXFP8) 2.85 PF/s LLM 训练形状下,相比 BF16 提升 1.6 倍
反向吞吐 (MXFP8) 2.00 PF/s LLM 训练形状下,相比 BF16 提升 1.52 倍
内部形状前向 (MX8) 2.54 PF/s 相比 BF16 提升 1.6 倍

Key Insights / 核心看点

  • 1 首次实现 Flash Attention 4 的前向与反向传播均原生支持 MXFP8,无需中间精度转换。
  • 2 在 LLM 训练形状上,MXFP8 前向吞吐达到 2.85 PF/s,相比 BF16 提升 1.6 倍。
  • 3 引入端到端零散列模块,确保 FP8 数据在 unpadded 位置保持,仅对微小的尺度因子进行复杂的内存管理。

PyTorch 发布 Flash Attention 4 MXFP8 版本:Blackwell 架构训练性能再创新高

更新背景

随着 NVIDIA Blackwell 架构的推出,其张量核心(Tensor Cores)原生支持微缩放格式(MXFP8, MXFP6, MXFP4),理论上可提供比 BF16 高 2-4 倍的计算吞吐量。然而,要在实际的大语言模型(LLM)训练中充分利用这一潜力,仅更换数据精度是不够的。Blackwell 架构对 TMEM(Tensor Memory)空间有严格限制,且量化过程中的尺度因子(Scale Factors)管理极为复杂。Meta 团队此次发布的 PyTorch 更新,旨在解决这些底层架构瓶颈,将 MXFP8 深度集成到 Flash Attention 4 内核中。

核心突破与功能特性

本次更新的核心在于实现了端到端(End-to-End)的 MXFP8 支持,涵盖前向传播与反向传播全过程,并针对 Blackwell 的硬件特性进行了深度优化:

1. 端到端零散列模块 (Zero-gather Jagged Module)

为了最大化利用 MXFP8 的带宽优势,团队开发了一种创新的模块设计。在该设计中,绝大多数激活值和计算结果保持为 FP8 格式,仅在需要处理稀疏或对齐要求时,将微小的尺度因子(Scale Factors)进行散列、填充和重排,使其适配 Blackwell 友好的 128 字节对齐地址。这种设计确保了数据在传输和计算过程中绝大部分时间都停留在 FP8 精度,避免了频繁的精度转换开销。

2. 极致优化的 TMEM 分配策略

Blackwell 的 TMEM 大小为 512 列,在现有 FA 内核中已被完全占满。新增 MXFP8 支持意味着必须为尺度因子预留空间,这极易引发资源冲突。团队通过精细的 Ping-Pong 计算策略和屏障同步(Barrier Synchronization),成功将尺度因子重叠存储在已有的 GEMM 结果区域中,实现了 TMEM 的零浪费利用,同时最小化了额外的同步开销。

3. 融合量化与归一化内核

为消除量化带来的额外延迟,团队开发了融合 RMSNorm 与量化、以及 GEMM 与量化的专用内核。这些内核能够在单次执行中同时输出 FP8 结果并生成双布局的尺度因子,彻底消除了传统流程中多次精度转换的瓶颈。

实际应用价值

该更新不仅是一个理论上的算法改进,更是 Meta 内部用于 GEM (Generative Engine for Meta) 训练的实际生产方案。对于开发者而言,这意味着:

  • 训练成本大幅降低:在同等硬件配置下,利用 MXFP8 可显著减少所需的算力资源,缩短训练周期。
  • 显存效率提升:低精度计算减少了中间激活值的显存占用,使得在有限显存上训练更大规模的模型成为可能。
  • 生产级落地:这是目前业界首个在大规模生产训练工作流中成功应用 MXFP8 Flash Attention 的前沿实现,为其他厂商提供了宝贵的参考范式。

关键亮点 (Key Highlights)

  • 全链路 MXFP8 支持:首次实现 Flash Attention 4 的前向与反向传播均原生支持 MXFP8,无需在中间步骤进行精度转换。
  • 生产级性能表现:在 LLM 训练形状上,MXFP8 前向吞吐达到 2.85 PF/s,反向吞吐达到 2 PF/s,相比 BF16 分别提升 1.6 倍 和 1.52 倍。
  • 创新零散列架构:引入端到端零散列模块,确保 FP8 数据在 unpadded 位置保持,仅对微小的尺度因子进行复杂的内存管理,最大化硬件利用率。

关键技术指标 (Metrics)

指标 数值 说明
前向吞吐 (MXFP8) 2.85 PF/s LLM 训练形状下,相比 BF16 提升 1.6 倍
反向吞吐 (MXFP8) 2.00 PF/s LLM 训练形状下,相比 BF16 提升 1.52 倍
内部形状前向吞吐 (MX8) 2.54 PF/s 内部测试形状,相比 BF16 提升 1.6 倍
内部形状反向吞吐 (MX8) 1.58 PF/s 内部测试形状,相比 BF16 提升 1.52 倍

官方引言: "To our knowledge, this is one of the first SoTA implementations of MXFP8 FA4 forward and backward being used in production training workloads. We have open sourced the code..." —— Meta 团队,PyTorch 官方博客

代码获取

开源代码已发布在 Meta 的 ADS 模型内核库中: https://github.com/facebookresearch/ads_model_kernel_library/tree/main/lp_fa4

“To our knowledge, this is one of the first SoTA implementations of MXFP8 FA4 forward and backward being used in production training workloads.”

— Meta 团队

同主题深度资讯

查看更多 →
官方动态 2026-10-1

PyTorch 官方最新动态:Modernizing Table Batched Embeddings with FBTriton

This post explores the FBTriton kernel design for Table Batched Embedding (TBE) forward and backward passes. These core operators handle embedding lookups across thousands of sharded GPUs within recom

PyTorch官方 / ADK编译 3分钟
官方动态 2026-10-08T18:13:46+00:00

PyTorch 官方最新动态:Session-Aware Agentic Inference with NVIDIA Dynamo

Agentic workloads change the traffic an inference server sees. Unlike single-turn chat, an agent session can involve a large initial prefill, repeated model calls, and subagents running in parallel, w

PyTorch官方 / ADK编译 3分钟
模型发布 2026-10-08

PyTorch 原生集成 Spyre 加速器:打造企业级推理新范式

PyTorch 官方宣布通过 `torch-spyre` 库将 IBM 的 Spyre 数据流 AI 加速器打造为 PyTorch 的原生设备。此次更新实现了从设备身份、内存分配器到编译图的全链路映射,支持 `tensor.to("spyre")` 等标准操作。核心突破在于利用 Spyre 的独立计算与数据移动流水线,在保持 PyTorch 易编程性的同时,实现了极低延迟的推理加速,为 IBM Z 及 Power 系统上的企业级 AI 应用提供了统一且高效的执行路径。

PyTorch官方 / ADK编译 5 分钟
模型发布 2026-10-08

Pixmax 发布全链路游戏素材生成平台,覆盖研发至买量全阶段

Pixmax 正式发布专为游戏行业打造的 AI 素材生成平台,深度理解游戏开发工作流,支持从研发立绘、场景贴图到宣发买量素材的全链路生产。平台通过“上传 - 生成 - 批量 - 适配”四步流程,解决风格统一与批量产出难题,显著降低美术成本与试错周期,助力中小团队提升生产效率。

Pixmax官方 / ADK编译 3 分钟
code · 免费+付费
★ 5.0 · 120评测
P

PyTorch

开源的机器学习库

PyTorch 是开源的机器学习库,主要用在深度学习研究和应用开发,以灵活性、易用性和强大的 GPU 加速功能而闻名。PyTorch 提供动态计算图,支持开发者在运行时动态修改模型结构,非常适合快速开发和实验。PyTorch 支持张量计算、自动微分(torch.autograd)和模块化的神经网络构建(torch.nn)。PyTorch 拥有丰富的社区支持和大量的预训练模型及教程,是学术界和工业界的首选深度学习框架之一。

查看 PyTorch 使用教程与功能