Support NextN = 2/4 in DSV32 (#24870)

Co-authored-by: b8zhong <b8zhong@users.noreply.github.com>
This commit is contained in:
Brayden Zhong
2026-06-02 19:06:06 -07:00
committed by GitHub
co-authored by b8zhong
parent 2d8cb87de7
commit 13852d3f31
2 changed files with 136 additions and 92 deletions
@@ -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,
@@ -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,
)