From c4bb3ce2730607400bb79df5d44a2579a1d1f2b1 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Thu, 7 May 2026 17:06:28 -0700 Subject: [PATCH] Fix stuck when enabling MTP on DSA models (#24635) --- .../srt/layers/attention/nsa_backend.py | 59 ++++++++++++++----- .../8-gpu-models/test_dsa_models_mtp.py | 1 - .../cp/test_deepseek_v32_cp_single_node.py | 1 - .../quant/test_deepseek_v32_fp4_mtp_4gpu.py | 1 - 4 files changed, 44 insertions(+), 18 deletions(-) diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index f515228d4..cb9d9eab3 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -68,14 +68,14 @@ else: def _to_2d_context_lens(seqlens_32: torch.Tensor, batch_size: int) -> torch.Tensor: + # Always normalize to (N_total, 1) layout, to avoid deadlock at deep_gemm.fp8_paged_mqa_logits if seqlens_32.dim() == 2: - return seqlens_32 - n = seqlens_32.numel() - assert ( - n % batch_size == 0 - ), f"seqlens_32 size {n} is not a multiple of batch_size {batch_size}" - next_n = n // batch_size - return seqlens_32.view(batch_size, next_n) + if seqlens_32.size(1) == 1: + return seqlens_32 + # Fall through and re-flatten if the caller already gave us a (bs, next_n) + # view — we want (N_total, 1) regardless. + seqlens_32 = seqlens_32.reshape(-1) + return seqlens_32.contiguous().view(-1, 1) # Reuse this workspace buffer across all NSA backend instances @@ -644,7 +644,7 @@ class NativeSparseAttnBackend( if is_cuda() and ( forward_batch.forward_mode.is_decode_or_idle() or forward_batch.forward_mode.is_target_verify() - or forward_batch.forward_mode.is_draft_extend() + or forward_batch.forward_mode.is_draft_extend(include_v2=True) ): try: import deep_gemm @@ -654,7 +654,7 @@ class NativeSparseAttnBackend( seqlens_expanded if ( forward_batch.forward_mode.is_target_verify() - or forward_batch.forward_mode.is_draft_extend() + or forward_batch.forward_mode.is_draft_extend(include_v2=True) ) else cache_seqlens_int32 ) @@ -933,7 +933,7 @@ class NativeSparseAttnBackend( if is_cuda() and ( forward_mode.is_decode_or_idle() or forward_mode.is_target_verify() - or forward_mode.is_draft_extend() + or forward_mode.is_draft_extend(include_v2=True) ): try: import deep_gemm @@ -942,7 +942,7 @@ class NativeSparseAttnBackend( seqlens_expanded if ( forward_mode.is_target_verify() - or forward_mode.is_draft_extend() + or forward_mode.is_draft_extend(include_v2=True) ) else cache_seqlens_int32 ) @@ -1084,7 +1084,7 @@ class NativeSparseAttnBackend( if is_cuda() and ( forward_mode.is_decode_or_idle() or forward_mode.is_target_verify() - or forward_mode.is_draft_extend() + or forward_mode.is_draft_extend(include_v2=True) ): try: import deep_gemm @@ -1093,7 +1093,7 @@ class NativeSparseAttnBackend( seqlens_expanded if ( forward_mode.is_target_verify() - or forward_mode.is_draft_extend() + or forward_mode.is_draft_extend(include_v2=True) ) else metadata.cache_seqlens_int32 ) @@ -1102,11 +1102,13 @@ class NativeSparseAttnBackend( seqlens_32_2d, 64, deep_gemm.get_num_sms() ) if metadata.paged_mqa_schedule_metadata is None: - metadata.paged_mqa_schedule_metadata = new_schedule + object.__setattr__( + metadata, "paged_mqa_schedule_metadata", new_schedule + ) else: metadata.paged_mqa_schedule_metadata.copy_(new_schedule) except (ImportError, ModuleNotFoundError): - metadata.paged_mqa_schedule_metadata = None + object.__setattr__(metadata, "paged_mqa_schedule_metadata", None) seqlens_expanded_size = seqlens_expanded.shape[0] assert ( metadata.nsa_cache_seqlens_int32 is not None @@ -1293,6 +1295,33 @@ class NativeSparseAttnBackend( flashmla_metadata = metadata.flashmla_metadata.slice(slice(0, size + 1)) flashmla_metadata.copy_(precomputed.flashmla_metadata) + # Refresh DeepGEMM paged MQA schedule metadata for the actual seqlens of + # this replay (the captured graph holds stale data otherwise, which can + # deadlock the kernel when the runtime work decomposition diverges from + # the captured one). + if is_cuda(): + try: + import deep_gemm + + if forward_mode.is_decode_or_idle(): + seqlens_32 = metadata.cache_seqlens_int32 + else: + seqlens_32 = metadata.nsa_seqlens_expanded[ + : precomputed.seqlens_expanded_size + ] + seqlens_32_2d = _to_2d_context_lens(seqlens_32, bs) + new_schedule = deep_gemm.get_paged_mqa_logits_metadata( + seqlens_32_2d, 64, deep_gemm.get_num_sms() + ) + if metadata.paged_mqa_schedule_metadata is None: + object.__setattr__( + metadata, "paged_mqa_schedule_metadata", new_schedule + ) + else: + metadata.paged_mqa_schedule_metadata.copy_(new_schedule) + except (ImportError, ModuleNotFoundError): + pass + self.forward_metadata = metadata def forward_extend( diff --git a/test/registered/8-gpu-models/test_dsa_models_mtp.py b/test/registered/8-gpu-models/test_dsa_models_mtp.py index 9bb05f36c..bf7dcae03 100644 --- a/test/registered/8-gpu-models/test_dsa_models_mtp.py +++ b/test/registered/8-gpu-models/test_dsa_models_mtp.py @@ -19,7 +19,6 @@ from sglang.test.test_utils import ( register_cuda_ci( est_time=1048, suite="stage-c-test-8-gpu-h200", - disabled="Disabled due to #24268. Should be fixed soon.", ) FULL_DEEPSEEK_V32_MODEL_PATH = "deepseek-ai/DeepSeek-V3.2" diff --git a/test/registered/cp/test_deepseek_v32_cp_single_node.py b/test/registered/cp/test_deepseek_v32_cp_single_node.py index dcfcb4d78..8bb1c6328 100644 --- a/test/registered/cp/test_deepseek_v32_cp_single_node.py +++ b/test/registered/cp/test_deepseek_v32_cp_single_node.py @@ -16,7 +16,6 @@ from sglang.test.test_utils import ( register_cuda_ci( est_time=616, suite="stage-c-test-deepep-8-gpu-h200", - disabled="Disabled due to #24268. Should be fixed soon.", ) DEEPSEEK_V32_MODEL_PATH = "deepseek-ai/DeepSeek-V3.2" diff --git a/test/registered/quant/test_deepseek_v32_fp4_mtp_4gpu.py b/test/registered/quant/test_deepseek_v32_fp4_mtp_4gpu.py index 14793c17a..ca259b2ae 100644 --- a/test/registered/quant/test_deepseek_v32_fp4_mtp_4gpu.py +++ b/test/registered/quant/test_deepseek_v32_fp4_mtp_4gpu.py @@ -18,7 +18,6 @@ from sglang.test.test_utils import ( register_cuda_ci( est_time=1060, suite="stage-c-test-4-gpu-b200", - disabled="Disabled due to #24268. Should be fixed soon.", ) FULL_DEEPSEEK_V3_FP4_MODEL_PATH = "nvidia/DeepSeek-V3.2-NVFP4"