From 383bec9c2a50426fd12776be7367f448b6c9fb32 Mon Sep 17 00:00:00 2001 From: Bortlesboat Date: Sun, 10 May 2026 15:38:23 -0400 Subject: [PATCH] [ROCm] Slice MLA decode fallback cache blocks Co-authored-by: OpenAI Codex Signed-off-by: Bortlesboat --- .../test_rocm_aiter_mla_sparse_fallback.py | 165 ++++++++++++++++++ .../v1/attention/ops/rocm_aiter_mla_sparse.py | 40 ++++- 2 files changed, 204 insertions(+), 1 deletion(-) create mode 100644 tests/v1/attention/test_rocm_aiter_mla_sparse_fallback.py diff --git a/tests/v1/attention/test_rocm_aiter_mla_sparse_fallback.py b/tests/v1/attention/test_rocm_aiter_mla_sparse_fallback.py new file mode 100644 index 000000000000..a24fdeafd778 --- /dev/null +++ b/tests/v1/attention/test_rocm_aiter_mla_sparse_fallback.py @@ -0,0 +1,165 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import importlib.util +import sys +from pathlib import Path +from types import ModuleType, SimpleNamespace + +import torch + + +def _package(name: str) -> ModuleType: + module = ModuleType(name) + module.__path__ = [] + return module + + +def _load_rocm_aiter_mla_sparse(monkeypatch): + class _FakeCurrentPlatform: + def is_rocm(self) -> bool: + return False + + def fp8_dtype(self) -> torch.dtype: + return torch.float8_e4m3fn + + class _FakeTriton: + def jit(self, fn=None, **kwargs): + def decorator(inner_fn): + return inner_fn + + return decorator(fn) if fn is not None else decorator + + tl = SimpleNamespace(constexpr=object) + stubs = { + "vllm": _package("vllm"), + "vllm.forward_context": SimpleNamespace(get_forward_context=lambda: None), + "vllm.platforms": SimpleNamespace(current_platform=_FakeCurrentPlatform()), + "vllm.triton_utils": SimpleNamespace(tl=tl, triton=_FakeTriton()), + "vllm.utils": _package("vllm.utils"), + "vllm.utils.torch_utils": SimpleNamespace(LayerNameType=str), + "vllm.v1": _package("vllm.v1"), + "vllm.v1.attention": _package("vllm.v1.attention"), + "vllm.v1.attention.backends": _package("vllm.v1.attention.backends"), + "vllm.v1.attention.backends.mla": _package("vllm.v1.attention.backends.mla"), + "vllm.v1.attention.backends.mla.indexer": SimpleNamespace( + DeepseekV32IndexerMetadata=object + ), + "vllm.v1.attention.ops": _package("vllm.v1.attention.ops"), + "vllm.v1.attention.ops.common": SimpleNamespace( + pack_seq_triton=lambda *args, **kwargs: None, + unpack_seq_triton=lambda *args, **kwargs: None, + ), + } + for name, module in stubs.items(): + monkeypatch.setitem(sys.modules, name, module) + + source_path = ( + Path(__file__).resolve().parents[3] + / "vllm" + / "v1" + / "attention" + / "ops" + / "rocm_aiter_mla_sparse.py" + ) + spec = importlib.util.spec_from_file_location( + "_test_rocm_aiter_mla_sparse", source_path + ) + assert spec is not None + assert spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def test_forward_decode_fallback_dequantizes_only_referenced_cache_blocks(monkeypatch): + module = _load_rocm_aiter_mla_sparse(monkeypatch) + + block_size = 4 + cache_width = 32 + swa_k_cache = torch.arange(6 * block_size * cache_width, dtype=torch.uint8).reshape( + 6, block_size, cache_width + ) + kv_cache = torch.arange(8 * block_size * cache_width, dtype=torch.uint8).reshape( + 8, block_size, cache_width + ) + + q = torch.zeros((2, 2, 6), dtype=torch.bfloat16) + output = torch.empty((2, 2, 4), dtype=torch.bfloat16) + swa_indices = torch.tensor( + [[2 * block_size + 1, 4 * block_size + 3, -1], [2 * block_size, -1, -1]], + dtype=torch.int64, + ) + topk_indices = torch.tensor( + [[[6 * block_size + 2, -1]], [[3 * block_size + 1, 6 * block_size]]], + dtype=torch.int64, + ) + swa_lens = torch.tensor([2, 1], dtype=torch.int32) + topk_lens = torch.tensor([1, 2], dtype=torch.int32) + + dequant_inputs = [] + captured_decode = {} + + def fake_dequantize(quant_k_cache, *, head_dim, nope_head_dim, rope_head_dim): + dequant_inputs.append(quant_k_cache.clone()) + return torch.empty( + (quant_k_cache.shape[0], quant_k_cache.shape[1], 1, head_dim), + dtype=torch.bfloat16, + ) + + def fake_sparse_decode( + *, + q, + blocked_k, + indices_in_kvcache, + topk_length, + scale, + head_dim, + attn_sink, + extra_blocked_k=None, + extra_indices_in_kvcache=None, + extra_topk_length=None, + ): + captured_decode["indices_in_kvcache"] = indices_in_kvcache.clone() + captured_decode["extra_indices_in_kvcache"] = ( + None + if extra_indices_in_kvcache is None + else extra_indices_in_kvcache.clone() + ) + return torch.full((q.shape[0], q.shape[2], head_dim), 7, dtype=torch.bfloat16) + + monkeypatch.setattr(module, "rocm_dequantize_blocked_k_cache", fake_dequantize) + monkeypatch.setattr(module, "rocm_ref_sparse_attn_decode", fake_sparse_decode) + + module.rocm_forward_decode_fallback( + q=q, + kv_cache=kv_cache, + swa_k_cache=swa_k_cache, + swa_only=False, + topk_indices=topk_indices, + topk_lens=topk_lens, + swa_indices=swa_indices, + swa_lens=swa_lens, + attn_sink=None, + scale=1.0, + head_dim=4, + nope_head_dim=4, + rope_head_dim=2, + output=output, + ) + + assert len(dequant_inputs) == 2 + torch.testing.assert_close(dequant_inputs[0], swa_k_cache[[2, 4]]) + torch.testing.assert_close(dequant_inputs[1], kv_cache[[3, 6]]) + + expected_swa_indices = torch.tensor( + [[[1, 7, -1]], [[0, -1, -1]]], dtype=torch.int64 + ) + expected_topk_indices = torch.tensor([[[6, -1]], [[1, 4]]], dtype=torch.int64) + torch.testing.assert_close( + captured_decode["indices_in_kvcache"], expected_swa_indices + ) + torch.testing.assert_close( + captured_decode["extra_indices_in_kvcache"], expected_topk_indices + ) + torch.testing.assert_close(output, torch.full_like(output, 7)) diff --git a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py index 5d0343ffd607..6979df61863b 100644 --- a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -1013,6 +1013,34 @@ def rocm_dequantize_blocked_k_cache( return result +def _gather_referenced_cache_blocks( + quant_k_cache: torch.Tensor, + indices_in_kvcache: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """Gather used cache blocks and remap flattened token indices.""" + block_size = quant_k_cache.shape[1] + valid_mask = indices_in_kvcache >= 0 + valid_indices = indices_in_kvcache[valid_mask] + if valid_indices.numel() == 0: + return quant_k_cache[:1], indices_in_kvcache + + valid_block_indices = torch.div( + valid_indices, + block_size, + rounding_mode="floor", + ) + block_indices = torch.unique(valid_block_indices, sorted=True) + compact_k_cache = quant_k_cache.index_select(0, block_indices.to(torch.long)) + + compact_block_indices = torch.searchsorted(block_indices, valid_block_indices) + remapped_indices = indices_in_kvcache.clone() + remapped_valid_indices = compact_block_indices * block_size + ( + valid_indices % block_size + ) + remapped_indices[valid_mask] = remapped_valid_indices.to(remapped_indices.dtype) + return compact_k_cache, remapped_indices + + def rocm_ref_sparse_attn_decode( q: torch.Tensor, blocked_k: torch.Tensor, @@ -1099,6 +1127,10 @@ def rocm_forward_decode_fallback( rope_head_dim: int, output: torch.Tensor, ) -> None: + swa_k_cache, swa_indices = _gather_referenced_cache_blocks( + swa_k_cache, + swa_indices, + ) blocked_swa = rocm_dequantize_blocked_k_cache( swa_k_cache, head_dim=head_dim, @@ -1106,8 +1138,14 @@ def rocm_forward_decode_fallback( rope_head_dim=rope_head_dim, ) blocked_extra = None + compact_topk_indices = None if not swa_only: assert kv_cache is not None + assert topk_indices is not None + kv_cache, compact_topk_indices = _gather_referenced_cache_blocks( + kv_cache, + topk_indices, + ) blocked_extra = rocm_dequantize_blocked_k_cache( kv_cache, head_dim=head_dim, @@ -1123,7 +1161,7 @@ def rocm_forward_decode_fallback( head_dim=head_dim, attn_sink=attn_sink[: q.shape[1]] if attn_sink is not None else None, extra_blocked_k=blocked_extra, - extra_indices_in_kvcache=topk_indices, + extra_indices_in_kvcache=compact_topk_indices, extra_topk_length=topk_lens, ) output.copy_(attn_out.to(output.dtype))