From ace1aded42a1adabd46227716ac7fa8d899ea3bb Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stig-Arne=20Gr=C3=B6nroos?= Date: Thu, 19 Mar 2026 12:11:59 +0200 Subject: [PATCH 1/2] [Bugfix][ROCm] Fix lru_cache on paged_mqa_logits_module MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fix lru_cache on paged_mqa_logits_module by moving to module scope The function was defined inside rocm_fp8_paged_mqa_logits, causing a new function object (and thus a new empty cache) to be created on every call, defeating the purpose of lru_cache. As a consequence, the module was reimported on each call. Signed-off-by: Stig-Arne Grönroos --- .../v1/attention/ops/rocm_aiter_mla_sparse.py | 38 +++++++++---------- 1 file changed, 19 insertions(+), 19 deletions(-) diff --git a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py index 878ae3aac521..b37415413100 100644 --- a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -273,6 +273,25 @@ def fp8_paged_mqa_logits_torch( return logits +@functools.lru_cache +def paged_mqa_logits_module(): + paged_mqa_logits_module_path = None + if importlib.util.find_spec("aiter.ops.triton.pa_mqa_logits") is not None: + paged_mqa_logits_module_path = "aiter.ops.triton.pa_mqa_logits" + elif ( + importlib.util.find_spec("aiter.ops.triton.attention.pa_mqa_logits") is not None + ): + paged_mqa_logits_module_path = "aiter.ops.triton.attention.pa_mqa_logits" + + if paged_mqa_logits_module_path is not None: + try: + module = importlib.import_module(paged_mqa_logits_module_path) + return module + except ImportError: + return None + return None + + def rocm_fp8_paged_mqa_logits( q_fp8: torch.Tensor, kv_cache_fp8: torch.Tensor, @@ -305,25 +324,6 @@ def rocm_fp8_paged_mqa_logits( """ from vllm._aiter_ops import rocm_aiter_ops - @functools.lru_cache - def paged_mqa_logits_module(): - paged_mqa_logits_module_path = None - if importlib.util.find_spec("aiter.ops.triton.pa_mqa_logits") is not None: - paged_mqa_logits_module_path = "aiter.ops.triton.pa_mqa_logits" - elif ( - importlib.util.find_spec("aiter.ops.triton.attention.pa_mqa_logits") - is not None - ): - paged_mqa_logits_module_path = "aiter.ops.triton.attention.pa_mqa_logits" - - if paged_mqa_logits_module_path is not None: - try: - module = importlib.import_module(paged_mqa_logits_module_path) - return module - except ImportError: - return None - return None - aiter_paged_mqa_logits_module = None if rocm_aiter_ops.is_enabled(): aiter_paged_mqa_logits_module = paged_mqa_logits_module() From 2b5fa072d2e83ef41879fd7f4364c03f001b69a9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stig-Arne=20Gr=C3=B6nroos?= Date: Thu, 19 Mar 2026 17:01:21 +0200 Subject: [PATCH 2/2] fix: another instance of the same pattern MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Stig-Arne Grönroos --- .../v1/attention/ops/rocm_aiter_mla_sparse.py | 39 ++++++++++--------- 1 file changed, 20 insertions(+), 19 deletions(-) diff --git a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py index b37415413100..9d1da5b53be5 100644 --- a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -400,6 +400,26 @@ def fp8_mqa_logits_torch( return logits +@functools.lru_cache +def mqa_logits_module(): + mqa_logits_module_path = None + if importlib.util.find_spec("aiter.ops.triton.fp8_mqa_logits") is not None: + mqa_logits_module_path = "aiter.ops.triton.fp8_mqa_logits" + elif ( + importlib.util.find_spec("aiter.ops.triton.attention.fp8_mqa_logits") + is not None + ): + mqa_logits_module_path = "aiter.ops.triton.attention.fp8_mqa_logits" + + if mqa_logits_module_path is not None: + try: + module = importlib.import_module(mqa_logits_module_path) + return module + except ImportError: + return None + return None + + def rocm_fp8_mqa_logits( q: torch.Tensor, kv: tuple[torch.Tensor, torch.Tensor], @@ -429,25 +449,6 @@ def rocm_fp8_mqa_logits( # path after aiter merge this kernel into main from vllm._aiter_ops import rocm_aiter_ops - @functools.lru_cache - def mqa_logits_module(): - mqa_logits_module_path = None - if importlib.util.find_spec("aiter.ops.triton.fp8_mqa_logits") is not None: - mqa_logits_module_path = "aiter.ops.triton.fp8_mqa_logits" - elif ( - importlib.util.find_spec("aiter.ops.triton.attention.fp8_mqa_logits") - is not None - ): - mqa_logits_module_path = "aiter.ops.triton.attention.fp8_mqa_logits" - - if mqa_logits_module_path is not None: - try: - module = importlib.import_module(mqa_logits_module_path) - return module - except ImportError: - return None - return None - aiter_mqa_logits_module = None if rocm_aiter_ops.is_enabled(): aiter_mqa_logits_module = mqa_logits_module()