diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 9e65c7f3c..a9166094e 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -650,6 +650,7 @@ class Envs: SGLANG_DSA_MQA_LOGITS_FREE_MEM_FRACTION = EnvFloat(0.2) SGLANG_ENABLE_PCG_DSV2_DUAL_STREAM = EnvBool(False) SGLANG_USE_FUSED_METADATA_COPY = EnvBool(True) + SGLANG_DSA_USE_FUSED_METADATA_GENERATION = EnvBool(True) SGLANG_DSA_TOPK_BROADCAST = EnvBool(False) SGLANG_DISABLE_DSA_INDEXER_FUSION = EnvBool(False) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py b/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py index dc3cb84ed..d75fe9611 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py @@ -11,11 +11,20 @@ from typing import TYPE_CHECKING, Optional import torch +from sglang.srt.environ import envs from sglang.srt.layers.attention.dsa.utils import compute_dsa_seqlens +from sglang.srt.layers.attention.utils import seqlens_expand_triton +from sglang.srt.utils import is_cuda, is_hip if TYPE_CHECKING: from sglang.srt.model_executor.forward_batch_info import ForwardMode +_is_cuda = is_cuda() +_is_hip = is_hip() +_USE_FUSED_METADATA_GENERATION = ( + envs.SGLANG_DSA_USE_FUSED_METADATA_GENERATION.get() and not _is_hip +) + @dataclass class PrecomputedMetadata: @@ -116,6 +125,69 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: """Precompute metadata for normal decode mode.""" max_len = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1] + if _USE_FUSED_METADATA_GENERATION and _is_cuda: + from sglang.srt.layers.attention.triton_ops.dsa_metadata import ( + fused_dsa_decode_metadata, + ) + + cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device) + cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device) + page_indices = torch.empty( + (bs, max_len), dtype=torch.int32, device=self.device + ) + dsa_cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device) + dsa_cu_seqlens_k = torch.empty( + bs + 1, dtype=torch.int32, device=self.device + ) + if self.real_page_size > 1: + real_cols = (max_len + self.real_page_size - 1) // self.real_page_size + real_page_table = torch.empty( + (bs, real_cols), dtype=torch.int32, device=self.device + ) + real_page_table_arg = real_page_table + else: + real_page_table = None + real_page_table_arg = page_indices + + fused_dsa_decode_metadata( + seq_lens=seq_lens, + req_pool_indices=req_pool_indices, + req_to_token=self.req_to_token, + cache_seqlens=cache_seqlens, + cu_seqlens_k=cu_seqlens_k, + page_table_1=page_indices, + dsa_cache_seqlens=dsa_cache_seqlens, + dsa_cu_seqlens_k=dsa_cu_seqlens_k, + real_page_table=real_page_table_arg, + bs=bs, + max_len=max_len, + dsa_index_topk=self.dsa_index_topk, + real_page_size=self.real_page_size, + ) + seqlens_expanded = cache_seqlens + seqlens_expanded_size = bs + + flashmla_metadata = None + if self.dsa_decode_impl == "flashmla_kv": + flashmla_metadata = self._compute_flashmla_metadata( + cache_seqlens=dsa_cache_seqlens, + seq_len_q=1, + ) + + return PrecomputedMetadata( + cache_seqlens=cache_seqlens, + cu_seqlens_k=cu_seqlens_k, + page_indices=page_indices, + real_page_table=real_page_table, + seqlens_expanded=seqlens_expanded, + dsa_cache_seqlens=dsa_cache_seqlens, + dsa_cu_seqlens_k=dsa_cu_seqlens_k, + seqlens_expanded_size=seqlens_expanded_size, + max_len=max_len, + max_seqlen_k=max_len, + flashmla_metadata=flashmla_metadata, + ) + # Convert to int32 and compute cumsum cache_seqlens = seq_lens.to(torch.int32) cu_seqlens_k = compute_cu_seqlens(cache_seqlens) @@ -169,9 +241,84 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: seq_lens_cpu: torch.Tensor, ) -> PrecomputedMetadata: """Precompute metadata for target verify mode.""" - max_seqlen_k = int( - seq_lens_cpu.max().item() + self.speculative_num_draft_tokens - ) + max_seqlen_k = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1] + seqlens_expanded_size = bs * self.speculative_num_draft_tokens + + if _USE_FUSED_METADATA_GENERATION and _is_cuda: + from sglang.srt.layers.attention.triton_ops.dsa_metadata import ( + fused_dsa_target_verify_metadata, + ) + + cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device) + cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device) + page_indices = torch.empty( + (seqlens_expanded_size, max_seqlen_k), + dtype=torch.int32, + device=self.device, + ) + seqlens_expanded = torch.empty( + seqlens_expanded_size, dtype=torch.int32, device=self.device + ) + dsa_cache_seqlens = torch.empty( + seqlens_expanded_size, dtype=torch.int32, device=self.device + ) + dsa_cu_seqlens_k = torch.empty( + seqlens_expanded_size + 1, + dtype=torch.int32, + device=self.device, + ) + if self.real_page_size > 1: + real_cols = ( + max_seqlen_k + self.real_page_size - 1 + ) // self.real_page_size + real_page_table = torch.empty( + (seqlens_expanded_size, real_cols), + dtype=torch.int32, + device=self.device, + ) + real_page_table_arg = real_page_table + else: + real_page_table = None + real_page_table_arg = page_indices + + fused_dsa_target_verify_metadata( + seq_lens=seq_lens, + req_pool_indices=req_pool_indices, + req_to_token=self.req_to_token, + cache_seqlens=cache_seqlens, + cu_seqlens_k=cu_seqlens_k, + page_table_1=page_indices, + seqlens_expanded=seqlens_expanded, + dsa_cache_seqlens=dsa_cache_seqlens, + dsa_cu_seqlens_k=dsa_cu_seqlens_k, + real_page_table=real_page_table_arg, + bs=bs, + max_seqlen_k=max_seqlen_k, + dsa_index_topk=self.dsa_index_topk, + real_page_size=self.real_page_size, + next_n=self.speculative_num_draft_tokens, + ) + + flashmla_metadata = None + if self.dsa_decode_impl == "flashmla_kv": + flashmla_metadata = self._compute_flashmla_metadata( + cache_seqlens=dsa_cache_seqlens, + seq_len_q=1, + ) + + return PrecomputedMetadata( + cache_seqlens=cache_seqlens, + cu_seqlens_k=cu_seqlens_k, + page_indices=page_indices, + real_page_table=real_page_table, + seqlens_expanded=seqlens_expanded, + dsa_cache_seqlens=dsa_cache_seqlens, + dsa_cu_seqlens_k=dsa_cu_seqlens_k, + seqlens_expanded_size=seqlens_expanded_size, + max_len=-1, + max_seqlen_k=max_seqlen_k, + flashmla_metadata=flashmla_metadata, + ) # Cache seqlens with draft tokens cache_seqlens = (seq_lens + self.speculative_num_draft_tokens).to(torch.int32) @@ -183,26 +330,19 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: page_indices, repeats=self.speculative_num_draft_tokens, dim=0 ).contiguous() - # Generate expanded seqlens - extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * bs - seqlens_int32_cpu = [ - self.speculative_num_draft_tokens + kv_len - for kv_len in seq_lens_cpu.tolist() - ] - seqlens_expanded = torch.cat( - [ - torch.arange( - kv_len - qo_len + 1, - kv_len + 1, - dtype=torch.int32, - device=self.device, - ) - for qo_len, kv_len in zip( - extend_seq_lens_cpu, - seqlens_int32_cpu, - strict=True, - ) - ] + # Generate expanded seqlens on device. seq_lens_cpu is optional for DSA + # CUDA graph replay, so this fallback must not require a host mirror. + extend_seq_lens = torch.full( + (bs,), + self.speculative_num_draft_tokens, + dtype=torch.int32, + device=self.device, + ) + seqlens_expanded = seqlens_expand_triton( + extend_seq_lens, + cache_seqlens, + bs * self.speculative_num_draft_tokens, + self.speculative_num_draft_tokens, ) # Compute DSA seqlens diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 13f78f253..1502b6a2b 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -117,6 +117,9 @@ global_workspace_buffer = None # Control whether to use fused metadata copy kernel for cuda graph replay (default: enabled) # Set SGLANG_USE_FUSED_METADATA_COPY=0 or false to disable _USE_FUSED_METADATA_COPY = envs.SGLANG_USE_FUSED_METADATA_COPY.get() and not _is_hip +_USE_FUSED_METADATA_GENERATION = ( + envs.SGLANG_DSA_USE_FUSED_METADATA_GENERATION.get() and not _is_hip +) @dataclass(frozen=True) @@ -588,6 +591,19 @@ class DeepseekSparseAttnBackend( return _to_2d_context_lens(seqlens_expanded, batch_size) return _to_2d_context_lens(cache_seqlens_int32, batch_size) + def _refresh_paged_mqa_schedule_metadata( + self, + metadata: DSAMetadata, + seqlens_32_2d: torch.Tensor, + ) -> None: + 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) + def _get_fused_topk_page_table(self, topk_indices: torch.Tensor) -> torch.Tensor: if ( self.dsa_topk_backend.is_sgl_kernel() @@ -1220,58 +1236,133 @@ class DeepseekSparseAttnBackend( # Normal Decode metadata: DSAMetadata = self.decode_cuda_graph_metadata[bs] + used_fused_metadata_generation = False + target_verify_ctx_lens_written = False if forward_mode.is_decode_or_idle(): # Normal Decode max_len = metadata.page_table_1.shape[1] - cache_seqlens = seq_lens.to(torch.int32) - metadata.cache_seqlens_int32.copy_(cache_seqlens) - metadata.cu_seqlens_k[1:].copy_( - torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32) - ) - page_indices = self.req_to_token[req_pool_indices, :max_len] - metadata.page_table_1[:, :max_len].copy_(page_indices) - dsa_cache_seqlens = compute_dsa_seqlens( - cache_seqlens, dsa_index_topk=self.dsa_index_topk - ) - metadata.dsa_cache_seqlens_int32.copy_(dsa_cache_seqlens) - seqlens_expanded = cache_seqlens + if _USE_FUSED_METADATA_GENERATION and is_cuda(): + from sglang.srt.layers.attention.triton_ops.dsa_metadata import ( + fused_dsa_decode_metadata, + ) + + fused_dsa_decode_metadata( + seq_lens=seq_lens, + req_pool_indices=req_pool_indices, + req_to_token=self.req_to_token, + cache_seqlens=metadata.cache_seqlens_int32, + cu_seqlens_k=metadata.cu_seqlens_k, + page_table_1=metadata.page_table_1, + dsa_cache_seqlens=metadata.dsa_cache_seqlens_int32, + dsa_cu_seqlens_k=metadata.dsa_cu_seqlens_k, + real_page_table=metadata.real_page_table, + bs=bs, + max_len=max_len, + dsa_index_topk=self.dsa_index_topk, + real_page_size=self.real_page_size, + ) + cache_seqlens = metadata.cache_seqlens_int32 + dsa_cache_seqlens = metadata.dsa_cache_seqlens_int32 + seqlens_expanded = cache_seqlens + page_indices = None + used_fused_metadata_generation = True + + if not used_fused_metadata_generation: + cache_seqlens = seq_lens.to(torch.int32) + metadata.cache_seqlens_int32.copy_(cache_seqlens) + metadata.cu_seqlens_k[1:].copy_( + torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32) + ) + page_indices = self.req_to_token[req_pool_indices, :max_len] + metadata.page_table_1[:, :max_len].copy_(page_indices) + dsa_cache_seqlens = compute_dsa_seqlens( + cache_seqlens, dsa_index_topk=self.dsa_index_topk + ) + metadata.dsa_cache_seqlens_int32.copy_(dsa_cache_seqlens) + seqlens_expanded = cache_seqlens elif forward_mode.is_target_verify(): max_seqlen_k = metadata.page_table_1.shape[1] - cache_seqlens = (seq_lens + self.speculative_num_draft_tokens).to( - torch.int32 - ) - metadata.cache_seqlens_int32.copy_(cache_seqlens) - metadata.cu_seqlens_k[1:].copy_( - torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32) - ) - page_indices = self.req_to_token[req_pool_indices, :max_seqlen_k] - page_indices = torch.repeat_interleave( - page_indices, repeats=self.speculative_num_draft_tokens, dim=0 - ) - metadata.page_table_1[:, :max_seqlen_k].copy_(page_indices) + if _USE_FUSED_METADATA_GENERATION and is_cuda(): + from sglang.srt.layers.attention.triton_ops.dsa_metadata import ( + fused_dsa_target_verify_metadata, + ) - # Fill the constant per-req qo lengths (num_draft_tokens) on-device; - # torch.tensor(list, device=cuda) does a pageable H2D copy that - # blocks the host on the whole queued stream. - extend_seq_lens = torch.full( - (bs,), - self.speculative_num_draft_tokens, - dtype=torch.int32, - device=self.device, - ) - seqlens_expanded = seqlens_expand_triton( - extend_seq_lens, - cache_seqlens, - self.speculative_num_draft_tokens * bs, - self.speculative_num_draft_tokens, - ) - metadata.dsa_seqlens_expanded.copy_(seqlens_expanded) - dsa_cache_seqlens = compute_dsa_seqlens( - seqlens_expanded, self.dsa_index_topk - ) - metadata.dsa_cache_seqlens_int32.copy_(dsa_cache_seqlens) + paged_mqa_ctx_lens_2d = None + if ( + self.speculative_num_draft_tokens >= 2 + and is_sm100_supported() + and metadata.paged_mqa_ctx_lens_2d is not None + and metadata.paged_mqa_ctx_lens_2d.dim() == 2 + and metadata.paged_mqa_ctx_lens_2d.size(0) == bs + and metadata.paged_mqa_ctx_lens_2d.size(1) + == self.speculative_num_draft_tokens + ): + paged_mqa_ctx_lens_2d = metadata.paged_mqa_ctx_lens_2d + + fused_dsa_target_verify_metadata( + seq_lens=seq_lens, + req_pool_indices=req_pool_indices, + req_to_token=self.req_to_token, + cache_seqlens=metadata.cache_seqlens_int32, + cu_seqlens_k=metadata.cu_seqlens_k, + page_table_1=metadata.page_table_1, + seqlens_expanded=metadata.dsa_seqlens_expanded, + dsa_cache_seqlens=metadata.dsa_cache_seqlens_int32, + dsa_cu_seqlens_k=metadata.dsa_cu_seqlens_k, + real_page_table=metadata.real_page_table, + bs=bs, + max_seqlen_k=max_seqlen_k, + dsa_index_topk=self.dsa_index_topk, + real_page_size=self.real_page_size, + next_n=self.speculative_num_draft_tokens, + paged_mqa_ctx_lens_2d=paged_mqa_ctx_lens_2d, + ) + target_verify_ctx_lens_written = paged_mqa_ctx_lens_2d is not None + cache_seqlens = metadata.cache_seqlens_int32 + seqlens_expanded = metadata.dsa_seqlens_expanded[ + : self.speculative_num_draft_tokens * bs + ] + dsa_cache_seqlens = metadata.dsa_cache_seqlens_int32[ + : self.speculative_num_draft_tokens * bs + ] + page_indices = None + used_fused_metadata_generation = True + + if not used_fused_metadata_generation: + cache_seqlens = (seq_lens + self.speculative_num_draft_tokens).to( + torch.int32 + ) + metadata.cache_seqlens_int32.copy_(cache_seqlens) + metadata.cu_seqlens_k[1:].copy_( + torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32) + ) + page_indices = self.req_to_token[req_pool_indices, :max_seqlen_k] + page_indices = torch.repeat_interleave( + page_indices, repeats=self.speculative_num_draft_tokens, dim=0 + ) + metadata.page_table_1[:, :max_seqlen_k].copy_(page_indices) + + # Fill the constant per-req qo lengths on-device; torch.tensor(list, + # device=cuda) does a pageable H2D copy that blocks the host. + extend_seq_lens = torch.full( + (bs,), + self.speculative_num_draft_tokens, + dtype=torch.int32, + device=self.device, + ) + seqlens_expanded = seqlens_expand_triton( + extend_seq_lens, + cache_seqlens, + self.speculative_num_draft_tokens * bs, + self.speculative_num_draft_tokens, + ) + metadata.dsa_seqlens_expanded.copy_(seqlens_expanded) + dsa_cache_seqlens = compute_dsa_seqlens( + seqlens_expanded, self.dsa_index_topk + ) + metadata.dsa_cache_seqlens_int32.copy_(dsa_cache_seqlens) elif forward_mode.is_draft_extend_v2(): # V2 draft-extend processes the full padded tree width # (speculative_num_draft_tokens) per req -- a static shape, like @@ -1280,17 +1371,7 @@ class DeepseekSparseAttnBackend( # the per-req accept length is handled downstream by output # selection, not by reshaping the page table here. max_seqlen_k = metadata.page_table_1.shape[1] - cache_seqlens = seq_lens.to(torch.int32) - metadata.cache_seqlens_int32.copy_(cache_seqlens) - metadata.cu_seqlens_k[1:].copy_( - torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32) - ) - - page_indices = self.req_to_token[req_pool_indices, :max_seqlen_k] - page_indices = torch.repeat_interleave( - page_indices, repeats=self.speculative_num_draft_tokens, dim=0 - ) - metadata.page_table_1[:, :max_seqlen_k].copy_(page_indices) + total_extend_len = self.speculative_num_draft_tokens * bs # See target-verify note: fill on-device to avoid the blocking # pageable H2D from torch.tensor(list, device=cuda). @@ -1300,17 +1381,65 @@ class DeepseekSparseAttnBackend( dtype=torch.int32, device=self.device, ) - seqlens_expanded = seqlens_expand_triton( - extend_seq_lens, - cache_seqlens, - self.speculative_num_draft_tokens * bs, - self.speculative_num_draft_tokens, - ) - metadata.dsa_seqlens_expanded.copy_(seqlens_expanded) - dsa_cache_seqlens = compute_dsa_seqlens( - seqlens_expanded, self.dsa_index_topk - ) - metadata.dsa_cache_seqlens_int32.copy_(dsa_cache_seqlens) + + if _USE_FUSED_METADATA_GENERATION and is_cuda(): + from sglang.srt.layers.attention.triton_ops.dsa_metadata import ( + fused_dsa_draft_extend_metadata, + ) + + fused_dsa_draft_extend_metadata( + seq_lens=seq_lens, + extend_seq_lens=extend_seq_lens, + req_pool_indices=req_pool_indices, + req_to_token=self.req_to_token, + cache_seqlens=metadata.cache_seqlens_int32, + cu_seqlens_k=metadata.cu_seqlens_k, + page_table_1=metadata.page_table_1, + seqlens_expanded=metadata.dsa_seqlens_expanded, + dsa_cache_seqlens=metadata.dsa_cache_seqlens_int32, + dsa_cu_seqlens_k=metadata.dsa_cu_seqlens_k, + real_page_table=metadata.real_page_table, + bs=bs, + total_len=total_extend_len, + max_seqlen_k=max_seqlen_k, + dsa_index_topk=self.dsa_index_topk, + real_page_size=self.real_page_size, + max_extend_len=self.speculative_num_draft_tokens, + max_total_len=bs * self.speculative_num_draft_tokens, + static_extend_len=True, + ) + cache_seqlens = metadata.cache_seqlens_int32 + seqlens_expanded = metadata.dsa_seqlens_expanded[:total_extend_len] + dsa_cache_seqlens = metadata.dsa_cache_seqlens_int32[:total_extend_len] + page_indices = None + used_fused_metadata_generation = True + + if not used_fused_metadata_generation: + cache_seqlens = seq_lens.to(torch.int32) + metadata.cache_seqlens_int32.copy_(cache_seqlens) + metadata.cu_seqlens_k[1:].copy_( + torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32) + ) + + page_indices = self.req_to_token[req_pool_indices, :max_seqlen_k] + page_indices = torch.repeat_interleave( + page_indices, repeats=self.speculative_num_draft_tokens, dim=0 + ) + metadata.page_table_1[:, :max_seqlen_k].copy_(page_indices) + + seqlens_expanded = seqlens_expand_triton( + extend_seq_lens, + cache_seqlens, + total_extend_len, + self.speculative_num_draft_tokens, + ) + metadata.dsa_seqlens_expanded[: seqlens_expanded.shape[0]].copy_( + seqlens_expanded + ) + dsa_cache_seqlens = compute_dsa_seqlens( + seqlens_expanded, self.dsa_index_topk + ) + metadata.dsa_cache_seqlens_int32.copy_(dsa_cache_seqlens) # Update DeepGEMM paged MQA schedule metadata outside the captured graph. if is_cuda() and ( @@ -1322,26 +1451,22 @@ class DeepseekSparseAttnBackend( schedule_seqlens_expanded = metadata.dsa_seqlens_expanded else: schedule_seqlens_expanded = seqlens_expanded - seqlens_32_2d = self._build_paged_mqa_schedule_2d_ctx_lens( - forward_mode, - metadata.cache_seqlens_int32, - schedule_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 + if target_verify_ctx_lens_written: + seqlens_32_2d = metadata.paged_mqa_ctx_lens_2d + else: + seqlens_32_2d = self._build_paged_mqa_schedule_2d_ctx_lens( + forward_mode, + metadata.cache_seqlens_int32, + schedule_seqlens_expanded, + bs, ) - else: - metadata.paged_mqa_schedule_metadata.copy_(new_schedule) + self._refresh_paged_mqa_schedule_metadata(metadata, seqlens_32_2d) # `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) + if not target_verify_ctx_lens_written: + 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 @@ -1349,17 +1474,19 @@ class DeepseekSparseAttnBackend( and self.dsa_index_topk is not None ) - metadata.dsa_cu_seqlens_k[1 : 1 + seqlens_expanded_size].copy_( - torch.cumsum(dsa_cache_seqlens, dim=0, dtype=torch.int32) - ) + if not used_fused_metadata_generation: + metadata.dsa_cu_seqlens_k[1 : 1 + seqlens_expanded_size].copy_( + torch.cumsum(dsa_cache_seqlens, dim=0, dtype=torch.int32) + ) # NOTE(dark): (dsa-) cu_seqlens_q is always arange, no need to copy assert self.real_page_size == metadata.page_size if self.real_page_size > 1: - real_table = self._transform_table_1_to_real(page_indices) - new_rows = real_table.shape[0] - new_cols = real_table.shape[1] - metadata.real_page_table[:new_rows, :new_cols].copy_(real_table) + if not used_fused_metadata_generation: + real_table = self._transform_table_1_to_real(page_indices) + new_rows = real_table.shape[0] + new_cols = real_table.shape[1] + metadata.real_page_table[:new_rows, :new_cols].copy_(real_table) else: assert metadata.real_page_table is metadata.page_table_1 @@ -1528,15 +1655,7 @@ class DeepseekSparseAttnBackend( metadata.dsa_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 - ) - else: - metadata.paged_mqa_schedule_metadata.copy_(new_schedule) + self._refresh_paged_mqa_schedule_metadata(metadata, seqlens_32_2d) if metadata.paged_mqa_ctx_lens_2d is None: object.__setattr__(metadata, "paged_mqa_ctx_lens_2d", seqlens_32_2d) else: diff --git a/python/sglang/srt/layers/attention/triton_ops/dsa_metadata.py b/python/sglang/srt/layers/attention/triton_ops/dsa_metadata.py new file mode 100644 index 000000000..38686484b --- /dev/null +++ b/python/sglang/srt/layers/attention/triton_ops/dsa_metadata.py @@ -0,0 +1,621 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit( + do_not_specialize=[ + "page_table_stride_0", + "real_page_table_stride_0", + "max_len", + ] +) +def _fused_dsa_decode_metadata_kernel( + seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + seq_lens_stride: tl.constexpr, + req_pool_indices_stride: tl.constexpr, + req_to_token_stride_0: tl.constexpr, + req_to_token_stride_1: tl.constexpr, + page_table_stride_0, + page_table_stride_1: tl.constexpr, + real_page_table_stride_0, + real_page_table_stride_1: tl.constexpr, + bs: tl.constexpr, + max_len, + dsa_index_topk: tl.constexpr, + real_page_size: tl.constexpr, + HAS_REAL_PAGE_TABLE: tl.constexpr, + BLOCK_BS: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + + if pid == 0: + offs_b = tl.arange(0, BLOCK_BS) + mask_b = offs_b < bs + seq = tl.load(seq_lens + offs_b * seq_lens_stride, mask=mask_b, other=0) + seq_i32 = seq.to(tl.int32) + dsa_seq = tl.minimum(seq_i32, dsa_index_topk) + + cu = tl.cumsum(seq_i32, 0) + dsa_cu = tl.cumsum(dsa_seq, 0) + + tl.store(cache_seqlens + offs_b, seq_i32, mask=mask_b) + tl.store(cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(cu_seqlens_k + 1 + offs_b, cu, mask=mask_b) + tl.store(dsa_cache_seqlens + offs_b, dsa_seq, mask=mask_b) + tl.store(dsa_cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(dsa_cu_seqlens_k + 1 + offs_b, dsa_cu, mask=mask_b) + return + + num_col_blocks = tl.cdiv(max_len, BLOCK_N) + page_pid = pid - 1 + row = page_pid // num_col_blocks + col_block = page_pid - row * num_col_blocks + offs_n = col_block * BLOCK_N + tl.arange(0, BLOCK_N) + mask = (row < bs) & (offs_n < max_len) + + req_idx = tl.load( + req_pool_indices + row * req_pool_indices_stride, + mask=row < bs, + other=0, + ) + vals = tl.load( + req_to_token + req_idx * req_to_token_stride_0 + offs_n * req_to_token_stride_1, + mask=mask, + other=0, + ).to(tl.int32) + tl.store( + page_table_1 + row * page_table_stride_0 + offs_n * page_table_stride_1, + vals, + mask=mask, + ) + + if HAS_REAL_PAGE_TABLE: + real_mask = mask & ((offs_n % real_page_size) == 0) + real_cols = offs_n // real_page_size + tl.store( + real_page_table + + row * real_page_table_stride_0 + + real_cols * real_page_table_stride_1, + vals // real_page_size, + mask=real_mask, + ) + + +def fused_dsa_decode_metadata( + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + req_to_token: torch.Tensor, + cache_seqlens: torch.Tensor, + cu_seqlens_k: torch.Tensor, + page_table_1: torch.Tensor, + dsa_cache_seqlens: torch.Tensor, + dsa_cu_seqlens_k: torch.Tensor, + real_page_table: torch.Tensor, + bs: int, + max_len: int, + dsa_index_topk: int, + real_page_size: int, +) -> None: + assert seq_lens.is_cuda + assert req_pool_indices.is_cuda + assert req_to_token.is_cuda + assert cache_seqlens.is_cuda + assert cu_seqlens_k.is_cuda + assert page_table_1.is_cuda + assert dsa_cache_seqlens.is_cuda + assert dsa_cu_seqlens_k.is_cuda + + if bs == 0: + cu_seqlens_k[:1].zero_() + dsa_cu_seqlens_k[:1].zero_() + return + + has_real_page_table = real_page_size > 1 + if has_real_page_table: + assert real_page_table is not None + assert real_page_table.is_cuda + else: + real_page_table = page_table_1 + + block_bs = triton.next_power_of_2(bs) + block_n = 128 + num_col_blocks = triton.cdiv(max_len, block_n) + grid = (1 + bs * num_col_blocks,) + + _fused_dsa_decode_metadata_kernel[grid]( + seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + seq_lens.stride(0), + req_pool_indices.stride(0), + req_to_token.stride(0), + req_to_token.stride(1), + page_table_1.stride(0), + page_table_1.stride(1), + real_page_table.stride(0) if has_real_page_table else 0, + real_page_table.stride(1) if has_real_page_table else 0, + bs, + max_len, + dsa_index_topk, + real_page_size, + has_real_page_table, + BLOCK_BS=block_bs, + BLOCK_N=block_n, + ) + + +@triton.jit( + do_not_specialize=[ + "page_table_stride_0", + "real_page_table_stride_0", + "max_seqlen_k", + ] +) +def _fused_dsa_target_verify_metadata_kernel( + seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + seqlens_expanded, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + paged_mqa_ctx_lens_2d, + seq_lens_stride: tl.constexpr, + req_pool_indices_stride: tl.constexpr, + req_to_token_stride_0: tl.constexpr, + req_to_token_stride_1: tl.constexpr, + page_table_stride_0, + page_table_stride_1: tl.constexpr, + real_page_table_stride_0, + real_page_table_stride_1: tl.constexpr, + paged_mqa_ctx_lens_stride_0: tl.constexpr, + paged_mqa_ctx_lens_stride_1: tl.constexpr, + bs: tl.constexpr, + max_seqlen_k, + dsa_index_topk: tl.constexpr, + real_page_size: tl.constexpr, + next_n: tl.constexpr, + HAS_REAL_PAGE_TABLE: tl.constexpr, + HAS_PAGED_MQA_CTX_LENS: tl.constexpr, + BLOCK_BS: tl.constexpr, + BLOCK_EXPANDED: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + expanded_size: tl.constexpr = bs * next_n + + if pid == 0: + offs_b = tl.arange(0, BLOCK_BS) + mask_b = offs_b < bs + seq = tl.load(seq_lens + offs_b * seq_lens_stride, mask=mask_b, other=0) + cache_seq = seq.to(tl.int32) + next_n + cu = tl.cumsum(cache_seq, 0) + + tl.store(cache_seqlens + offs_b, cache_seq, mask=mask_b) + tl.store(cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(cu_seqlens_k + 1 + offs_b, cu, mask=mask_b) + + offs_e = tl.arange(0, BLOCK_EXPANDED) + mask_e = offs_e < expanded_size + req_row = offs_e // next_n + draft_off = offs_e - req_row * next_n + base_seq = tl.load( + seq_lens + req_row * seq_lens_stride, + mask=mask_e, + other=0, + ).to(tl.int32) + expanded_seq = base_seq + draft_off + 1 + expanded_seq = tl.where(mask_e, expanded_seq, 0) + dsa_seq = tl.minimum(expanded_seq, dsa_index_topk) + dsa_cu = tl.cumsum(dsa_seq, 0) + + tl.store(seqlens_expanded + offs_e, expanded_seq, mask=mask_e) + tl.store(dsa_cache_seqlens + offs_e, dsa_seq, mask=mask_e) + tl.store(dsa_cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(dsa_cu_seqlens_k + 1 + offs_e, dsa_cu, mask=mask_e) + + if HAS_PAGED_MQA_CTX_LENS: + tl.store( + paged_mqa_ctx_lens_2d + + req_row * paged_mqa_ctx_lens_stride_0 + + draft_off * paged_mqa_ctx_lens_stride_1, + base_seq + next_n, + mask=mask_e, + ) + return + + num_col_blocks = tl.cdiv(max_seqlen_k, BLOCK_N) + page_pid = pid - 1 + out_row = page_pid // num_col_blocks + col_block = page_pid - out_row * num_col_blocks + offs_n = col_block * BLOCK_N + tl.arange(0, BLOCK_N) + mask = (out_row < expanded_size) & (offs_n < max_seqlen_k) + + req_row = out_row // next_n + req_idx = tl.load( + req_pool_indices + req_row * req_pool_indices_stride, + mask=out_row < expanded_size, + other=0, + ) + vals = tl.load( + req_to_token + req_idx * req_to_token_stride_0 + offs_n * req_to_token_stride_1, + mask=mask, + other=0, + ).to(tl.int32) + tl.store( + page_table_1 + out_row * page_table_stride_0 + offs_n * page_table_stride_1, + vals, + mask=mask, + ) + + if HAS_REAL_PAGE_TABLE: + real_mask = mask & ((offs_n % real_page_size) == 0) + real_cols = offs_n // real_page_size + tl.store( + real_page_table + + out_row * real_page_table_stride_0 + + real_cols * real_page_table_stride_1, + vals // real_page_size, + mask=real_mask, + ) + + +def fused_dsa_target_verify_metadata( + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + req_to_token: torch.Tensor, + cache_seqlens: torch.Tensor, + cu_seqlens_k: torch.Tensor, + page_table_1: torch.Tensor, + seqlens_expanded: torch.Tensor, + dsa_cache_seqlens: torch.Tensor, + dsa_cu_seqlens_k: torch.Tensor, + real_page_table: torch.Tensor, + bs: int, + max_seqlen_k: int, + dsa_index_topk: int, + real_page_size: int, + next_n: int, + paged_mqa_ctx_lens_2d: torch.Tensor = None, +) -> None: + assert seq_lens.is_cuda + assert req_pool_indices.is_cuda + assert req_to_token.is_cuda + assert cache_seqlens.is_cuda + assert cu_seqlens_k.is_cuda + assert page_table_1.is_cuda + assert seqlens_expanded.is_cuda + assert dsa_cache_seqlens.is_cuda + assert dsa_cu_seqlens_k.is_cuda + + if bs == 0: + cu_seqlens_k[:1].zero_() + dsa_cu_seqlens_k[:1].zero_() + return + assert next_n > 0 + + has_real_page_table = real_page_size > 1 + if has_real_page_table: + assert real_page_table is not None + assert real_page_table.is_cuda + else: + real_page_table = page_table_1 + + has_paged_mqa_ctx_lens = paged_mqa_ctx_lens_2d is not None + if has_paged_mqa_ctx_lens: + assert paged_mqa_ctx_lens_2d.is_cuda + assert paged_mqa_ctx_lens_2d.dtype == torch.int32 + assert paged_mqa_ctx_lens_2d.dim() == 2 + assert paged_mqa_ctx_lens_2d.size(0) == bs + assert paged_mqa_ctx_lens_2d.size(1) == next_n + else: + paged_mqa_ctx_lens_2d = page_table_1 + + expanded_size = bs * next_n + block_bs = triton.next_power_of_2(bs) + block_expanded = triton.next_power_of_2(expanded_size) + block_n = 128 + num_col_blocks = triton.cdiv(max_seqlen_k, block_n) + grid = (1 + expanded_size * num_col_blocks,) + + _fused_dsa_target_verify_metadata_kernel[grid]( + seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + seqlens_expanded, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + paged_mqa_ctx_lens_2d, + seq_lens.stride(0), + req_pool_indices.stride(0), + req_to_token.stride(0), + req_to_token.stride(1), + page_table_1.stride(0), + page_table_1.stride(1), + real_page_table.stride(0) if has_real_page_table else 0, + real_page_table.stride(1) if has_real_page_table else 0, + paged_mqa_ctx_lens_2d.stride(0) if has_paged_mqa_ctx_lens else 0, + paged_mqa_ctx_lens_2d.stride(1) if has_paged_mqa_ctx_lens else 0, + bs, + max_seqlen_k, + dsa_index_topk, + real_page_size, + next_n, + has_real_page_table, + has_paged_mqa_ctx_lens, + BLOCK_BS=block_bs, + BLOCK_EXPANDED=block_expanded, + BLOCK_N=block_n, + ) + + +@triton.jit( + do_not_specialize=[ + "page_table_stride_0", + "real_page_table_stride_0", + "total_len", + "max_seqlen_k", + ] +) +def _fused_dsa_draft_extend_metadata_kernel( + seq_lens, + extend_seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + seqlens_expanded, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + seq_lens_stride: tl.constexpr, + extend_seq_lens_stride: tl.constexpr, + req_pool_indices_stride: tl.constexpr, + req_to_token_stride_0: tl.constexpr, + req_to_token_stride_1: tl.constexpr, + page_table_stride_0, + page_table_stride_1: tl.constexpr, + real_page_table_stride_0, + real_page_table_stride_1: tl.constexpr, + bs: tl.constexpr, + total_len, + max_seqlen_k, + dsa_index_topk: tl.constexpr, + real_page_size: tl.constexpr, + HAS_REAL_PAGE_TABLE: tl.constexpr, + STATIC_EXTEND_LEN: tl.constexpr, + BLOCK_BS: tl.constexpr, + BLOCK_EXPANDED: tl.constexpr, + BLOCK_ROWS: tl.constexpr, + BLOCK_N: tl.constexpr, +): + pid = tl.program_id(0) + + if pid == 0: + offs_b = tl.arange(0, BLOCK_BS) + mask_b = offs_b < bs + seq = tl.load(seq_lens + offs_b * seq_lens_stride, mask=mask_b, other=0) + cache_seq = seq.to(tl.int32) + cu = tl.cumsum(cache_seq, 0) + + tl.store(cache_seqlens + offs_b, cache_seq, mask=mask_b) + tl.store(cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(cu_seqlens_k + 1 + offs_b, cu, mask=mask_b) + + offs_e = tl.arange(0, BLOCK_EXPANDED) + mask_e = offs_e < total_len + if STATIC_EXTEND_LEN: + static_qo_len = tl.load(extend_seq_lens).to(tl.int32) + req_row = offs_e // static_qo_len + local_off = offs_e - req_row * static_qo_len + qo_len_for_row = tl.zeros((BLOCK_EXPANDED,), tl.int32) + static_qo_len + else: + req_row = tl.full((BLOCK_EXPANDED,), 0, tl.int32) + local_off = tl.full((BLOCK_EXPANDED,), 0, tl.int32) + qo_len_for_row = tl.full((BLOCK_EXPANDED,), 1, tl.int32) + prefix = tl.full((), 0, tl.int32) + + for i in tl.range(0, bs): + qo_len = tl.load(extend_seq_lens + i * extend_seq_lens_stride).to( + tl.int32 + ) + in_row = (offs_e >= prefix) & (offs_e < prefix + qo_len) + req_row = tl.where(in_row, i, req_row) + local_off = tl.where(in_row, offs_e - prefix, local_off) + qo_len_for_row = tl.where(in_row, qo_len, qo_len_for_row) + prefix += qo_len + + base_seq = tl.load( + seq_lens + req_row * seq_lens_stride, + mask=mask_e, + other=0, + ).to(tl.int32) + expanded_seq = base_seq - qo_len_for_row + local_off + 1 + expanded_seq = tl.where(mask_e, expanded_seq, 0) + dsa_seq = tl.minimum(expanded_seq, dsa_index_topk) + dsa_cu = tl.cumsum(dsa_seq, 0) + + tl.store(seqlens_expanded + offs_e, expanded_seq, mask=mask_e) + tl.store(dsa_cache_seqlens + offs_e, dsa_seq, mask=mask_e) + tl.store(dsa_cu_seqlens_k, tl.full((), 0, tl.int32)) + tl.store(dsa_cu_seqlens_k + 1 + offs_e, dsa_cu, mask=mask_e) + return + + num_col_blocks = tl.cdiv(max_seqlen_k, BLOCK_N) + page_pid = pid - 1 + req_row = page_pid // num_col_blocks + col_block = page_pid - req_row * num_col_blocks + offs_n = col_block * BLOCK_N + tl.arange(0, BLOCK_N) + + qo_len = tl.load( + extend_seq_lens + req_row * extend_seq_lens_stride, + mask=req_row < bs, + other=0, + ).to(tl.int32) + if STATIC_EXTEND_LEN: + prefix = req_row * qo_len + else: + prefix = tl.full((), 0, tl.int32) + for i in tl.range(0, bs): + prev_qo_len = tl.load(extend_seq_lens + i * extend_seq_lens_stride).to( + tl.int32 + ) + prefix += tl.where(i < req_row, prev_qo_len, 0) + offs_r = tl.arange(0, BLOCK_ROWS) + out_rows = prefix + offs_r + row_mask = (req_row < bs) & (offs_r < qo_len) & (out_rows < total_len) + col_mask = offs_n < max_seqlen_k + has_rows = (req_row < bs) & (qo_len > 0) + mask = row_mask[:, None] & col_mask[None, :] + + req_idx = tl.load( + req_pool_indices + req_row * req_pool_indices_stride, + mask=has_rows, + other=0, + ) + vals = tl.load( + req_to_token + req_idx * req_to_token_stride_0 + offs_n * req_to_token_stride_1, + mask=col_mask & has_rows, + other=0, + ).to(tl.int32) + tl.store( + page_table_1 + + out_rows[:, None] * page_table_stride_0 + + offs_n[None, :] * page_table_stride_1, + vals[None, :], + mask=mask, + ) + + if HAS_REAL_PAGE_TABLE: + real_mask = mask & ((offs_n[None, :] % real_page_size) == 0) + real_cols = offs_n // real_page_size + tl.store( + real_page_table + + out_rows[:, None] * real_page_table_stride_0 + + real_cols[None, :] * real_page_table_stride_1, + (vals // real_page_size)[None, :], + mask=real_mask, + ) + + +def fused_dsa_draft_extend_metadata( + seq_lens: torch.Tensor, + extend_seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + req_to_token: torch.Tensor, + cache_seqlens: torch.Tensor, + cu_seqlens_k: torch.Tensor, + page_table_1: torch.Tensor, + seqlens_expanded: torch.Tensor, + dsa_cache_seqlens: torch.Tensor, + dsa_cu_seqlens_k: torch.Tensor, + real_page_table: torch.Tensor, + bs: int, + total_len: int, + max_seqlen_k: int, + dsa_index_topk: int, + real_page_size: int, + max_extend_len: int, + max_total_len: int, + static_extend_len: bool = False, +) -> None: + assert seq_lens.is_cuda + assert extend_seq_lens.is_cuda + assert req_pool_indices.is_cuda + assert req_to_token.is_cuda + assert cache_seqlens.is_cuda + assert cu_seqlens_k.is_cuda + assert page_table_1.is_cuda + assert seqlens_expanded.is_cuda + assert dsa_cache_seqlens.is_cuda + assert dsa_cu_seqlens_k.is_cuda + + if bs == 0: + cu_seqlens_k[:1].zero_() + dsa_cu_seqlens_k[:1].zero_() + return + if total_len == 0: + cache = seq_lens.to(torch.int32) + cache_seqlens.copy_(cache) + cu_seqlens_k[:1].zero_() + cu_seqlens_k[1 : bs + 1].copy_(torch.cumsum(cache, dim=0, dtype=torch.int32)) + dsa_cu_seqlens_k[:1].zero_() + return + assert total_len <= max_total_len + # Caller-owned graph metadata guarantees each request accepts at most + # max_extend_len tokens. Avoid checking extend_seq_lens.max() here because + # that would sync in the replay hot path. + assert max_extend_len > 0 + assert total_len <= bs * max_extend_len + + has_real_page_table = real_page_size > 1 + if has_real_page_table: + assert real_page_table is not None + assert real_page_table.is_cuda + else: + real_page_table = page_table_1 + + block_bs = triton.next_power_of_2(bs) + block_expanded = triton.next_power_of_2(max_total_len) + block_rows = triton.next_power_of_2(max_extend_len) + block_n = 128 + num_col_blocks = triton.cdiv(max_seqlen_k, block_n) + grid = (1 + bs * num_col_blocks,) + + _fused_dsa_draft_extend_metadata_kernel[grid]( + seq_lens, + extend_seq_lens, + req_pool_indices, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table_1, + seqlens_expanded, + dsa_cache_seqlens, + dsa_cu_seqlens_k, + real_page_table, + seq_lens.stride(0), + extend_seq_lens.stride(0), + req_pool_indices.stride(0), + req_to_token.stride(0), + req_to_token.stride(1), + page_table_1.stride(0), + page_table_1.stride(1), + real_page_table.stride(0) if has_real_page_table else 0, + real_page_table.stride(1) if has_real_page_table else 0, + bs, + total_len, + max_seqlen_k, + dsa_index_topk, + real_page_size, + has_real_page_table, + static_extend_len, + BLOCK_BS=block_bs, + BLOCK_EXPANDED=block_expanded, + BLOCK_ROWS=block_rows, + BLOCK_N=block_n, + ) diff --git a/test/registered/kernels/test_dsa_metadata.py b/test/registered/kernels/test_dsa_metadata.py new file mode 100644 index 000000000..e9592e3ff --- /dev/null +++ b/test/registered/kernels/test_dsa_metadata.py @@ -0,0 +1,474 @@ +import unittest + +import torch + +from sglang.srt.layers.attention.triton_ops.dsa_metadata import ( + fused_dsa_decode_metadata, + fused_dsa_draft_extend_metadata, + fused_dsa_target_verify_metadata, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=15, stage="base-b", runner_config="1-gpu-large") + + +def _cu_seqlens(seqlens: torch.Tensor) -> torch.Tensor: + out = torch.empty(seqlens.numel() + 1, dtype=torch.int32, device=seqlens.device) + out[:1].zero_() + out[1:].copy_(torch.cumsum(seqlens.to(torch.int32), dim=0, dtype=torch.int32)) + return out + + +def _dsa_seqlens(seqlens: torch.Tensor, topk: int) -> torch.Tensor: + return torch.minimum( + seqlens.to(torch.int32), torch.tensor(topk, device=seqlens.device) + ) + + +def _real_page_table(page_table_1: torch.Tensor, real_page_size: int) -> torch.Tensor: + if real_page_size == 1: + return page_table_1 + return page_table_1[:, ::real_page_size] // real_page_size + + +def _make_req_to_token( + pool_size: int, max_len: int, device: torch.device +) -> torch.Tensor: + # Row-dependent values catch accidental row reuse, while monotonic columns make + # real-page-table checks easy to reason about. + cols = torch.arange(max_len, dtype=torch.int32, device=device) + rows = torch.arange(pool_size, dtype=torch.int32, device=device).view(-1, 1) + return rows * (max_len + 17) + cols + + +def _assert_equal(actual: torch.Tensor, expected: torch.Tensor, name: str) -> None: + torch.testing.assert_close(actual, expected, rtol=0, atol=0, msg=name) + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA is required for this test.") +class TestDSAMetadataKernels(CustomTestCase): + def setUp(self): + super().setUp() + self.device = torch.device("cuda") + + def _check_decode( + self, + seq_lens_values, + *, + max_len: int, + dsa_index_topk: int, + real_page_size: int, + ): + bs = len(seq_lens_values) + pool_size = max(bs + 3, 8) + seq_lens = torch.tensor(seq_lens_values, dtype=torch.int64, device=self.device) + req_pool_indices = torch.arange(bs, dtype=torch.int64, device=self.device) * 2 + req_to_token = _make_req_to_token(pool_size * 2, max_len, self.device) + + cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device) + cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device) + page_table_1 = torch.empty((bs, max_len), dtype=torch.int32, device=self.device) + dsa_cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device) + dsa_cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device) + real_page_table = ( + torch.empty( + (bs, (max_len + real_page_size - 1) // real_page_size), + dtype=torch.int32, + device=self.device, + ) + if real_page_size > 1 + else page_table_1 + ) + + fused_dsa_decode_metadata( + seq_lens=seq_lens, + req_pool_indices=req_pool_indices, + req_to_token=req_to_token, + cache_seqlens=cache_seqlens, + cu_seqlens_k=cu_seqlens_k, + page_table_1=page_table_1, + dsa_cache_seqlens=dsa_cache_seqlens, + dsa_cu_seqlens_k=dsa_cu_seqlens_k, + real_page_table=real_page_table, + bs=bs, + max_len=max_len, + dsa_index_topk=dsa_index_topk, + real_page_size=real_page_size, + ) + + expected_cache = seq_lens.to(torch.int32) + expected_page_table = req_to_token[req_pool_indices, :max_len].contiguous() + expected_dsa = _dsa_seqlens(expected_cache, dsa_index_topk) + + _assert_equal(cache_seqlens, expected_cache, "decode cache_seqlens") + _assert_equal(cu_seqlens_k, _cu_seqlens(expected_cache), "decode cu_seqlens_k") + _assert_equal(page_table_1, expected_page_table, "decode page_table_1") + _assert_equal(dsa_cache_seqlens, expected_dsa, "decode dsa_cache_seqlens") + _assert_equal( + dsa_cu_seqlens_k, _cu_seqlens(expected_dsa), "decode dsa_cu_seqlens_k" + ) + if real_page_size > 1: + _assert_equal( + real_page_table, + _real_page_table(expected_page_table, real_page_size), + "decode real_page_table", + ) + + def _check_target_verify( + self, + seq_lens_values, + *, + max_seqlen_k: int, + dsa_index_topk: int, + real_page_size: int, + next_n: int, + fill_ctx_lens: bool, + ): + bs = len(seq_lens_values) + expanded_size = bs * next_n + pool_size = max(bs + 5, 8) + seq_lens = torch.tensor(seq_lens_values, dtype=torch.int64, device=self.device) + req_pool_indices = torch.arange(bs, dtype=torch.int64, device=self.device) + 1 + req_to_token = _make_req_to_token(pool_size + 2, max_seqlen_k, self.device) + + cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device) + cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device) + page_table_1 = torch.empty( + (expanded_size, max_seqlen_k), dtype=torch.int32, device=self.device + ) + seqlens_expanded = torch.empty( + expanded_size, dtype=torch.int32, device=self.device + ) + dsa_cache_seqlens = torch.empty( + expanded_size, dtype=torch.int32, device=self.device + ) + dsa_cu_seqlens_k = torch.empty( + expanded_size + 1, dtype=torch.int32, device=self.device + ) + real_page_table = ( + torch.empty( + ( + expanded_size, + (max_seqlen_k + real_page_size - 1) // real_page_size, + ), + dtype=torch.int32, + device=self.device, + ) + if real_page_size > 1 + else page_table_1 + ) + paged_mqa_ctx_lens_2d = ( + torch.empty((bs, next_n), dtype=torch.int32, device=self.device) + if fill_ctx_lens + else None + ) + + fused_dsa_target_verify_metadata( + seq_lens=seq_lens, + req_pool_indices=req_pool_indices, + req_to_token=req_to_token, + cache_seqlens=cache_seqlens, + cu_seqlens_k=cu_seqlens_k, + page_table_1=page_table_1, + seqlens_expanded=seqlens_expanded, + dsa_cache_seqlens=dsa_cache_seqlens, + dsa_cu_seqlens_k=dsa_cu_seqlens_k, + real_page_table=real_page_table, + bs=bs, + max_seqlen_k=max_seqlen_k, + dsa_index_topk=dsa_index_topk, + real_page_size=real_page_size, + next_n=next_n, + paged_mqa_ctx_lens_2d=paged_mqa_ctx_lens_2d, + ) + + expected_cache = (seq_lens + next_n).to(torch.int32) + base_page_table = req_to_token[req_pool_indices, :max_seqlen_k].contiguous() + expected_page_table = torch.repeat_interleave( + base_page_table, repeats=next_n, dim=0 + ).contiguous() + draft_offsets = torch.arange(next_n, dtype=torch.int32, device=self.device) + expected_expanded = seq_lens.to(torch.int32).view(-1, 1) + draft_offsets + 1 + expected_expanded = expected_expanded.reshape(-1).contiguous() + expected_dsa = _dsa_seqlens(expected_expanded, dsa_index_topk) + + _assert_equal(cache_seqlens, expected_cache, "target cache_seqlens") + _assert_equal(cu_seqlens_k, _cu_seqlens(expected_cache), "target cu_seqlens_k") + _assert_equal(page_table_1, expected_page_table, "target page_table_1") + _assert_equal(seqlens_expanded, expected_expanded, "target seqlens_expanded") + _assert_equal(dsa_cache_seqlens, expected_dsa, "target dsa_cache_seqlens") + _assert_equal( + dsa_cu_seqlens_k, _cu_seqlens(expected_dsa), "target dsa_cu_seqlens_k" + ) + if real_page_size > 1: + _assert_equal( + real_page_table, + _real_page_table(expected_page_table, real_page_size), + "target real_page_table", + ) + if fill_ctx_lens: + expected_ctx = expected_cache.view(bs, 1).expand(bs, next_n).contiguous() + _assert_equal( + paged_mqa_ctx_lens_2d, expected_ctx, "target paged_mqa_ctx_lens_2d" + ) + + def _check_draft_extend( + self, + seq_lens_values, + extend_seq_lens_values, + *, + max_seqlen_k: int, + dsa_index_topk: int, + real_page_size: int, + max_extend_len: int, + max_total_len: int, + static_extend_len: bool, + ): + bs = len(seq_lens_values) + total_len = sum(extend_seq_lens_values) + pool_size = max(bs + 4, 8) + seq_lens = torch.tensor(seq_lens_values, dtype=torch.int64, device=self.device) + extend_seq_lens = torch.tensor( + extend_seq_lens_values, dtype=torch.int32, device=self.device + ) + req_pool_indices = torch.arange(bs, dtype=torch.int64, device=self.device) + 2 + req_to_token = _make_req_to_token(pool_size + 4, max_seqlen_k, self.device) + + cache_seqlens = torch.empty(bs, dtype=torch.int32, device=self.device) + cu_seqlens_k = torch.empty(bs + 1, dtype=torch.int32, device=self.device) + page_table_1 = torch.empty( + (max_total_len, max_seqlen_k), dtype=torch.int32, device=self.device + ) + seqlens_expanded = torch.empty( + max_total_len, dtype=torch.int32, device=self.device + ) + dsa_cache_seqlens = torch.empty( + max_total_len, dtype=torch.int32, device=self.device + ) + dsa_cu_seqlens_k = torch.empty( + max_total_len + 1, dtype=torch.int32, device=self.device + ) + real_page_table = ( + torch.empty( + (max_total_len, (max_seqlen_k + real_page_size - 1) // real_page_size), + dtype=torch.int32, + device=self.device, + ) + if real_page_size > 1 + else page_table_1 + ) + + fused_dsa_draft_extend_metadata( + seq_lens=seq_lens, + extend_seq_lens=extend_seq_lens, + req_pool_indices=req_pool_indices, + req_to_token=req_to_token, + cache_seqlens=cache_seqlens, + cu_seqlens_k=cu_seqlens_k, + page_table_1=page_table_1, + seqlens_expanded=seqlens_expanded, + dsa_cache_seqlens=dsa_cache_seqlens, + dsa_cu_seqlens_k=dsa_cu_seqlens_k, + real_page_table=real_page_table, + bs=bs, + total_len=total_len, + max_seqlen_k=max_seqlen_k, + dsa_index_topk=dsa_index_topk, + real_page_size=real_page_size, + max_extend_len=max_extend_len, + max_total_len=max_total_len, + static_extend_len=static_extend_len, + ) + + expected_cache = seq_lens.to(torch.int32) + base_page_table = req_to_token[req_pool_indices, :max_seqlen_k].contiguous() + expected_page_table = torch.repeat_interleave( + base_page_table, repeats=extend_seq_lens, dim=0 + ).contiguous() + expanded_parts = [] + for seq_len, qo_len in zip(seq_lens, extend_seq_lens, strict=True): + expanded_parts.append( + torch.arange( + seq_len.item() - qo_len.item() + 1, + seq_len.item() + 1, + dtype=torch.int32, + device=self.device, + ) + ) + expected_expanded = ( + torch.cat(expanded_parts, dim=0) + if expanded_parts + else torch.empty(0, dtype=torch.int32, device=self.device) + ) + expected_dsa = _dsa_seqlens(expected_expanded, dsa_index_topk) + + _assert_equal(cache_seqlens, expected_cache, "draft cache_seqlens") + _assert_equal(cu_seqlens_k, _cu_seqlens(expected_cache), "draft cu_seqlens_k") + _assert_equal( + page_table_1[:total_len], expected_page_table, "draft page_table_1" + ) + _assert_equal( + seqlens_expanded[:total_len], expected_expanded, "draft seqlens_expanded" + ) + _assert_equal( + dsa_cache_seqlens[:total_len], expected_dsa, "draft dsa_cache_seqlens" + ) + _assert_equal( + dsa_cu_seqlens_k[: total_len + 1], + _cu_seqlens(expected_dsa), + "draft dsa_cu_seqlens_k", + ) + if real_page_size > 1: + _assert_equal( + real_page_table[:total_len], + _real_page_table(expected_page_table, real_page_size), + "draft real_page_table", + ) + + def test_decode_matches_eager_reference(self): + for real_page_size in (1, 64): + with self.subTest(real_page_size=real_page_size): + self._check_decode( + [1, 7, 65, 513], + max_len=769, + dsa_index_topk=64, + real_page_size=real_page_size, + ) + + def test_target_verify_matches_eager_reference(self): + for real_page_size, fill_ctx_lens in ((1, False), (64, True)): + with self.subTest( + real_page_size=real_page_size, fill_ctx_lens=fill_ctx_lens + ): + self._check_target_verify( + [5, 63, 128], + max_seqlen_k=257, + dsa_index_topk=64, + real_page_size=real_page_size, + next_n=4, + fill_ctx_lens=fill_ctx_lens, + ) + + def test_draft_extend_static_width_matches_eager_reference(self): + self._check_draft_extend( + [16, 31, 80], + [4, 4, 4], + max_seqlen_k=193, + dsa_index_topk=64, + real_page_size=1, + max_extend_len=4, + max_total_len=12, + static_extend_len=True, + ) + + def test_draft_extend_variable_width_defensive_path(self): + # The production draft-extend-v2 replay path uses static_extend_len=True. + # Keep this case to guard the generic variable-width kernel branch. + self._check_draft_extend( + [12, 31, 80], + [3, 5, 2], + max_seqlen_k=193, + dsa_index_topk=64, + real_page_size=64, + max_extend_len=5, + max_total_len=10, + static_extend_len=False, + ) + + def test_draft_extend_partial_fill(self): + self._check_draft_extend( + [12, 31, 80], + [3, 5, 2], + max_seqlen_k=193, + dsa_index_topk=64, + real_page_size=64, + max_extend_len=5, + max_total_len=16, + static_extend_len=False, + ) + + def test_empty_batch(self): + self._check_decode( + [], + max_len=8, + dsa_index_topk=64, + real_page_size=64, + ) + self._check_target_verify( + [], + max_seqlen_k=8, + dsa_index_topk=64, + real_page_size=64, + next_n=4, + fill_ctx_lens=True, + ) + self._check_draft_extend( + [], + [], + max_seqlen_k=8, + dsa_index_topk=64, + real_page_size=64, + max_extend_len=1, + max_total_len=0, + static_extend_len=True, + ) + + def test_large_shape_coverage(self): + max_len = 1_000_003 + self._check_decode( + [1_000_000, 999_983], + max_len=max_len, + dsa_index_topk=4096, + real_page_size=64, + ) + self._check_target_verify( + [1_000_000], + max_seqlen_k=max_len, + dsa_index_topk=4096, + real_page_size=64, + next_n=2, + fill_ctx_lens=True, + ) + self._check_draft_extend( + [1_000_000], + [4], + max_seqlen_k=max_len, + dsa_index_topk=4096, + real_page_size=64, + max_extend_len=4, + max_total_len=4, + static_extend_len=True, + ) + + def test_large_batch_coverage(self): + bs = 16 * 1024 + seq_lens = (torch.arange(bs, dtype=torch.int64) % 257 + 1).tolist() + self._check_decode( + seq_lens, + max_len=1, + dsa_index_topk=64, + real_page_size=1, + ) + self._check_target_verify( + seq_lens, + max_seqlen_k=1, + dsa_index_topk=64, + real_page_size=1, + next_n=1, + fill_ctx_lens=False, + ) + self._check_draft_extend( + seq_lens, + [1] * bs, + max_seqlen_k=1, + dsa_index_topk=64, + real_page_size=1, + max_extend_len=1, + max_total_len=bs, + static_extend_len=True, + ) + + +if __name__ == "__main__": + unittest.main()