From 03459dee261fb5894c02916ed816dd435236ab98 Mon Sep 17 00:00:00 2001 From: zhangxjohn Date: Fri, 29 May 2026 09:53:53 +0000 Subject: [PATCH] [Bugfix] Align MLA indexer block table with MTP speculative decode Fix GLM-5.1 MTP=4 crashes under concurrent load: - Pad expanded_block_table_buffer columns like MultiGroupBlockTable to avoid expand size mismatch (e.g. 1669 vs 1670) when flattening MTP decode tokens into the indexer path. - Drop padded MTP decode slots so seq_lens and block_table row counts stay consistent in the flatten path. - Normalize DeepGEMM context_lens to 2D contiguous tensors and flatten 2D seq_lens for persistent_topk in sparse_attn_indexer. Fixes #41094 (DeepGEMM contiguous assertion) together with the block table alignment that caused EngineDeadError on high-concurrency MTP. Co-authored-by: Cursor --- .../test_indexer_block_table_padding.py | 43 +++++++++++++++++++ .../layers/sparse_attn_indexer.py | 11 ++++- vllm/utils/deep_gemm.py | 3 ++ vllm/v1/attention/backends/mla/indexer.py | 21 +++++++-- 4 files changed, 73 insertions(+), 5 deletions(-) create mode 100644 tests/v1/attention/test_indexer_block_table_padding.py diff --git a/tests/v1/attention/test_indexer_block_table_padding.py b/tests/v1/attention/test_indexer_block_table_padding.py new file mode 100644 index 000000000000..1e7e18ef642a --- /dev/null +++ b/tests/v1/attention/test_indexer_block_table_padding.py @@ -0,0 +1,43 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import torch + +from tests.v1.attention.utils import create_vllm_config +from vllm.v1.attention.backends.mla.indexer import DeepseekV32IndexerMetadataBuilder +from vllm.v1.kv_cache_interface import MLAAttentionSpec + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_indexer_expanded_block_table_matches_multigroup_block_table_padding(): + """Regression: expanded_block_table_buffer width must match block_table_tensor. + + MultiGroupBlockTable pads max_num_blocks_per_req to a multiple of + (128 // block_size). Without the same padding in the MLA indexer buffer, + MTP flatten decode can fail with expand size mismatch (e.g. 1669 vs 1670). + """ + device = torch.device("cuda") + max_model_len = 106816 + block_size = 64 + # ceil(106816 / 64) = 1669 -> padded to 1670 for block_size=64 + expected_blocks_per_req = 1670 + + kv_cache_spec = MLAAttentionSpec( + block_size=block_size, + num_kv_heads=1, + head_size=128, + dtype=torch.bfloat16, + ) + vllm_config = create_vllm_config( + max_model_len=max_model_len, + block_size=block_size, + ) + builder = DeepseekV32IndexerMetadataBuilder( + kv_cache_spec=kv_cache_spec, + layer_names=["dummy"], + vllm_config=vllm_config, + device=device, + ) + + assert builder.expanded_block_table_buffer.shape[1] == expected_blocks_per_req diff --git a/vllm/model_executor/layers/sparse_attn_indexer.py b/vllm/model_executor/layers/sparse_attn_indexer.py index 9597708b62e7..f72835a7f7b5 100644 --- a/vllm/model_executor/layers/sparse_attn_indexer.py +++ b/vllm/model_executor/layers/sparse_attn_indexer.py @@ -297,6 +297,13 @@ def sparse_attn_indexer( next_n = padded_q_quant_decode_tokens.shape[1] num_padded_tokens = batch_size * next_n seq_lens = decode_metadata.seq_lens[:batch_size] + # persistent_topk expects 1D per-row context lengths in row-major order. + if seq_lens.ndim == 2 and seq_lens.shape[1] > 1: + seq_lens_for_topk = seq_lens.reshape(-1).contiguous() + else: + seq_lens_for_topk = seq_lens.contiguous() + if seq_lens_for_topk.ndim == 2: + seq_lens_for_topk = seq_lens_for_topk.squeeze(-1) # seq_lens is always 2D: (B, next_n) for native spec decode, (B, 1) # otherwise. deep_gemm fp8_fp4_paged_mqa_logits requires 2D context_lens; # the downstream topk kernels accept both 1D and 2D. @@ -341,7 +348,7 @@ def sparse_attn_indexer( ) torch.ops._C.persistent_topk( logits, - seq_lens, + seq_lens_for_topk, topk_indices, topk_workspace, topk_tokens, @@ -351,7 +358,7 @@ def sparse_attn_indexer( ops.top_k_per_row_decode( logits, next_n, - seq_lens, + seq_lens_for_topk, topk_indices, num_rows, logits.stride(0), diff --git a/vllm/utils/deep_gemm.py b/vllm/utils/deep_gemm.py index 6b89f5c33203..fb503a4293c8 100644 --- a/vllm/utils/deep_gemm.py +++ b/vllm/utils/deep_gemm.py @@ -401,6 +401,9 @@ def get_paged_mqa_logits_metadata( _lazy_init() if _get_paged_mqa_logits_metadata_impl is None: return _missing() + if context_lens.dim() == 1: + context_lens = context_lens.unsqueeze(-1) + context_lens = context_lens.contiguous() return _get_paged_mqa_logits_metadata_impl(context_lens, block_size, num_sms) diff --git a/vllm/v1/attention/backends/mla/indexer.py b/vllm/v1/attention/backends/mla/indexer.py index 2870ec9a15c0..9f6d56b04cde 100644 --- a/vllm/v1/attention/backends/mla/indexer.py +++ b/vllm/v1/attention/backends/mla/indexer.py @@ -304,10 +304,19 @@ def __init__(self, *args, **kwargs): dtype=torch.int32, device=self.device, ) + block_size = self.kv_cache_spec.block_size max_num_blocks_per_req = cdiv( self.vllm_config.model_config.max_model_len, - self.kv_cache_spec.block_size * get_total_cp_world_size(), + block_size * get_total_cp_world_size(), ) + # Match MultiGroupBlockTable padding (block_table.py): without this, + # block_table_tensor has more columns than expanded_block_table_buffer + # and MTP flatten hits expand errors like [28, 1669] vs [28, 1670]. + if block_size <= 128: + max_num_blocks_per_req = ( + cdiv(max_num_blocks_per_req, 128 // block_size) + * (128 // block_size) + ) self.expanded_block_table_buffer = torch.zeros( ( scheduler_config.max_num_batched_tokens, @@ -420,7 +429,6 @@ def _prepare_decode_tensors( expanded_offsets + self.arange_buffer[:actual_expanded] + 1 ) self.decode_seq_lens_buffer[actual_expanded:] = 0 - seq_lens = self.decode_seq_lens_buffer[:num_decode_tokens] # Give each of the flattened entries the same block table row as the # original request. @@ -433,6 +441,9 @@ def _prepare_decode_tensors( self.expanded_block_table_buffer[ actual_expanded:num_decode_tokens, 0 ] = 0 + # Drop padded decode slots so seq_lens/logits row counts match. + num_decode_tokens = actual_expanded + seq_lens = self.decode_seq_lens_buffer[:num_decode_tokens] block_table = self.expanded_block_table_buffer[:num_decode_tokens] # All reqs now have decode_len=1 @@ -583,6 +594,9 @@ def build( max_decode_len=max_decode_len, ) ) + # Flatten path may drop padded MTP slots; keep metadata in sync. + if not use_native and batch_size != num_decode_tokens: + num_decode_tokens = batch_size # For DeepseekV4 (compress_ratio > 1), the indexer KV cache stores # compressed tokens. Convert uncompressed seq_lens to compressed. @@ -610,8 +624,9 @@ def build( # DeepGEMM is required for the paged MQA logits on CUDA devices if current_platform.is_cuda() and has_deep_gemm(): + metadata_context_lens = seq_lens.contiguous() self.scheduler_metadata_buffer[:] = get_paged_mqa_logits_metadata( - seq_lens, + metadata_context_lens, self.kv_cache_spec.storage_block_size, self.num_sms, )