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")