From 2dc206aad5007469a16d290f716e1a33ab73a1ba Mon Sep 17 00:00:00 2001 From: Anthony Shoumikhin Date: Fri, 2 Oct 2026 18:18:26 -0700 Subject: [PATCH] Skip two SDPA fallback tests on ROCm, where AOTInductor cannot lower them The tests for SDPA calls the Triton kernel does not accept export each call through the regular AOTInductor lowering. On ROCm, two of those exports fail inside AOTInductor's C++ wrapper codegen, after the replacement pass has correctly left the call alone: AssertionError: ([torch.bfloat16, torch.float32, torch.float32, Integer], ['in_ptr0', 'out_ptr0', 'xnumel']) The calls are SDPA with an additive float mask and SDPA with dropout. A generated kernel's argument types and names differ in length. The same exports pass on CUDA, and the other regular-lowering cases pass on ROCm. Skip just those two tests on ROCm, as the file already does for split-K. The root cause is tracked in pytorch/pytorch#199619, and the skip carries a TODO to remove it once that is fixed. --- backends/cuda/tests/test_sdpa_splitk_replacement.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/backends/cuda/tests/test_sdpa_splitk_replacement.py b/backends/cuda/tests/test_sdpa_splitk_replacement.py index 0527d46ba04..e41ffd09201 100644 --- a/backends/cuda/tests/test_sdpa_splitk_replacement.py +++ b/backends/cuda/tests/test_sdpa_splitk_replacement.py @@ -31,6 +31,14 @@ def _require_splitk(tc: unittest.TestCase) -> None: tc.skipTest("split-K is off on ROCm") +def _require_regular_lowering_on_rocm(tc: unittest.TestCase) -> None: + # On ROCm, AOTInductor's C++ wrapper fails to compile the regular lowering of + # some SDPA calls: a generated kernel's argument types and names differ. + # TODO: remove once https://github.com/pytorch/pytorch/issues/199619 is fixed. + if torch.version.hip is not None: + tc.skipTest("AOTInductor cannot lower this SDPA call on ROCm") + + class SDPAModule(nn.Module): """Single-layer model with SDPA and a static KV cache buffer.""" @@ -244,6 +252,7 @@ def assertStays(self, msgs, reason): self.assertTrue(any("Replaced 0 nodes" in m for m in msgs), msgs) def test_float_mask_is_not_replaced(self): + _require_regular_lowering_on_rocm(self) q = _bf16(1, 2, 8, 16) mask = torch.zeros(1, 1, 8, 8, dtype=torch.bfloat16) self.assertStays(self._logs(q, q, q, mask), "attn_mask must have dtype") @@ -279,6 +288,7 @@ def test_causal_with_other_lengths_is_not_replaced(self): self.assertStays(self._logs(q, kv, kv, is_causal=True), "Causal masking") def test_dropout_is_not_replaced(self): + _require_regular_lowering_on_rocm(self) q = _bf16(1, 2, 8, 16) self.assertStays(self._logs(q, q, q, dropout_p=0.1), "dropout_p must be 0.0")