Support NextN = 2/4 in DSV32 (#24870)
Co-authored-by: b8zhong <b8zhong@users.noreply.github.com>
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user