PyTorch 原生支持 AMD FP8 训练:TorchTitan 与 TorchAO 协同实现 MoE 架构性能飞跃
在 PyTorch Conference 2025 上,PyTorch 团队展示了利用 AMD Primus-Turbo 库在 Instinct 集群上实现千卡线性扩展的能力。如今,这些关键的 AMD 优化已全面开源并合并至上游 PyTorch 生态,使得 TorchTitan 能够原生支持 AMD Instinct GPU,并开箱即用地提供具有竞争力的 FP8 训练性能。
核心突破:从格式兼容到 MoE 架构加速
此次更新的核心在于解决了 AMD GPU 特有的 FNUZ (Finite, No NaN, Unsigned Zero) 数值格式兼容性问题,并针对 Mixture-of-Experts (MoE) 架构进行了深度的内核优化。
1. 原生 AMD FP8 格式支持
此前,TorchAO 默认使用 NVIDIA 的 FP8 格式,导致在 AMD Instinct GPU 上因最大数值(Max Value)差异(AMD 为 240,NVIDIA 为 4369)而产生静默错误。激活值被截断,梯度损坏,且由于无 NaN/Inf 编码,错误未被捕获。
- 自动检测机制:TorchAO 新增硬件自动检测功能,自动选择正确的 FP8 数据类型和最大值,无需手动硬编码。
- 精度修复:修复了平台特定的损失基准线,确保 FNUZ 数值逻辑的正确性。
- 内核扩展:为 MI300 和 MI350 GPU 贡献了 Blockwise 量化内核支持,支持 Tensorwise、Rowwise、Blockwise 及 MXFP8 四种量化策略。
2. MoE 架构的分组 GEMM 优化
对于 DeepSeek-V3 等 MoE 模型,FP8 量化带来了显著的开销。团队通过以下手段大幅回收了这部分开销:
- 分组 GEMM 启用:在 ROCm 上启用分组 GEMM,适配 AMD 后端,实现单次 Launch 完成量化与分发。
- Triton 内核融合:将 FP8 量化流水线中的
absmax计算、Scale 推导、Clamp 和 Cast 步骤融合为单个 Triton 内核,消除了中间张量对高带宽内存 (HBM) 的占用。 - 性能提升:
- DeepSeek-V3 671B:单个 MoE 层速度提升 6.2 倍 (7,290µs → 1,170µs)。
- 整体吞吐:端到端训练速度提升 17%,成功回收 89% 的 FP8 量化开销。
3. 密集模型性能增益
在 Llama3-8B 等密集模型上,通过行向 FP8 (Rowwise FP8) 训练,相比 BF16 实现了 13.4% 的吞吐量提升。虽然峰值内存占用相似,但得益于 AMD GPU 更快的 FP8 矩阵核心,计算效率显著提升。
实际应用价值
- 开发者:无需额外配置即可在 AMD Instinct 集群上运行最新的大模型训练任务,享受与 NVIDIA 相当甚至更优的 FP8 训练效率。
- 企业用户:降低了使用 AMD GPU 进行大规模模型训练的成本,特别是在处理超大规模 MoE 模型时,显著缩短了训练周期。
- 生态影响:进一步打破了单一硬件生态的垄断,推动了 PyTorch 在异构计算领域的标准化与统一性。
“我们将 PyTorch Conference 上展示的 AMD 优化全面开源,旨在让所有开发者都能平等地利用 AMD Instinct GPU 的强大算力,特别是在 FP8 训练这一关键领域。” —— PyTorch 官方团队
关键技术指标
| 指标 | 数值/描述 |
|---|---|
| Llama3-8B 吞吐提升 | +13.4% (FP8 vs BF16) |
| DeepSeek-V3 单层加速 | 6.2 倍 (7,290µs → 1,170µs) |
| MoE 量化开销回收 | 89% |
| 端到端训练加速 | +17% |
| 支持硬件 | AMD Instinct (MI300X, MI325X, MI350X) |
| 支持格式 | e4m3fnuz (FNUZ) |