viable/strict/1790977859: [ROCm] Enable test_grouped_mm on ROCm (#199048)
PyTorch PR #199048 移除了 DistMatrixOpsTest.test_grouped_mm 在 ROCm 上的 @unittest.skipIf(TEST_WITH_ROCM, "ROCm doesn't support CUTLASS") 跳过标记。证据称该测试为纯 bf16,在 ROCm 上 backend="cutlass" 参数化实际走架构无关的 _grouped_mm 回退路径,CUTLASS 从未被调用;backend="cublaslt" 变体由既有 CUDA-Toolkit 版本守卫自行跳过。相邻 SM90OrLater 守卫在 gfx942/gfx950/gfx90a 上满足。测试计划在 gfx942(MI300X)与 gfx90a(MI200)上运行 pytest test/distributed/tensor/test_matrix_ops.py -k test_grouped_mm -v,结果均为 4 passed(cutlass 变体)、4 skipped(cublaslt 自跳过);直接 _grouped_mm bf16 探针与逐组参考循环相比 max abs diff 为 0.0。该 PR 修复 #168447 与 #168448,由 jeffdaily 批准,作者声明在 Claude(AI)协助下完成。
发展脉络
- 首次出现viable/strict/1790977859: [ROCm] Enable test_grouped_mm on ROCm (#199048)PyTorch Core
- 当前判断这是一次测试覆盖面的修正,而非硬件或算子能力声明。它反映 AMD ROCm 与 PyTorch 分布式张量算子测试矩阵在持续对齐,但证据未提供性能、吞吐或生产可用性数据,不宜外推为 ROCm 在 grouped GEMM 上已具备与 CUDA 对等能力。Agent Pulse · 分析
PyTorch 合并 PR #199048,移除 ROCm 上对 test_grouped_mm 的跳过。证据给出的理由是:该测试为纯 bf16,ROCm 上 backend="cutlass" 参数化实际路由到架构无关的 _grouped_mm 回退,CUTLASS 并未被调用,因此原跳过理由属于误称;backend="cublaslt" 变体则由既有 CUDA-Toolkit 版本守卫自行跳过。验证在 gfx942(MI300X)与 gfx90a(MI200)上进行,结果为 4 passed、4 skipped,直接 bf16 探针与逐组参考循环 max abs diff 为 0.0。该改动修复 #168447 与 #168448,由 jeffdaily 批准,作者声明在 Claude(AI)协助下完成。
从证据看,关键点不是 ROCm 获得了 CUTLASS 支持,而是测试的跳过条件与真实执行路径不一致:backend="cutlass" 在 ROCm 上落到 _grouped_mm 回退,CUTLASS 未被调用。可验证的下一信号是查看该回退路径在 gfx942/gfx90a 之外架构上的覆盖情况,以及 cublaslt 变体的 CUDA-Toolkit 版本守卫是否仍会长期遮蔽测试。
这是一次测试覆盖面的修正,而非硬件或算子能力声明。它反映 AMD ROCm 与 PyTorch 分布式张量算子测试矩阵在持续对齐,但证据未提供性能、吞吐或生产可用性数据,不宜外推为 ROCm 在 grouped GEMM 上已具备与 CUDA 对等能力。
对使用 ROCm(MI300X/MI200)运行 PyTorch 分布式张量算子的团队,该改动意味着 grouped_mm 的 bf16 回退路径进入常规 CI 覆盖,降低回归漏检风险;但证据未给出性能收益,采购或迁移决策不应据此推断成本优势。
若后续 PR 继续清理以 CUTLASS 为名的 ROCm 跳过条件,可视为 ROCm 测试覆盖系统性收敛的信号;反之若 cublaslt 变体长期自跳过,则说明该路径在 ROCm 上仍未被实际验证。