From 13852d3f311750811c24d8ed0a5eeb1b6ada4688 Mon Sep 17 00:00:00 2001 From: Brayden Zhong Date: Tue, 2 Jun 2026 19:06:06 -0700 Subject: [PATCH] Support NextN = 2/4 in DSV32 (#24870) Co-authored-by: b8zhong --- .../srt/layers/attention/dsa/dsa_indexer.py | 50 ++++- .../srt/layers/attention/dsa_backend.py | 178 ++++++++++-------- 2 files changed, 136 insertions(+), 92 deletions(-) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 318a147ba..85fcd4b9e 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -602,11 +602,31 @@ 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) - # 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: + + 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: seqlens_32_2d = seqlens_32 else: seqlens_32_2d = seqlens_32.unsqueeze(-1) @@ -616,8 +636,6 @@ 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 @@ -628,12 +646,10 @@ 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), @@ -651,7 +667,21 @@ 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 e3b8eebd5..eb2c50cb4 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -50,7 +50,10 @@ 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 +from sglang.srt.utils import is_cuda, is_hip, is_sm100_supported + +if is_cuda(): + import deep_gemm if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -153,6 +156,9 @@ 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 @@ -199,6 +205,7 @@ 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: @@ -380,6 +387,31 @@ 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() @@ -654,33 +686,24 @@ class DeepseekSparseAttnBackend( dsa_cu_seqlens_q = self.get_device_int32_arange(len(dsa_cu_seqlens_k)) paged_mqa_schedule_metadata = None - # DeepGEMM paged MQA logits path needs a schedule metadata tensor. - # Compute it once per forward batch and reuse it across layers. + paged_mqa_ctx_lens_2d = None 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) ): - 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 + 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() + ) metadata = DSAMetadata( page_size=self.real_page_size, @@ -701,6 +724,7 @@ 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, @@ -947,28 +971,18 @@ 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) ): - 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 + 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() + ) metadata = DSAMetadata( page_size=self.real_page_size, @@ -980,6 +994,7 @@ 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, @@ -1120,29 +1135,26 @@ class DeepseekSparseAttnBackend( or forward_mode.is_target_verify() or forward_mode.is_draft_extend(include_v2=True) ): - 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 + 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 ) - 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) + 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_expanded_size = seqlens_expanded.shape[0] assert ( metadata.dsa_cache_seqlens_int32 is not None @@ -1334,27 +1346,28 @@ class DeepseekSparseAttnBackend( # 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.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() + 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, ) - 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 + 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) self.forward_metadata = metadata @@ -2319,6 +2332,7 @@ 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, )