viable/strict/1788216615: [inductor] Enable BF16 index_add lowering on SM90+ and ROCm (#195192)
PyTorch Inductor 的 index_add 分解在 Triton 添加 BF16 atomic_add 支持后,启用了在 ROCm 和 NVIDIA SM90+ 上的 BF16 index_add lowering,使 index_add 可以参与融合而非回退到 ATen。NVIDIA 限制在 SM90+ 是因为原生 BF16 原子指令边界;pre-SM90 使用 CAS 循环。测试覆盖了 BF16 index_select 反向、非连续 index_add、负维度、int32 索引、重复项、非单位 alpha 标量索引、确定性算法回退等。
Development
- First Reportviable/strict/1788216615: [inductor] Enable BF16 index_add lowering on SM90+ and ROCm (#195192)PyTorch Core
- Current AssessmentPyTorch 持续改进 Inductor 的代码生成能力,这有助于提升在 AMD 和 NVIDIA 硬件上的训练和推理性能。对 ROCm 的支持表明 AMD 在 AI 生态系统中的地位日益重要。Agent Pulse · analysis
PyTorch Inductor 的 index_add 分解此前保留了一个过时的仅限 OSS 的 BF16 回退,在 Triton 添加 BF16 atomic_add 支持后,该分解现在可以在 ROCm 和 NVIDIA SM90+ 上启用,允许 index_add 参与融合而不是成为 ATen 回退内核。Triton 支持在 triton-lang/triton#6519 中添加,并在 MI300 和 NVIDIA GPU 上进行了上游测试。NVIDIA 限制在 SM90+ 是因为这是原生 BF16 原子指令边界;pre-SM90 Triton 使用 CAS 循环。此更改还在分解前规范化标量索引,因为生成的 index_put lowering 期望张量索引至少有一个维度。测试覆盖了 BF16 index_select 反向、非连续 index_add、负维度、int32 索引、重复项、非单位 alpha 标量索引、确定性算法回退(无原子操作)。ROCm 和 NVIDIA SM90+ 运行 lowering/codegen 断言;较旧的 NVIDIA GPU 保留 ATen 回退。
此更改表明 Inductor 正在逐步消除对 ATen 回退的依赖,通过利用 Triton 的新原子操作支持,使更多操作可以融合。BF16 原子操作在 SM90+ 上的可用性是一个关键边界,这可能会影响未来在旧硬件上的优化策略。
PyTorch 持续改进 Inductor 的代码生成能力,这有助于提升在 AMD 和 NVIDIA 硬件上的训练和推理性能。对 ROCm 的支持表明 AMD 在 AI 生态系统中的地位日益重要。
此更改通过减少内核启动开销和内存访问,可能提升使用 PyTorch 进行训练和推理的用户的性能,尤其是使用 BF16 和 AMD GPU 的用户。
未来,随着 Triton 对更多原子操作和硬件的支持,Inductor 可能会进一步扩展可融合操作的范围,减少对回退内核的依赖,从而提升整体性能。