diff --git a/tests/kernels/attention/test_rocm_aiter_mla_op_registration.py b/tests/kernels/attention/test_rocm_aiter_mla_op_registration.py index 892a9ac76bb9..659211521a07 100644 --- a/tests/kernels/attention/test_rocm_aiter_mla_op_registration.py +++ b/tests/kernels/attention/test_rocm_aiter_mla_op_registration.py @@ -1,11 +1,11 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""ROCm custom op schema test for AITER MLA decode. +"""ROCm custom op schema tests for AITER MLA decode. -A single ``opcheck`` call on ``torch.ops.vllm.rocm_aiter_mla_decode_fwd`` -verifies that the custom op is registered and that its schema and fake -implementation are consistent with the real kernel: fake-tensor support for -torch.compile tracing and the ``mutates_args=["o"]`` in-place output aliasing. +``opcheck`` verifies that the decode ops are registered and that their schemas +and fake implementations are consistent with the real kernels: fake-tensor +support for torch.compile tracing and ``mutates_args=["o"]`` in-place output +aliasing. """ import pytest @@ -14,16 +14,9 @@ from tests.kernels.utils import opcheck from vllm.platforms import current_platform -_SKIP_NON_MI3XX = True -if current_platform.is_rocm(): - from vllm.platforms.rocm import on_mi3xx - - _SKIP_NON_MI3XX = not on_mi3xx() - -pytestmark = [ - pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm-specific tests"), - pytest.mark.skipif(_SKIP_NON_MI3XX, reason="MI300/MI350 ROCm only"), -] +pytestmark = pytest.mark.skipif( + not current_platform.is_rocm(), reason="ROCm-specific tests" +) Q_HEAD_DIM = 576 # kv_lora_rank + qk_rope_head_dim V_HEAD_DIM = 512 # kv_lora_rank @@ -31,6 +24,10 @@ def _require_aiter(): from vllm._aiter_ops import is_aiter_found_and_supported + from vllm.platforms.rocm import get_cdna_version + + if get_cdna_version() not in (3, 4): + pytest.skip("AITER MLA requires CDNA 3 or 4") if not is_aiter_found_and_supported(): pytest.skip("aiter is required on supported ROCm hardware for this test") @@ -77,3 +74,34 @@ def test_mla_decode_fwd_op_schema() -> None: "reduce_partial_map": None, }, ) + + +@torch.inference_mode() +def test_mla_decode_fwd_lse_op_schema() -> None: + """Validate graph registration and mutation schema for LSE decode.""" + _require_aiter() + # Import ensures the custom op is registered. + from vllm._aiter_ops import rocm_aiter_ops # noqa: F401 + + batch_size, nhead = 2, 16 + q = torch.randn(batch_size, nhead, Q_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + kv_buffer = torch.randn(32, 1, 1, Q_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + o = torch.zeros(batch_size, nhead, V_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device="cuda") + kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device="cuda") * 16 + kv_indices = torch.arange(32, dtype=torch.int32, device="cuda") + kv_last_page_lens = torch.ones(batch_size, dtype=torch.int32, device="cuda") + + opcheck( + torch.ops.vllm.rocm_aiter_mla_decode_fwd_lse, + (q, kv_buffer, o, qo_indptr, 1), + { + "kv_indptr": kv_indptr, + "kv_indices": kv_indices, + "kv_last_page_lens": kv_last_page_lens, + "sm_scale": Q_HEAD_DIM**-0.5, + "logit_cap": 0.0, + "q_scale": None, + "kv_scale": None, + }, + ) diff --git a/tests/kernels/attention/test_rocm_aiter_mla_sink.py b/tests/kernels/attention/test_rocm_aiter_mla_sink.py new file mode 100644 index 000000000000..6de46afe3e87 --- /dev/null +++ b/tests/kernels/attention/test_rocm_aiter_mla_sink.py @@ -0,0 +1,512 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Correctness tests for ROCm AITER sparse MLA attention sinks.""" + +from types import SimpleNamespace + +import pytest +import torch + +from vllm.platforms import current_platform +from vllm.utils.torch_utils import set_random_seed + +pytestmark = pytest.mark.skipif( + not current_platform.is_rocm(), reason="ROCm-specific tests" +) + +Q_HEAD_DIM = 576 +V_HEAD_DIM = 512 +SM_SCALE = Q_HEAD_DIM**-0.5 + + +def _gamma_fp32(operations: int) -> float: + u = torch.finfo(torch.float32).eps / 2 + return operations * u / (1 - operations * u) + + +def _sink_reference(q, keys, sinks, scale): + # Reference the stored operands, so input quantization is not kernel error. + q, keys, sinks = q.double(), keys.double(), sinks.double() + scores = q @ keys.T * scale + logits = torch.cat((scores, sinks[:, None]), dim=-1) + probabilities = logits.softmax(dim=-1)[:, :-1] + values = keys[:, :V_HEAD_DIM] + expected = probabilities @ values + value_scale = probabilities @ values.abs() + # FP32 dot accumulation and conversion/multiplication of the scale. + score_error = _gamma_fp32(q.shape[-1] + 2) * (q.abs() @ keys.abs().T) * abs(scale) + score_error = score_error.amax(dim=-1) if keys.shape[0] else torch.zeros_like(sinks) + return expected, logits.logsumexp(dim=-1), value_scale, score_error + + +def _assert_sink_output_close( + actual, expected, value_scale, score_error, probability_dtype, native +): + assert actual.shape == expected.shape + # Budget the identifiable rounding stages against the absolute value + # contributions: output-relative ULPs become unbounded when values cancel. + # Rounding P before P@V costs u(P) * sum(P*abs(V)). + # AITER rounds output before and after sink scaling; Triton rounds once. + u_probability = torch.finfo(probability_dtype).eps / 2 + u_output = torch.finfo(actual.dtype).eps / 2 + output_rounds = 2 if native else 1 + amplification = (1 + u_output) ** output_rounds + relative_error = (1 + u_probability) * torch.exp(2 * score_error) - 1 + allowance = ( + relative_error.unsqueeze(-1) * value_scale * amplification + + (amplification - 1) * expected.abs() + + output_rounds * torch.finfo(actual.dtype).tiny * u_output + ) + actual = actual.to(device=expected.device, dtype=torch.float64) + assert torch.isfinite(actual).all() + assert torch.count_nonzero(actual[value_scale == 0]) == 0 + error = (actual - expected).abs() + max_ratio = (error / allowance.clamp_min(1e-300)).max().item() + assert max_ratio <= 1, f"exceeded rounding budget by {max_ratio:.3f}x" + + +def _require_aiter() -> None: + from vllm._aiter_ops import is_aiter_found_and_supported + from vllm.platforms.rocm import get_cdna_version + + if get_cdna_version() not in (3, 4): + pytest.skip("AITER MLA requires CDNA 3 or 4") + + if not is_aiter_found_and_supported(): + pytest.skip("aiter is required on supported ROCm hardware for this test") + + +@pytest.mark.parametrize( + ("real_heads", "cache_kind"), + [ + (1, "bf16"), + (4, "bf16"), + (8, "bf16"), + (8, "fp8"), + (32, "fp8"), + (64, "fp8"), + (12, "bf16"), + (16, "bf16"), + (24, "bf16"), + (32, "bf16"), + (40, "bf16"), + (48, "bf16"), + (64, "bf16"), + (80, "bf16"), + ], +) +@torch.inference_mode() +def test_sparse_mla_sink_matches_ragged_reference( + real_heads: int, cache_kind: str +) -> None: + """Exercise ragged sink decode, head padding, and FP8 scale forwarding.""" + _require_aiter() + from vllm.v1.attention.backends.mla.rocm_aiter_mla import AiterMLAHelper + from vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseImpl, + ) + + set_random_seed(real_heads * 17 + (cache_kind == "fp8")) + device = torch.device("cuda") + seq_lens = [0, 1, 5, 23, 17] + batch_size = len(seq_lens) + + q_source = torch.randn(batch_size, real_heads, Q_HEAD_DIM, device=device) * 0.2 + + pool_size = sum(seq_lens) + 52 + kv_source = torch.randn(pool_size, 1, Q_HEAD_DIM, device=device) * 0.2 + if cache_kind == "fp8": + fp8_dtype = current_platform.fp8_dtype() + q_scale = torch.tensor(0.5, dtype=torch.float32, device=device) + kv_scale = torch.tensor(0.25, dtype=torch.float32, device=device) + q_real = (q_source / q_scale).to(fp8_dtype) + kv = (kv_source / kv_scale).to(fp8_dtype) + q_ref = q_real.float() * q_scale + kv_ref = kv.float() * kv_scale + else: + q_scale = kv_scale = None + q_real = q_source.to(torch.bfloat16) + kv = kv_source.to(torch.bfloat16) + q_ref = q_real.float() + kv_ref = kv.float() + q = AiterMLAHelper.get_mla_padded_q(real_heads, q_real) + + indices = torch.randperm(pool_size, device=device)[: sum(seq_lens)].to(torch.int32) + kv_indptr = torch.tensor( + [0] + [sum(seq_lens[:i]) for i in range(1, len(seq_lens) + 1)], + dtype=torch.int32, + device=device, + ) + + # Non-None garbage proves the sink path cannot accidentally select the + # gfx942 persistent kernel, which has no return-LSE code object. + metadata = SimpleNamespace( + attn_out_dtype=torch.bfloat16, + qo_indptr=torch.arange(batch_size + 1, dtype=torch.int32, device=device), + paged_kv_indptr=kv_indptr, + paged_kv_indices=indices, + paged_kv_last_page_len=torch.ones(batch_size, dtype=torch.int32, device=device), + work_meta_data=torch.tensor([123], dtype=torch.int32), + work_indptr=None, + work_info_set=None, + reduce_indptr=None, + reduce_final_map=None, + reduce_partial_map=None, + num_prefills=0, + num_decodes=batch_size, + num_decode_tokens=batch_size, + max_query_len=1, + ) + sinks = torch.linspace(-2.0, 6.0, real_heads, device=device) + + impl = object.__new__(ROCMAiterMLASparseImpl) + impl.num_heads = real_heads + impl.kv_lora_rank = V_HEAD_DIM + impl.kv_cache_dtype = cache_kind + impl.scale = SM_SCALE + impl.sinks = sinks + + output, lse = impl._forward_mla( + SimpleNamespace(_q_scale=q_scale, _k_scale=kv_scale), q, kv, metadata + ) + kv_flat = kv_ref[:, 0] + references = [] + start = 0 + for batch_idx, seq_len in enumerate(seq_lens): + rows = kv_flat[indices[start : start + seq_len].long()] + references.append(_sink_reference(q_ref[batch_idx], rows, sinks, SM_SCALE)) + start += seq_len + expected, expected_lse, value_scale, score_error = ( + torch.stack(values) for values in zip(*references) + ) + + assert output.dtype == torch.bfloat16 + if lse is not None: + assert lse.dtype == torch.float32 + assert lse.shape == expected_lse.shape + lse_error = (lse.double() - expected_lse).abs() + rounded_lse = expected_lse.float() + positive_inf = torch.full_like(rounded_lse, float("inf")) + ulp = torch.maximum( + torch.nextafter(rounded_lse, positive_inf) - rounded_lse, + rounded_lse - torch.nextafter(rounded_lse, -positive_inf), + ).double() + lse_allowance = score_error + 2 * ulp + assert torch.all(lse_error <= lse_allowance) + torch.testing.assert_close(lse[0].double(), expected_lse[0], atol=0, rtol=0) + else: + from vllm.platforms.rocm import on_gfx942 + + assert cache_kind == "bf16" and real_heads in (40, 48, 64) and on_gfx942() + _assert_sink_output_close( + output, + expected, + value_scale, + score_error, + q_real.dtype, + native=lse is not None, + ) + + +@pytest.mark.parametrize( + ("q_dtype", "kv_dtype"), + [ + (torch.float16, torch.bfloat16), + (torch.bfloat16, torch.float16), + ], +) +def test_sparse_mla_sink_rejects_unsupported_aiter_dtypes( + q_dtype: torch.dtype, kv_dtype: torch.dtype +) -> None: + from vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseImpl, + ) + + impl = object.__new__(ROCMAiterMLASparseImpl) + impl.num_heads = 16 + impl.kv_lora_rank = V_HEAD_DIM + impl.kv_cache_dtype = "auto" + impl.scale = SM_SCALE + impl.sinks = torch.zeros(16, dtype=torch.float32) + metadata = SimpleNamespace( + attn_out_dtype=torch.bfloat16, + num_prefills=0, + num_decodes=1, + num_decode_tokens=1, + max_query_len=1, + ) + + with pytest.raises(ValueError, match="both use BF16 or both use FP8"): + impl._forward_mla( + SimpleNamespace(_q_scale=None, _k_scale=None), + torch.empty(1, 16, Q_HEAD_DIM, dtype=q_dtype), + torch.empty(1, 1, Q_HEAD_DIM, dtype=kv_dtype), + metadata, + ) + + +def _make_noncontiguous_sink() -> torch.Tensor: + return torch.empty(8, dtype=torch.float32)[::2] + + +@pytest.mark.parametrize( + ("sinks", "match"), + [ + (torch.empty(4, dtype=torch.bfloat16), "must be float32"), + (torch.empty(2, 2, dtype=torch.float32), "must have shape"), + (_make_noncontiguous_sink(), "must be contiguous"), + ], +) +def test_sparse_mla_sink_validation(sinks: torch.Tensor, match: str) -> None: + from vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseImpl, + ) + + with pytest.raises(ValueError, match=match): + ROCMAiterMLASparseImpl( + num_heads=4, + head_size=Q_HEAD_DIM, + scale=SM_SCALE, + num_kv_heads=1, + alibi_slopes=None, + sliding_window=None, + kv_cache_dtype="bfloat16", + logits_soft_cap=None, + attn_type="decoder", + kv_sharing_target_layer_name=None, + sinks=sinks, + kv_lora_rank=V_HEAD_DIM, + ) + + +def test_sparse_mla_backend_reports_sink_support_for_current_hardware() -> None: + from vllm.platforms.rocm import get_cdna_version + from vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseBackend, + ) + + assert ROCMAiterMLASparseBackend.supports_sink() == (get_cdna_version() in (3, 4)) + + +def test_sparse_mla_backend_rejects_dcp() -> None: + from vllm.platforms.rocm import RocmPlatform + from vllm.v1.attention.backends.registry import AttentionBackendEnum + from vllm.v1.attention.selector import AttentionSelectorConfig + + selector_config = AttentionSelectorConfig( + head_size=Q_HEAD_DIM, + dtype=torch.bfloat16, + kv_cache_dtype="bfloat16", + block_size=16, + use_mla=True, + has_sink=True, + use_sparse=True, + use_mm_prefix=False, + use_per_head_quant_scales=False, + attn_type="decoder", + use_dcp=True, + ) + + with pytest.raises(ValueError, match="DCP not supported"): + RocmPlatform.get_attn_backend_cls( + selected_backend=AttentionBackendEnum.ROCM_AITER_MLA_SPARSE, + attn_selector_config=selector_config, + ) + + +@pytest.mark.parametrize("num_heads,block_size", [(8, 16), (20, 32)]) +@pytest.mark.parametrize("interleaved_pages", [False, True]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@torch.inference_mode() +def test_sparse_mla_sink_matches_dense_attention_with_empty_rows_and_paged_cache( + num_heads: int, + block_size: int, + interleaved_pages: bool, + dtype: torch.dtype, +) -> None: + """Preserve head-specific sinks and latent values across ragged cache pages.""" + _require_aiter() + from vllm.v1.attention.backends.mla.rocm_aiter_mla import AiterMLAHelper + from vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseImpl, + ) + from vllm.v1.attention.backends.mla.sparse_utils import flat_kv_row_view + + head_dim, value_dim = 576, 512 + generator = torch.Generator().manual_seed(421) + # The empty row exercises sink-only attention; the last row crosses both + # page boundaries and the kernel's 16-key reduction tile. + selected_rows = [ + [], + [(2, 7)], + [(3, 1), (0, block_size - 1), (1, 0), (2, 4)], + [(i % 4, (i * 7) % block_size) for i in range(23)], + ] + q_cpu = torch.randn( + len(selected_rows), num_heads, head_dim, generator=generator + ).to(dtype) + cache_cpu = torch.randn(4, block_size, head_dim, generator=generator).to(dtype) + sinks_cpu = torch.linspace(-4.0, 6.0, num_heads) + + # Layers can share a backing allocation, leaving unused rows between pages. + num_layers = 3 if interleaved_pages else 1 + backing = torch.full( + (4, num_layers, block_size, head_dim), + float("nan"), + dtype=dtype, + device="cuda", + ) + cache = backing[:, num_layers - 1] + cache.copy_(cache_cpu) + indices = [ + page * num_layers * block_size + offset + for rows in selected_rows + for page, offset in rows + ] + lengths = torch.tensor([0, *(len(rows) for rows in selected_rows)]) + metadata = SimpleNamespace( + block_size=block_size, + num_prefills=0, + num_decodes=len(selected_rows), + num_decode_tokens=len(selected_rows), + max_query_len=1, + qo_indptr=torch.arange( + len(selected_rows) + 1, dtype=torch.int32, device="cuda" + ), + paged_kv_last_page_len=torch.ones( + len(selected_rows), dtype=torch.int32, device="cuda" + ), + work_meta_data=None, + attn_out_dtype=dtype, + paged_kv_indices=torch.tensor(indices, dtype=torch.int32, device="cuda"), + paged_kv_indptr=lengths.cumsum(0).to(device="cuda", dtype=torch.int32), + ) + impl = ROCMAiterMLASparseImpl.__new__(ROCMAiterMLASparseImpl) + impl.num_heads = num_heads + impl.head_size = head_dim + impl.kv_lora_rank = value_dim + impl.scale = head_dim**-0.5 + impl.kv_cache_dtype = "auto" + impl.sinks = sinks_cpu.cuda() + padded_q = AiterMLAHelper.get_mla_padded_q(num_heads, q_cpu.cuda()) + + kv_rows, _ = flat_kv_row_view(cache, block_size) + actual, _ = impl._forward_mla( + SimpleNamespace(_q_scale=None, _k_scale=None), + padded_q, + kv_rows.unsqueeze(1), + metadata, + ) + + references = [] + for query_idx, rows in enumerate(selected_rows): + keys = ( + torch.stack([cache_cpu[page, offset] for page, offset in rows]) + if rows + else cache_cpu.new_empty((0, head_dim)) + ) + references.append( + _sink_reference(q_cpu[query_idx], keys, sinks_cpu, impl.scale) + ) + expected, _, value_scale, score_error = ( + torch.stack(values) for values in zip(*references) + ) + assert actual.shape == (len(selected_rows), num_heads, value_dim) + assert actual.dtype == dtype + _assert_sink_output_close( + actual, + expected, + value_scale, + score_error, + dtype, + native=dtype == torch.bfloat16, + ) + + +@pytest.mark.parametrize("layout", ["LBNHC", "LBHNC", "BLNHC"]) +def test_sparse_mla_backend_resolves_only_contiguous_layer_layouts(monkeypatch, layout): + from vllm.config import CacheConfig + from vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseBackend, + ) + from vllm.v1.attention.backends.utils import resolve_kv_cache_layout + + monkeypatch.setenv("VLLM_KV_CACHE_LAYOUT", layout) + config = SimpleNamespace(cache_config=CacheConfig()) + supported = [ + [x.name for x in ROCMAiterMLASparseBackend.supported_kv_cache_layouts()] + ] + if layout == "BLNHC": + with pytest.raises(ValueError, match="does not satisfy"): + resolve_kv_cache_layout(config, supported) + else: + assert resolve_kv_cache_layout(config, supported).name == layout + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@torch.inference_mode() +def test_sparse_mla_sink_forward_mqa_preserves_split_query(dtype): + """The public sparse forward joins latent/RoPE queries in the model dtype.""" + _require_aiter() + from vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseImpl, + ) + + set_random_seed(412) + num_tokens, num_heads, block_size = 2, 8, 16 + q = torch.randn(num_tokens, num_heads, Q_HEAD_DIM, device="cuda").to(dtype) + kv = torch.randn(2 * block_size, 1, Q_HEAD_DIM, device="cuda").to(dtype) + selected = torch.tensor([[1, 7, 18], [0, 3, 20]], dtype=torch.int32, device="cuda") + sinks = torch.linspace(-3.0, 5.0, num_heads, device="cuda") + impl = object.__new__(ROCMAiterMLASparseImpl) + impl.num_heads = num_heads + impl.kv_lora_rank = V_HEAD_DIM + impl.kv_cache_dtype = "auto" + impl.scale = SM_SCALE + impl.sinks = sinks + impl.topk_indices_buffer = torch.full( + (num_tokens, 128), -1, dtype=torch.int32, device="cuda" + ) + impl.topk_indices_buffer[:, : selected.shape[1]] = selected + impl.q_concat_buffer = torch.empty_like(q) + metadata = SimpleNamespace( + attn_out_dtype=dtype, + num_actual_tokens=num_tokens, + num_prefills=0, + num_decodes=num_tokens, + num_decode_tokens=num_tokens, + max_query_len=1, + block_size=block_size, + topk_tokens=impl.topk_indices_buffer.shape[1], + req_id_per_token=torch.zeros(num_tokens, dtype=torch.int32, device="cuda"), + block_table=torch.tensor([[0, 1]], dtype=torch.int32, device="cuda"), + qo_indptr=torch.arange(num_tokens + 1, dtype=torch.int32, device="cuda"), + paged_kv_indptr=torch.tensor([0, 3, 6], dtype=torch.int32, device="cuda"), + paged_kv_indices=torch.empty( + selected.numel(), dtype=torch.int32, device="cuda" + ), + paged_kv_last_page_len=torch.ones(num_tokens, dtype=torch.int32, device="cuda"), + work_meta_data=None, + ) + actual, _ = impl.forward_mqa( + (q[..., :V_HEAD_DIM], q[..., V_HEAD_DIM:]), + kv, + metadata, + SimpleNamespace(_q_scale=None, _k_scale=None), + ) + references = [ + _sink_reference(q[i], kv[:, 0][selected[i].long()], sinks, SM_SCALE) + for i in range(num_tokens) + ] + expected, _, value_scale, score_error = ( + torch.stack(values) for values in zip(*references) + ) + assert actual.dtype == dtype + _assert_sink_output_close( + actual, + expected, + value_scale, + score_error, + dtype, + native=dtype == torch.bfloat16, + ) diff --git a/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py b/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py index d2f5251c62de..99da0af7d57b 100644 --- a/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py +++ b/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py @@ -59,6 +59,7 @@ def _make_builder(): builder.paged_kv_indptr = torch.zeros( max_num_batched_tokens + 1, dtype=torch.int32, device="cpu" ) + builder._use_persistent_metadata = True builder._num_attention_heads = 16 builder._num_compute_units = current_platform.num_compute_units() builder._mla_work_meta_data = torch.empty(1, dtype=torch.int32, device="cpu") @@ -130,6 +131,7 @@ def fake_generate_sparse_seqlen_triton( synchronize=lambda: events.append("sync") if events is not None else None ), ) + return fake_aiter def test_build_populates_decode_only_split_fields(monkeypatch): @@ -165,6 +167,25 @@ def test_build_populates_mixed_split_fields(monkeypatch): assert md.prefill is None +def test_sink_build_skips_persistent_metadata(monkeypatch): + builder = _make_builder() + builder._use_persistent_metadata = False + fake_aiter = _patch_build_deps(monkeypatch) + + md = builder.build( + common_prefix_len=0, common_attn_metadata=_make_common_metadata() + ) + + fake_aiter.get_mla_metadata_v1.assert_not_called() + assert md.work_meta_data is None + assert md.work_indptr is None + assert md.work_info_set is None + assert md.reduce_indptr is None + assert md.reduce_final_map is None + assert md.reduce_partial_map is None + assert builder._prev_metadata_key is None + + def test_sparse_persistent_metadata_syncs_only_after_recompute(monkeypatch): builder = _make_builder() common_metadata = _make_common_metadata() diff --git a/tests/models/test_initialization.py b/tests/models/test_initialization.py index 4100d0c8cf93..3b68f424788f 100644 --- a/tests/models/test_initialization.py +++ b/tests/models/test_initialization.py @@ -210,6 +210,12 @@ def test_can_initialize_large_subset(model_arch: str, monkeypatch: pytest.Monkey This test covers the complement of the tests covered in the "small subset" test. """ + if model_arch in ("HYV4ForCausalLM", "HYV4MTPModel"): + from vllm.platforms import current_platform + + if current_platform.is_rocm(): + pytest.skip("HY V4 ROCm initialization requires #54405") + can_initialize(model_arch, monkeypatch, HF_EXAMPLE_MODELS) diff --git a/tests/v1/attention/test_rocm_glm5next_sparse.py b/tests/v1/attention/test_rocm_glm5next_sparse.py index 35b51c855790..509d8f1faea4 100644 --- a/tests/v1/attention/test_rocm_glm5next_sparse.py +++ b/tests/v1/attention/test_rocm_glm5next_sparse.py @@ -1,11 +1,14 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from types import SimpleNamespace + import pytest import torch from vllm.platforms import current_platform from vllm.triton_utils import tl, triton +from vllm.v1.attention.backends.mla import rocm_aiter_mla_sparse as sparse_mod from vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( _use_rocm_sparse_triton, fit_kpool_indices_to_aiter, @@ -99,6 +102,55 @@ def test_rocm_sparse_triton_route( ) +@pytest.mark.parametrize("num_heads", [8, 12]) +def test_rocm_sparse_triton_route_preserves_padded_sinks(monkeypatch, num_heads): + captured = {} + + def fake_rocm_sparse_attn_prefill(**kwargs): + output = kwargs["output"] + captured["attn_sink"] = kwargs["attn_sink"] + output.copy_( + captured["attn_sink"].to(output.dtype).view(1, -1, 1).expand_as(output) + ) + + monkeypatch.setattr( + sparse_mod, "rocm_sparse_attn_prefill", fake_rocm_sparse_attn_prefill + ) + + impl = object.__new__(sparse_mod.ROCMAiterMLASparseImpl) + impl.num_heads = num_heads + impl.kv_lora_rank = 512 + impl.kv_cache_dtype = "auto" + impl.scale = 512**-0.5 + impl.sinks = torch.arange(num_heads, dtype=torch.float32) + + q = torch.zeros(2, 16, 512, dtype=torch.bfloat16) + kv = torch.zeros(4, 1, 512, dtype=torch.bfloat16) + metadata = SimpleNamespace( + attn_out_dtype=torch.bfloat16, + num_prefills=1, + num_decodes=0, + num_decode_tokens=0, + max_query_len=2, + paged_kv_indices=torch.empty(0, dtype=torch.int32), + paged_kv_indptr=torch.zeros(3, dtype=torch.int32), + ) + + output, lse = impl._forward_mla(SimpleNamespace(), q, kv, metadata) + + if num_heads == 8: + expected_sinks = impl.sinks.repeat_interleave(2) + else: + expected_sinks = torch.cat((impl.sinks, impl.sinks[:4])) + torch.testing.assert_close(captured["attn_sink"], expected_sinks) + assert output.shape == (2, num_heads, 512) + torch.testing.assert_close( + output[:, :, 0].float(), + impl.sinks.expand(2, -1), + ) + assert lse is None + + def test_rocm_sparse_attention_accepts_glm_nope_dimensions(): _validate_sparse_dims(512, 512, 0, "test") diff --git a/vllm/_aiter_ops.py b/vllm/_aiter_ops.py index 1b9bf1573585..2055fad22f0b 100644 --- a/vllm/_aiter_ops.py +++ b/vllm/_aiter_ops.py @@ -624,6 +624,110 @@ def _rocm_aiter_mla_decode_fwd_impl( ) +def _rocm_aiter_mla_decode_fwd_lse_impl( + q: torch.Tensor, + kv_buffer: torch.Tensor, + o: torch.Tensor, + qo_indptr: torch.Tensor, + max_seqlen_qo: int, + kv_indptr: torch.Tensor | None = None, + kv_indices: torch.Tensor | None = None, + kv_last_page_lens: torch.Tensor | None = None, + sm_scale: float = 1.0, + logit_cap: float = 0.0, + q_scale: torch.Tensor | None = None, + kv_scale: torch.Tensor | None = None, +) -> torch.Tensor: + """Run non-persistent AITER MLA decode and return natural-log LSE. + + gfx942 persistent MLA kernels do not provide an LSE code object. Keeping + this as a separate op makes that constraint structural: callers that need + LSE cannot accidentally forward persistent work metadata. + """ + from aiter.mla import get_meta_param, mla_decode_fwd + + kwargs: dict[str, float | int | torch.Tensor | None | bool] = { + "sm_scale": sm_scale, + "logit_cap": logit_cap, + "return_lse": True, + } + if _check_aiter_mla_fp8_support(): + kwargs["q_scale"] = q_scale + kwargs["kv_scale"] = kv_scale + + if ( + q.dtype == torch.bfloat16 + and kv_buffer.dtype == torch.bfloat16 + and q.shape[1] == 32 + ): + from vllm.platforms.rocm import on_gfx950 + + if on_gfx950(): + # The gfx950 H32 BF16 kernel mishandles ragged tails with split KV. + # Its single-split path returns correct output and natural-log LSE. + kwargs["num_kv_splits"] = 1 + kwargs["num_kv_splits_indptr"] = torch.arange( + qo_indptr.shape[0], dtype=torch.int32, device=q.device + ) + + if q.dtype == FP8_DTYPE: + assert kv_indices is not None + num_kv_splits, num_kv_splits_indptr = get_meta_param( + None, + qo_indptr.shape[0] - 1, + kv_indices.numel(), + q.shape[1], + max_seqlen_qo, + q.dtype, + ) + if num_kv_splits == 1: + # gfx942's one-split FP8 asm writes the final output directly but + # does not write either LSE buffer. Force the normal split reducer, + # which produces both the same output and an accurate natural LSE. + num_kv_splits = 2 + num_kv_splits_indptr = torch.arange( + 0, + (qo_indptr.shape[0]) * num_kv_splits, + num_kv_splits, + dtype=torch.int32, + device=q.device, + ) + kwargs["num_kv_splits"] = num_kv_splits + kwargs["num_kv_splits_indptr"] = num_kv_splits_indptr + + _, final_lse = mla_decode_fwd( + q, + kv_buffer.view(-1, 1, 1, q.shape[-1]), + o, + qo_indptr, + kv_indptr, + kv_indices, + kv_last_page_lens, + max_seqlen_qo, + **kwargs, + ) + if final_lse is None: + raise RuntimeError("AITER MLA decode did not return the requested LSE") + return final_lse + + +def _rocm_aiter_mla_decode_fwd_lse_fake( + q: torch.Tensor, + kv_buffer: torch.Tensor, + o: torch.Tensor, + qo_indptr: torch.Tensor, + max_seqlen_qo: int, + kv_indptr: torch.Tensor | None = None, + kv_indices: torch.Tensor | None = None, + kv_last_page_lens: torch.Tensor | None = None, + sm_scale: float = 1.0, + logit_cap: float = 0.0, + q_scale: torch.Tensor | None = None, + kv_scale: torch.Tensor | None = None, +) -> torch.Tensor: + return torch.empty(q.shape[:-1], dtype=torch.float32, device=q.device) + + def _rocm_aiter_w8a8_gemm_impl( A: torch.Tensor, B: torch.Tensor, @@ -2170,6 +2274,13 @@ def register_ops_once() -> None: mutates_args=["o"], ) + direct_register_custom_op( + op_name="rocm_aiter_mla_decode_fwd_lse", + op_func=_rocm_aiter_mla_decode_fwd_lse_impl, + mutates_args=["o"], + fake_impl=_rocm_aiter_mla_decode_fwd_lse_fake, + ) + direct_register_custom_op( op_name="rocm_aiter_w8a8_gemm", op_func=_rocm_aiter_w8a8_gemm_impl, @@ -2756,6 +2867,37 @@ def mla_decode_fwd( reduce_partial_map=reduce_partial_map, ) + @staticmethod + def mla_decode_fwd_lse( + q: torch.Tensor, + kv_buffer: torch.Tensor, + o: torch.Tensor, + sm_scale: float, + qo_indptr: torch.Tensor, + max_seqlen_qo: int, + kv_indptr: torch.Tensor | None = None, + kv_indices: torch.Tensor | None = None, + kv_last_page_lens: torch.Tensor | None = None, + logit_cap: float = 0.0, + q_scale: torch.Tensor | None = None, + kv_scale: torch.Tensor | None = None, + ) -> torch.Tensor: + """Run MLA decode without persistent metadata and return its LSE.""" + return torch.ops.vllm.rocm_aiter_mla_decode_fwd_lse( + q, + kv_buffer.view(-1, 1, 1, q.shape[-1]), + o, + qo_indptr, + max_seqlen_qo, + kv_indptr, + kv_indices, + kv_last_page_lens, + sm_scale=sm_scale, + logit_cap=logit_cap, + q_scale=q_scale, + kv_scale=kv_scale, + ) + @staticmethod def per_tensor_quant( x: torch.Tensor, diff --git a/vllm/v1/attention/backend.py b/vllm/v1/attention/backend.py index 2d545a8d8a7b..f6ff83206ebb 100644 --- a/vllm/v1/attention/backend.py +++ b/vllm/v1/attention/backend.py @@ -226,6 +226,13 @@ def supports_pcp(cls) -> bool: except NotImplementedError: return False + @classmethod + def supports_dcp(cls) -> bool: + try: + return cls.get_impl_cls().supports_dcp + except NotImplementedError: + return False + @classmethod def supports_non_causal_dcp(cls) -> bool: builder_cls = cls.get_builder_cls() @@ -324,6 +331,8 @@ def validate_configuration( invalid_reasons.append("KV connector not supported") if use_pcp and not cls.supports_pcp(): invalid_reasons.append("PCP not supported") + if use_dcp and not cls.supports_dcp(): + invalid_reasons.append("DCP not supported") if ( use_adaptive_verification and not cls.supports_device_cpu_query_lens_mismatch() diff --git a/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py b/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py index ef707d6ee1ae..3a9a4c92b350 100644 --- a/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py @@ -35,7 +35,7 @@ from vllm.v1.attention.ops.rocm_aiter_mla_sparse import ( rocm_sparse_attn_prefill, ) -from vllm.v1.kv_cache_interface import AttentionSpec +from vllm.v1.kv_cache_interface import AttentionSpec, KVCacheLayout from vllm.v1.worker.workspace import current_workspace_manager if TYPE_CHECKING: @@ -350,6 +350,17 @@ def is_mla(cls) -> bool: def is_sparse(cls) -> bool: return True + @classmethod + def supported_kv_cache_layouts(cls) -> tuple[KVCacheLayout, ...]: + # Global index conversion assumes contiguous pages within each layer. + return (KVCacheLayout.LBNHC, KVCacheLayout.LBHNC) + + @classmethod + def supports_sink(cls) -> bool: + from vllm.platforms.rocm import on_mi3xx + + return on_mi3xx() + @dataclass class ROCMAiterMLASparseMetadata(AttentionMetadata): @@ -418,6 +429,14 @@ def __init__( self.num_heads = self.model_config.get_num_attention_heads(parallel_config) self.mla_dims = get_mla_dims(self.model_config) self.topk_tokens = vllm_config.model_config.hf_config.index_topk + attention_context = vllm_config.compilation_config.static_forward_context + # Sink decode must use AITER's nonpersistent path. In particular, + # gfx942 has no persistent+LSE kernel, and its metadata heuristic + # terminates for HY-V4's TP1 H64 shape. + self._use_persistent_metadata = all( + getattr(attention_context[name].impl, "sinks", None) is None + for name in layer_names + ) # Bounds the KV-split heuristic (see `_sparse_decode_max_split`). self._num_compute_units = current_platform.num_compute_units() self.max_model_len_tensor = torch.tensor( @@ -591,14 +610,13 @@ def build( paged_kv_indices = self.paged_kv_indices[: num_tokens * self.topk_tokens] # ----- Compute persistent MLA metadata ----- - # The aiter sparse decode kernel uses qseqlen=1 (each query token is - # treated as its own batch entry), so persistent metadata can always - # be precomputed here. The kernel switches to the persistent - # work-stealing path automatically when work_meta_data is non-None. - # The output is a deterministic function of the per-request query and - # context lengths (both clamped to topk_tokens, past which per-token KV - # length saturates) and num_heads; fingerprint those CPU-side and skip - # the launch when nothing changed. + # The AITER sparse decode kernel uses qseqlen=1 (each query token is + # treated as its own batch entry). Build its persistent work metadata + # only when AITER is selected and no layer needs the nonpersistent LSE + # path for attention sinks. The output is a deterministic function of + # the per-request query and context lengths (both clamped to + # topk_tokens, past which per-token KV length saturates) and num_heads; + # fingerprint those CPU-side and skip the launch when nothing changed. head_size = self.mla_dims.kv_lora_rank + self.mla_dims.qk_rope_head_dim use_triton_sparse = _use_rocm_sparse_triton( kv_cache_dtype=self.kv_cache_dtype, @@ -615,7 +633,7 @@ def build( reduce_indptr = None reduce_final_map = None reduce_partial_map = None - if not use_triton_sparse: + if self._use_persistent_metadata and not use_triton_sparse: num_reqs = common_attn_metadata.num_reqs clamped_seq_lens = np.minimum( common_attn_metadata.seq_lens_cpu[:num_reqs].numpy(), @@ -729,6 +747,7 @@ def log2sumexp2(a: torch.Tensor, dim: int) -> torch.Tensor: class ROCMAiterMLASparseImpl(MLAAttentionImpl[ROCMAiterMLASparseMetadata]): is_sparse = True supports_dense_mha_prefill = False + supports_dcp = False def __init__( self, @@ -753,6 +772,20 @@ def __init__( self.head_size = head_size self.scale = float(scale) self.num_kv_heads = num_kv_heads + sinks = mla_args.pop("sinks", None) + if sinks is not None: + if sinks.dtype != torch.float32: + raise ValueError( + f"ROCm AITER MLA sinks must be float32, got {sinks.dtype}" + ) + if sinks.ndim != 1 or sinks.numel() != num_heads: + raise ValueError( + "ROCm AITER MLA sinks must have shape " + f"({num_heads},), got {tuple(sinks.shape)}" + ) + if not sinks.is_contiguous(): + raise ValueError("ROCm AITER MLA sinks must be contiguous") + self.sinks: torch.Tensor | None = sinks self.kv_cache_dtype = kv_cache_dtype self.kv_lora_rank: int = mla_args["kv_lora_rank"] self.softmax_scale = scale @@ -776,16 +809,31 @@ def _forward_mla( q: torch.Tensor, # [sq, heads, d_qk] kv_c_and_k_pe_cache: torch.Tensor, # [blocks, heads, d_qk] attn_metadata: ROCMAiterMLASparseMetadata, - ) -> torch.Tensor: + ) -> tuple[torch.Tensor, torch.Tensor | None]: num_tokens = q.shape[0] - mla_num_heads = AiterMLAHelper.get_actual_mla_num_heads(self.num_heads) - output = torch.empty( - [num_tokens, mla_num_heads, self.kv_lora_rank], - dtype=attn_metadata.attn_out_dtype, - device=q.device, + base_mla_num_heads = AiterMLAHelper.get_actual_mla_num_heads(self.num_heads) + mla_num_heads = base_mla_num_heads + need_lse = self.sinks is not None + from vllm.platforms.rocm import on_gfx942 + + # Keep sink attention available for dtypes/head shapes without an + # AITER return-LSE kernel. Sink layers never need persistent metadata. + triton_sink_fallback = ( + need_lse + and q.dtype == kv_c_and_k_pe_cache.dtype + and ( + q.dtype == torch.float16 + or ( + q.dtype == torch.bfloat16 + and ( + mla_num_heads > 128 + or (on_gfx942() and 32 < mla_num_heads <= 64) + ) + ) + ) ) - if _use_rocm_sparse_triton( + if triton_sink_fallback or _use_rocm_sparse_triton( kv_cache_dtype=self.kv_cache_dtype, head_size=q.shape[-1], kv_lora_rank=self.kv_lora_rank, @@ -794,6 +842,18 @@ def _forward_mla( num_decode_tokens=attn_metadata.num_decode_tokens, max_query_len=attn_metadata.max_query_len, ): + output = torch.empty( + [num_tokens, q.shape[1], self.kv_lora_rank], + dtype=attn_metadata.attn_out_dtype, + device=q.device, + ) + triton_sinks = None + if self.sinks is not None: + triton_sinks = AiterMLAHelper.get_mla_padded_q( + self.num_heads, + self.sinks.reshape(1, self.num_heads, 1), + q.shape[1], + ).reshape(-1) rocm_sparse_attn_prefill( q=q, kv=kv_c_and_k_pe_cache.view(-1, 1, q.shape[-1]), @@ -803,44 +863,149 @@ def _forward_mla( head_dim=q.shape[-1], nope_head_dim=self.kv_lora_rank, rope_head_dim=q.shape[-1] - self.kv_lora_rank, - attn_sink=None, + attn_sink=triton_sinks, output=output, ragged_indices=attn_metadata.paged_kv_indices, ragged_indptr=attn_metadata.paged_kv_indptr, ) - return AiterMLAHelper.get_mla_unpadded_o(self.num_heads, output) - - # Build kwargs and forward the persistent MLA metadata when it has - # been computed. The aiter mla_decode_fwd switches to its - # work-stealing persistent kernel path when work_meta_data is given. - mla_kwargs: dict = dict( - q_scale=layer._q_scale, - kv_scale=layer._k_scale, - ) - if attn_metadata.work_meta_data is not None: - mla_kwargs.update( - work_meta_data=attn_metadata.work_meta_data, - work_indptr=attn_metadata.work_indptr, - work_info_set=attn_metadata.work_info_set, - reduce_indptr=attn_metadata.reduce_indptr, - reduce_final_map=attn_metadata.reduce_final_map, - reduce_partial_map=attn_metadata.reduce_partial_map, - ) + output = AiterMLAHelper.get_mla_unpadded_o(self.num_heads, output) + return output, None + + # AITER's nonpersistent return-LSE dispatch has discrete head kernels. + supported_head_buckets: tuple[int, ...] | None = None + head_dtype_name = "" + if need_lse: + if ( + q.dtype == torch.bfloat16 + and kv_c_and_k_pe_cache.dtype == torch.bfloat16 + ): + supported_head_buckets = (16, 32, 64, 128) + head_dtype_name = "BF16" + elif ( + q.dtype == current_platform.fp8_dtype() + and kv_c_and_k_pe_cache.dtype == current_platform.fp8_dtype() + ): + supported_head_buckets = (16, 128) + head_dtype_name = "FP8" + else: + raise ValueError( + "ROCm AITER MLA attention sinks require query and KV to " + "both use BF16 or both use FP8, got " + f"query={q.dtype}, KV={kv_c_and_k_pe_cache.dtype}" + ) - rocm_aiter_ops.mla_decode_fwd( - q, - kv_c_and_k_pe_cache, - output, - self.scale, - attn_metadata.qo_indptr, - 1, - attn_metadata.paged_kv_indptr, - attn_metadata.paged_kv_indices, - attn_metadata.paged_kv_last_page_len, - **mla_kwargs, + if supported_head_buckets is not None: + from vllm.platforms.rocm import on_mi3xx + + if on_mi3xx(): + supported_heads = next( + ( + heads + for heads in supported_head_buckets + if heads >= mla_num_heads + ), + None, + ) + if supported_heads is None: + raise ValueError( + "ROCm AITER MLA attention sinks support at most 128 " + f"padded local {head_dtype_name} heads; increase " + "tensor_parallel_size" + ) + if supported_heads != mla_num_heads: + q = AiterMLAHelper.get_mla_padded_q( + mla_num_heads, q, supported_heads + ) + mla_num_heads = supported_heads + output = torch.empty( + [num_tokens, mla_num_heads, self.kv_lora_rank], + dtype=attn_metadata.attn_out_dtype, + device=q.device, ) - return AiterMLAHelper.get_mla_unpadded_o(self.num_heads, output) + if need_lse: + # gfx942 has no persistent MLA code object that writes final LSE. + # The split-KV path consumes the same ragged indices and is exact. + lse = rocm_aiter_ops.mla_decode_fwd_lse( + q, + kv_c_and_k_pe_cache, + output, + self.scale, + attn_metadata.qo_indptr, + 1, + attn_metadata.paged_kv_indptr, + attn_metadata.paged_kv_indices, + attn_metadata.paged_kv_last_page_len, + q_scale=layer._q_scale, + kv_scale=layer._k_scale, + ) + else: + # Preserve the persistent work-stealing fast path for models that + # do not use sinks. + mla_kwargs: dict = dict( + q_scale=layer._q_scale, + kv_scale=layer._k_scale, + ) + if attn_metadata.work_meta_data is not None: + mla_kwargs.update( + work_meta_data=attn_metadata.work_meta_data, + work_indptr=attn_metadata.work_indptr, + work_info_set=attn_metadata.work_info_set, + reduce_indptr=attn_metadata.reduce_indptr, + reduce_final_map=attn_metadata.reduce_final_map, + reduce_partial_map=attn_metadata.reduce_partial_map, + ) + + rocm_aiter_ops.mla_decode_fwd( + q, + kv_c_and_k_pe_cache, + output, + self.scale, + attn_metadata.qo_indptr, + 1, + attn_metadata.paged_kv_indptr, + attn_metadata.paged_kv_indices, + attn_metadata.paged_kv_last_page_len, + **mla_kwargs, + ) + lse = None + + if mla_num_heads != base_mla_num_heads: + if mla_num_heads % base_mla_num_heads == 0: + head_stride = mla_num_heads // base_mla_num_heads + output = output[:, ::head_stride] + if lse is not None: + lse = lse[:, ::head_stride] + else: + output = output[:, :base_mla_num_heads] + if lse is not None: + lse = lse[:, :base_mla_num_heads] + + output = AiterMLAHelper.get_mla_unpadded_o(self.num_heads, output) + if lse is not None: + lse = AiterMLAHelper.get_mla_unpadded_o( + self.num_heads, lse.unsqueeze(-1) + ).squeeze(-1) + + if self.sinks is not None: + assert lse is not None + # Empty ragged rows have only sink mass and no value contribution. + # AITER can return NaN output/LSE for those rows; do not multiply it + # by a zero normalization factor and propagate the NaN. + has_keys = ( + attn_metadata.paged_kv_indptr[1:] > attn_metadata.paged_kv_indptr[:-1] + ).unsqueeze(-1) + lse = torch.where(has_keys, lse, float("-inf")) + sink_lse = torch.logaddexp(lse, self.sinks) + sink_scale = torch.exp(lse - sink_lse) + output = torch.where( + has_keys.unsqueeze(-1), + output.float() * sink_scale.unsqueeze(-1), + 0.0, + ).to(output.dtype) + lse = sink_lse + + return output, lse def forward_mqa( self, @@ -863,6 +1028,8 @@ def forward_mqa( q = self.q_concat_buffer[: ql_nope.shape[0]] if q_pe.shape[-1] == 0: q.copy_(ql_nope) + elif q.dtype == torch.float16: + torch.cat((ql_nope, q_pe), dim=-1, out=q) else: ops.concat_mla_q(ql_nope, q_pe, q) @@ -892,7 +1059,6 @@ def forward_mqa( q, _ = ops.scaled_fp8_quant(q.view(q.shape[0], -1), layer._q_scale) q = q.view(original_q_shape) mla_padded_q = AiterMLAHelper.get_mla_padded_q(self.num_heads, q) - attn_out = self._forward_mla( + return self._forward_mla( layer, mla_padded_q, kv_c_and_k_pe_cache, attn_metadata ) - return attn_out, None diff --git a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py index 33f6386e346e..58bef917a2f8 100644 --- a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -2610,7 +2610,7 @@ def _rocm_sparse_attn_prefill_ragged_triton( block_d = triton.next_power_of_2(head_dim) block_k = 16 if head_dim >= 256 else 32 num_warps = 4 - out = torch.empty_like(q, dtype=torch.bfloat16) + out = torch.empty_like(q) _sparse_attn_prefill_ragged_kernel[(num_queries, triton.cdiv(num_heads, block_h))]( q, kv, @@ -3195,7 +3195,7 @@ def rocm_sparse_attn_prefill( rope_head_dim=rope_head_dim, topk_length=topk_length, ) - output.copy_(output_chunk.to(output.dtype)) + output.copy_(output_chunk[..., : output.shape[-1]].to(output.dtype)) def rocm_sparse_attn_decode(