FP8 Training on AMD GPUs with TorchTitan and TorchAO: Upstreaming Performance Improvements

TL;DR · AI 摘要
PyTorch Blog文章展示了在AMD GPU上使用FP8训练的性能提升,包括13.4%吞吐量增益和6.2倍加速,通过Triton融合等技术实现。
核心要点
- FP8训练在Llama3-8B模型上实现13.4%吞吐量提升,内存占用与BF16接近。
- Triton融合技术恢复DeepSeek-V3 671B MoE架构89%的FP8量化开销。
- AMD FNUZ格式支持使TorchTitan原生兼容MI300X等GPU。
结构提纲
按章节快速跳转。
思维导图
用一张图看清主题之间的关系。
查看大纲文本(无障碍 / 无 JS 友好)
- AMD GPU FP8训练优化
- 性能提升
- Llama3-8B +13.4%吞吐量
- DeepSeek-V3 6.2×加速
- 关键技术
- AMD FNUZ格式支持
- Triton融合管道
- ROCm分组GEMM
金句 / Highlights
值得收藏与分享的关键句。
FP8训练在Llama3-8B上实现13.4%吞吐量增益,内存占用与BF16接近。
Triton融合技术使DeepSeek-V3 671B MoE层达到6.2×加速。
FNUZ格式(无NaN/Inf)在MI300X GPU上实现高效FP8计算。
在 AMD GPU 上使用 TorchTitan 和 TorchAO 进行 FP8 训练:将性能改进上游化 – PyTorch
特色项目
在 2025 年 PyTorch 大会上,我们通过 Primus-Turbo(AMD 优化库)展示了在 AMD Instinct 集群上使用 TorchTitan 进行训练时,超过 1000 块 GPU 的线性扩展能力。此后我们已将这些 AMD 优化方案上游化,使 TorchTitan 可直接支持 AMD Instinct(™) GPU,并具备开箱即用的 FP8 性能优势。所有提到的贡献均已合并至上游 pytorch/AO 和 pytorch/TorchTitan 项目。
如图 1 所示,在密集型模型上,FP8 训练在 Llama3-8B(#2736)上相比 BF16 实现了 13.4% 的吞吐量提升。在 DeepSeek-V3 671B 等 MoE 架构中,FP8 量化初期带来了显著的性能开销。通过融合 Triton 量化内核,我们在 DeepSeek-V3 671B MoE 架构(#4311)上成功回收了 89% 的 FP8 量化开销,单个内核优化甚至实现了 6.2 倍的加速(#4113)。
本文将介绍实现这些性能提升的 FP8 优化方案,涵盖从核心加速到 Triton 融合流水线的完整技术栈。实现过程包含以下三个关键工作:
- 添加对 AMD FP8 数值格式的原生支持
- 在 ROCm 上启用 MoE 模型的分组 GEMM
- 构建 Triton 融合流水线以降低量化开销
| 工作负载 | 优化措施 | 结果 | PR | |------------------|-------------------------------|---------------------------|------| | Llama3-8B (密集型) | 行维度 FP8 vs BF16 | +13.4% 吞吐量 | #2736 | | DeepSeek-MoE-16B | 反向转置移除 + 融合 | 4.2 倍反向传递速度 | #3972 | | | | | #4069 | | DeepSeek-V3 671B | 列维度缩放合并 | 每 MoE 层 6.2 倍加速(7,290→1,170µs) | #4113 | | 前向传递融合 | +17% 端到端性能;回收 89% FP8 性能差距 | | #4311 |
图 1:在 8×MI300X GPU 上使用 Llama3-8B 进行行维度 FP8 训练的吞吐量(批大小 1,序列长度 8192,100 步,torch.compile,FSDP2,每操作选择性激活检查点)。采用高精度权重-梯度方案的行维度 FP8(权重更新 GEMM 保持在 BF16,前向和梯度输入 GEMM 使用 FP8)相比 BF16 实现了 13.4% 的吞吐量提升,峰值内存几乎相同(~39 GB)。性能提升源于更快的 FP8 矩阵核心而非内存节省。所有数据均来自 TorchAO PR #2736。
TorchAO 中的 AMD FP8 格式
每个线性层执行三次矩阵乘法:前向传递、梯度输入和梯度权重更新。FP8 训练将这些操作从 16 位量化到 8 位,显著提升吞吐量。AMD Instinct GPU 实现了一种称为 FNUZ(有限值、无 NaN、无符号零)的 FP8 格式变体,如下表所示:
| 属性 | e4m3fnuz (AMD) | |--------------|-------------| | 最大值 | 240 | | NaN/Inf 编码 | 无 | | 硬件支持 | MI300X, MI325X, MI350X |
Primus-Turbo 展示的 FP8 能力已直接上游化至 TorchAO 和 TorchTitan,如图 2 所示,涵盖三个领域:硬件感知的 FP8 格式支持、ROCm 上 MoE 分组 GEMM 启用、降低量化开销的 Triton 内核融合流水线。
图 2:ROCm 上 TorchTitan 的 FP8 训练软件栈。AMD 的上游贡献涵盖 TorchAO(FP8 数据类型支持、Triton 内核优化)和 TorchTitan(MFU 修复、损失基线、扩展配方)
TorchAO 库最初并未支持 AMD 的相同数值格式,因此它使用了不同的最大值进行缩放计算。在 AMD Instinct GPU 上,由于 e4m3fnuz 的最大值为 240,这会导致静默错误:张量被缩放到超出硬件可表示值的范围,导致激活值被截断并破坏梯度。由于 e4m3fnuz 没有 NaN/Inf 编码,溢出不会引发错误,而是导致模型质量下降。因此选择正确的格式是正确性要求,而非调优选项。我们添加了硬件自动检测功能,使 TorchAO 能够自动选择正确的格式。实现这一目标需要在 TorchAO 和 TorchTitan 中进行一系列格式正确性修复:
- 自动检测平台并选择正确的 FP8 数据类型和最大值,而非硬编码 NVIDIA e4m3fn:TorchAO #1142 , #1150 , #2225
- 报告 MI300X 的正确峰值 FLOPS 以确保 MFU 数值准确:TorchTitan #920
- 为 FNUZ 数值添加平台特定的损失基线:TorchTitan #2156
FP8 缩放可以应用于不同粒度:每个张量一个缩放值(张量级,最快但最粗粒度)、每行一个缩放值(行级,精度更高)、固定大小的块(块级),或与数据打包在一起的组(MXFP8)。TorchAO 和 TorchTitan 支持所有四种策略。对于 AMD GPU,我们确保每种量化策略都能与 AMD 特定数值正确配合,并为 MI300 和 MI350 GPU 贡献了块级内核支持(#3996)。
将 FP8 扩展到 MoE 架构
像 DeepSeek V3 和 Llama 4 这样的专家混合(MoE)模型将每个 token 路由到专家子集,生成需要通过分组 GEMM 处理的可变大小批次(图 3)。与所有线性层形状相同的密集模型不同,分组 GEMM 需要激活值的每行缩放、权重的每专家列缩放,以及通过偏移张量将行路由到正确专家。
我们通过调整量化流水线以使用正确的数据类型和 AMD 的 Composable Kernel 后端调度,在 ROCm 上启用了 FP8 分组 GEMM(#3955)。
图 3:ROCm 上的 MoE FP8 分组 GEMM 流水线。(A)通过偏移将 token 路由到专家,由融合 Triton 内核量化,通过 Composable Kernel 单次启动完成调度。(B)分组 GEMM 需要每行、每专家列缩放和基于偏移的路由,比密集 GEMM 的统一缩放更复杂。
Triton 内核优化
在确保正确性后,我们转向性能优化。TorchAO 中的 FP8 量化流水线通过多步骤链将张量转换为 FP8:
- 计算每行/列的绝对最大值(absmax)
- 推导缩放因子并应用
- 截断并转换为 FP8
每一步都是单独的内核启动,并在步骤之间将中间张量写入高带宽内存(HBM)。对于每层有数十个专家权重张量的 MoE 模型,这些额外的往返会主导 FP8 开销。在这些形状上,FP8 量化受内存限制:8 位计算成本低,但围绕它的内核启动和 HBM 往返成本高。以下优化通过减少数据移动而非算术运算,在三个粒度级别上提升性能:减少内核启动次数(一级)、使剩余内核高效移动内存(二级)、移除不必要的低级同步(三级)。
一级:减少内核启动次数:
反向传播:
图4展示了通过融合改进FP8量化的方式。反向传播存在两个复合问题。首先,.t().contiguous().t()模式迫使通过HBM进行完整的张量复制,以转换权重布局以实现GEMM兼容性。我们在#3972中移除了这些冗余复制。其次,多步骤的scale-and-cast链式操作在各步骤之间通过HBM显式化中间张量并启动独立内核。我们在多个位置(#4069)将该链式操作融合为单个Triton内核。在8xMI300X GPU上运行DeepSeek-MoE-16B时,这些反向传播融合使反向传播吞吐量提升了4.2倍。
图4:优化前后反向传播FP8量化对比。上游代码在每次量化调用中启动五个通用内核,并在各步骤间通过HBM显式化中间张量。PR #3972消除了转置操作,PR #4069将剩余链式操作融合为单个内核,并通过双内核实现grad_output与激活量化的同时处理。
正向传播:
相同的多内核模式也应用于正向传播路径。量化专家权重时每次调用启动五个通用内核,每个步骤进行24次调用,这增加了约90ms/step的开销。我们用单个融合的Triton内核(#4311)替换了整个链式操作,该内核在专家和输出维度块之间并行处理(图5)。将五次启动合并为一次也使周围GEMM能更早发出。在8x MI325X GPU上运行DeepSeek-V3 671B时,这一改进使端到端吞吐量提升了17%(5,996 → 7,027 tok/s)。
优化前FP8正向传播
融合后FP8正向传播
图5:Perfetto轨迹对比(8xMI325X GPU,DeepSeek-V3 671B)。左侧:FP8(正向优化前)显示每个专家重复的5内核急切链式操作。右侧:融合后FP8(优化后)显示单个triton_fp8_colwise_3d_scale_and_cast内核替代了链式操作。正向传播性能从约19ms提升至约7ms。
图6:在8xMI325X GPU上运行DeepSeek-V3 671B(4层MoE)时,三种配置的GPU时间分类统计。FP8上游配置(V2)每步增加127ms,其中92%消耗在"Others"(通用量化内核)中。融合Triton内核(V4)消除了大部分开销,恢复了89%的BF16→FP8性能差距。
第二层级:优化每个内核的内存移动效率:反向传播中使用的列优先scale内核存在非对齐内存写入(#4113):连续的SIMD通道向间隔K字节的地址写入,每个操作都触发独立内存事务。我们通过在存储前通过LDS(本地数据共享)转置输出块解决了这个问题,并添加了消除冗余HBM读取的单次遍历融合变体。在MI300X GPU上运行DeepSeek-V3 671B时,每MoE层从7,290μs降至1,170μs(6.2倍加速)。
第三层级:移除硬件不必要的同步操作:我们还解决了硬件层级的效率问题。Triton的原子操作(atomic_add、atomic_max、atomic_min)默认使用acquire-release内存顺序,在AMD GPU上会在每个原子操作前后插入内存栅栏,这些昂贵的同步点对于可交换归约操作是不必要的。我们在AMD GPU上将这些操作切换为宽松顺序(#3945),通过torch.version.hip检查确保NVIDIA行为保持不变。
未奏效的尝试:自动调优搜索空间。我们扩展了Triton自动调优搜索空间中MoE FP8内核的候选配置数量,从1个增加到8–16个(#3952),期望更宽泛的搜索能在AMD基于wavefront的架构中找到更快的tile尺寸。然而,在MI300X GPU上对Llama 4模型形状进行基准测试时,未观察到可测量的性能提升,且额外配置增加了首次迭代的编译时间。我们已回滚该更改(#4024)。经验教训:自动调优搜索空间应由硬件约束(wavefront大小、LDS容量、寄存器压力)决定,而非默认扩展更多候选配置。
在三个层面同时优化数据移动,可叠加在图1所示的基线FP8吞吐量提升基础上。在DeepSeek-V3 671B模型形状上,仅通过前向传递内核融合就恢复了89%的量化开销(5,996 → 7,027 tok/s,对比8xMI325X GPU上7,156 BF16基线),而列方向缩放优化使每个MoE层的加速达到6.2倍。
总结与后续计划
本文介绍了我们在TorchAO和TorchTitan中如何针对AMD Instinct GPU优化FP8训练:内核加速、数值稳定性修复以及对MoE架构的支持。
下一代硬件的优化工作仍在继续。我们正在为MI355X GPU开发MXFP8分组GEMM和量化内核,用于前向和反向传播;相关结果将在后续博客中发布。内核融合流水线(#3972 → #4069 → #4113 → #4311)将持续推进,包含更多Triton优化。
这些FP8优化现已集成到标准PyTorch训练栈中:配备AMD Instinct GPU的团队只需升级TorchAO和TorchTitan即可获得,无需安装任何AMD专有组件。这项工作由AMD与Meta/PyTorch工程师团队合作完成。所有贡献已合并至pytorch/ao和pytorch/torchtitan主分支,确保AMD GPU上的FP8训练对更广泛的PyTorch社区开箱即用。
附加资源
- TorchAO float8 README,包含基准测试复现说明
- TorchTitan
- AMD ROCm文档:了解ROCm软件栈的更多细节
- AMD AI开发者计划:获取免费云GPU资源和开发者工具,快速上手ROCm与PyTorch在AMD Instinct GPU上的开发
/post-content /inner-wrap