From 73b53e7a873821c1bfd1c1a56732d22e1357b28b Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Wed, 3 Jun 2026 01:29:52 -0700 Subject: [PATCH] Revert "Support NextN = 2/4 in DSV32" (#27138) --- .../srt/layers/attention/dsa/dsa_indexer.py | 50 +---- .../srt/layers/attention/dsa_backend.py | 178 ++++++++---------- 2 files changed, 92 insertions(+), 136 deletions(-) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 85fcd4b9e..318a147ba 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -602,31 +602,11 @@ class Indexer(MultiPlatformOp): # Reuse pre-computed schedule metadata if available (from init_forward_metadata), # otherwise fall back to computing it here. schedule_metadata = getattr(metadata, "paged_mqa_schedule_metadata", None) - - assert len(q_fp8.shape) == 3 - # attn_tp_size > 1 or MAX_LEN padding mode can leave padding in the - # hidden states; q_offset is the real (unpadded) q length. - q_offset = sum(metadata.get_dsa_extend_len_cpu()) - - # DG-native q=[B,next_n,H,D] is faster than expanded q=[B*next_n,1,H,D] - # for target_verify with next_n>=2 (bigger MMA tile, fewer atoms). The - # precomputed ctx_lens_2d's shape is the single source of truth — if - # dsa_backend chose the per-token layout (e.g. non-SM100), fall through - # to the expanded path. - B = metadata.get_seqlens_int32().shape[0] - next_n = q_offset // B if B > 0 else 0 - ctx_2d = getattr(metadata, "paged_mqa_ctx_lens_2d", None) - use_dg_native = ( - _is_cuda - and forward_batch.forward_mode.is_target_verify() - and next_n >= 2 - and ctx_2d is not None - and ctx_2d.shape == (B, next_n) - ) - - if use_dg_native: - seqlens_32_2d = ctx_2d - elif seqlens_32.dim() == 2: + # DeepGEMM release-0426 requires context_lens of shape [batch_size, next_n] + # to match q.shape = [batch_size, next_n, heads, head_dim]. The indexer uses + # next_n=1 with batch_size=N_total via q_fp8.unsqueeze(1) below, so mirror + # that layout here. + if seqlens_32.dim() == 2: seqlens_32_2d = seqlens_32 else: seqlens_32_2d = seqlens_32.unsqueeze(-1) @@ -636,6 +616,8 @@ class Indexer(MultiPlatformOp): seqlens_32_2d, blocksize, self.sm_count ) + assert len(q_fp8.shape) == 3 + q_fp8 = q_fp8.unsqueeze(1) # the next_n dim is 1 now assert len(kv_cache_fp8.shape) == 2 block_kv = page_size num_heads_kv = 1 @@ -646,10 +628,12 @@ class Indexer(MultiPlatformOp): assert len(weights.shape) == 3 weights = weights.squeeze(2) + # When attn_tp_size > 1 or in the MAX_LEN padding mode, padding may exist in the hidden states, + # and it is necessary to extract the actual q length. + q_offset = sum(metadata.get_dsa_extend_len_cpu()) if _is_hip: from aiter.ops.triton.pa_mqa_logits import deepgemm_fp8_paged_mqa_logits - q_fp8 = q_fp8.unsqueeze(1) batch_size, next_n, heads, _ = q_fp8.shape logits = torch.empty( (batch_size * next_n, max_seq_len), @@ -667,21 +651,7 @@ class Indexer(MultiPlatformOp): Preshuffle=_use_aiter_preshuffle, KVBlockSize=block_kv, ) - elif use_dg_native: - # block_tables[::next_n] de-expands dsa_backend's repeat_interleave - # without a copy (DG only checks `stride(1) == 1`). - logits = deep_gemm.fp8_paged_mqa_logits( - q_fp8[:q_offset].view(B, next_n, q_fp8.shape[1], q_fp8.shape[2]), - kv_cache_fp8, - weights[:q_offset], - seqlens_32_2d, - block_tables[::next_n], - schedule_metadata, - max_seq_len, - clean_logits=False, - ) else: - q_fp8 = q_fp8.unsqueeze(1) logits = deep_gemm.fp8_paged_mqa_logits( q_fp8[:q_offset], kv_cache_fp8, diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index eb2c50cb4..e3b8eebd5 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -50,10 +50,7 @@ from sglang.srt.layers.attention.utils import ( ) from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.utils import is_cuda, is_hip, is_sm100_supported - -if is_cuda(): - import deep_gemm +from sglang.srt.utils import is_cuda, is_hip if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -156,9 +153,6 @@ class DSAMetadata: # DeepGEMM schedule metadata for paged MQA logits (decode/target_verify/draft_extend only). # Precomputed once per forward batch and reused across layers. paged_mqa_schedule_metadata: Optional[torch.Tensor] = None - # 2D context_lens used to build the schedule above; the indexer reuses it - # as DG's `context_lens` arg so the broadcast doesn't rebuild per layer. - paged_mqa_ctx_lens_2d: Optional[torch.Tensor] = None # The sum of sequence lengths for key, prefill only seq_lens_sum: Optional[int] = None # The flattened 1D page table with shape (seq_lens_sum,), prefill only @@ -205,7 +199,6 @@ class DSAIndexerMetadata(BaseIndexerMetadata): topk_transform_method: TopkTransformMethod topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL paged_mqa_schedule_metadata: Optional[torch.Tensor] = None - paged_mqa_ctx_lens_2d: Optional[torch.Tensor] = None force_unfused_topk: bool = False def get_seqlens_int32(self) -> torch.Tensor: @@ -387,31 +380,6 @@ class DeepseekSparseAttnBackend( else: self.workspace_buffer = None - def _build_paged_mqa_schedule_2d_ctx_lens( - self, - forward_mode: ForwardMode, - cache_seqlens_int32: torch.Tensor, - seqlens_expanded: torch.Tensor, - batch_size: int, - ) -> torch.Tensor: - # target_verify with next_n>=2 uses DG-native q=[B,next_n,H,D] which - # needs a [B, next_n] schedule; everything else stays per-token. - # TODO: SM90 supports DG-native next_n in {1,2} too — enable once - # validated; for now DG-native is SM100+ only. - next_n = self.speculative_num_draft_tokens - if ( - forward_mode.is_target_verify() - and next_n - and next_n >= 2 - and is_sm100_supported() - ): - return cache_seqlens_int32.view(-1, 1).expand(-1, next_n).contiguous() - if forward_mode.is_target_verify() or forward_mode.is_draft_extend( - include_v2=True - ): - return _to_2d_context_lens(seqlens_expanded, batch_size) - return _to_2d_context_lens(cache_seqlens_int32, batch_size) - def _get_fused_topk_page_table(self, topk_indices: torch.Tensor) -> torch.Tensor: if ( self.dsa_topk_backend.is_sgl_kernel() @@ -686,24 +654,33 @@ class DeepseekSparseAttnBackend( dsa_cu_seqlens_q = self.get_device_int32_arange(len(dsa_cu_seqlens_k)) paged_mqa_schedule_metadata = None - paged_mqa_ctx_lens_2d = None + # DeepGEMM paged MQA logits path needs a schedule metadata tensor. + # Compute it once per forward batch and reuse it across layers. 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(include_v2=True) ): - paged_mqa_ctx_lens_2d = self._build_paged_mqa_schedule_2d_ctx_lens( - forward_batch.forward_mode, - cache_seqlens_int32, - seqlens_expanded, - forward_batch.batch_size, - ) - # NOTE: block_kv arg must be 64 here — DG computes SPLIT_KV = - # block_kv * 4 and both DG's and the indexer's compute kernels - # require SPLIT_KV = 256; this is independent of the cache page size. - paged_mqa_schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata( - paged_mqa_ctx_lens_2d, 64, deep_gemm.get_num_sms() - ) + try: + import deep_gemm + + # NOTE: DeepGEMM paged path uses block_size=64. + seqlens_32 = ( + seqlens_expanded + if ( + forward_batch.forward_mode.is_target_verify() + or forward_batch.forward_mode.is_draft_extend(include_v2=True) + ) + else cache_seqlens_int32 + ) + seqlens_32_2d = _to_2d_context_lens( + seqlens_32, forward_batch.batch_size + ) + paged_mqa_schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata( + seqlens_32_2d, 64, deep_gemm.get_num_sms() + ) + except (ImportError, ModuleNotFoundError): + paged_mqa_schedule_metadata = None metadata = DSAMetadata( page_size=self.real_page_size, @@ -724,7 +701,6 @@ class DeepseekSparseAttnBackend( else None ), paged_mqa_schedule_metadata=paged_mqa_schedule_metadata, - paged_mqa_ctx_lens_2d=paged_mqa_ctx_lens_2d, dsa_cache_seqlens_int32=dsa_cache_seqlens_int32, dsa_cu_seqlens_q=dsa_cu_seqlens_q, dsa_cu_seqlens_k=dsa_cu_seqlens_k, @@ -971,18 +947,28 @@ class DeepseekSparseAttnBackend( real_page_table = self._transform_table_1_to_real(page_table_1) paged_mqa_schedule_metadata = None - paged_mqa_ctx_lens_2d = None if is_cuda() and ( forward_mode.is_decode_or_idle() or forward_mode.is_target_verify() or forward_mode.is_draft_extend(include_v2=True) ): - paged_mqa_ctx_lens_2d = self._build_paged_mqa_schedule_2d_ctx_lens( - forward_mode, cache_seqlens_int32, seqlens_expanded, bs - ) - paged_mqa_schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata( - paged_mqa_ctx_lens_2d, 64, deep_gemm.get_num_sms() - ) + try: + import deep_gemm + + seqlens_32 = ( + seqlens_expanded + if ( + forward_mode.is_target_verify() + or forward_mode.is_draft_extend(include_v2=True) + ) + else cache_seqlens_int32 + ) + seqlens_32_2d = _to_2d_context_lens(seqlens_32, bs) + paged_mqa_schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata( + seqlens_32_2d, 64, deep_gemm.get_num_sms() + ) + except (ImportError, ModuleNotFoundError): + paged_mqa_schedule_metadata = None metadata = DSAMetadata( page_size=self.real_page_size, @@ -994,7 +980,6 @@ class DeepseekSparseAttnBackend( page_table_1=page_table_1, flashmla_metadata=flashmla_metadata, paged_mqa_schedule_metadata=paged_mqa_schedule_metadata, - paged_mqa_ctx_lens_2d=paged_mqa_ctx_lens_2d, dsa_cache_seqlens_int32=dsa_cache_seqlens_int32, dsa_cu_seqlens_q=dsa_cu_seqlens_q, dsa_cu_seqlens_k=dsa_cu_seqlens_k, @@ -1135,26 +1120,29 @@ class DeepseekSparseAttnBackend( or forward_mode.is_target_verify() or forward_mode.is_draft_extend(include_v2=True) ): - seqlens_32_2d = self._build_paged_mqa_schedule_2d_ctx_lens( - forward_mode, - metadata.cache_seqlens_int32, - seqlens_expanded, - 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 + try: + import deep_gemm + + seqlens_32 = ( + seqlens_expanded + if ( + forward_mode.is_target_verify() + or forward_mode.is_draft_extend(include_v2=True) + ) + else metadata.cache_seqlens_int32 ) - else: - metadata.paged_mqa_schedule_metadata.copy_(new_schedule) - # `copy_` preserves the buffer's data_ptr that the captured graph captured. - if metadata.paged_mqa_ctx_lens_2d is None: - object.__setattr__(metadata, "paged_mqa_ctx_lens_2d", seqlens_32_2d) - else: - metadata.paged_mqa_ctx_lens_2d.copy_(seqlens_32_2d) + 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): + object.__setattr__(metadata, "paged_mqa_schedule_metadata", None) seqlens_expanded_size = seqlens_expanded.shape[0] assert ( metadata.dsa_cache_seqlens_int32 is not None @@ -1346,28 +1334,27 @@ class DeepseekSparseAttnBackend( # deadlock the kernel when the runtime work decomposition diverges from # the captured one). if is_cuda(): - if forward_mode.is_decode_or_idle(): - seqlens_32_2d = _to_2d_context_lens(metadata.cache_seqlens_int32, bs) - else: - seqlens_32_2d = self._build_paged_mqa_schedule_2d_ctx_lens( - forward_mode, - metadata.cache_seqlens_int32, - metadata.dsa_seqlens_expanded[: precomputed.seqlens_expanded_size], - bs, + try: + import deep_gemm + + if forward_mode.is_decode_or_idle(): + seqlens_32 = metadata.cache_seqlens_int32 + else: + seqlens_32 = metadata.dsa_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() ) - 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) - if metadata.paged_mqa_ctx_lens_2d is None: - object.__setattr__(metadata, "paged_mqa_ctx_lens_2d", seqlens_32_2d) - else: - metadata.paged_mqa_ctx_lens_2d.copy_(seqlens_32_2d) + 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 @@ -2332,7 +2319,6 @@ class DeepseekSparseAttnBackend( ), topk_backend=self.dsa_topk_backend, paged_mqa_schedule_metadata=self.forward_metadata.paged_mqa_schedule_metadata, - paged_mqa_ctx_lens_2d=self.forward_metadata.paged_mqa_ctx_lens_2d, force_unfused_topk=force_unfused, )