Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions backends/cuda/tests/test_sdpa_splitk_replacement.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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")

Expand Down
Loading