[DSA] Optimize DSA CUDA graph replay metadata generation (#29499)

This commit is contained in:
Mohammad Miadh Angkad
2026-06-29 19:53:23 -07:00
committed by GitHub
parent 25b6051c70
commit 3a72d02415
5 changed files with 1478 additions and 123 deletions
+1
View File
@@ -650,6 +650,7 @@ class Envs:
SGLANG_DSA_MQA_LOGITS_FREE_MEM_FRACTION = EnvFloat(0.2) SGLANG_DSA_MQA_LOGITS_FREE_MEM_FRACTION = EnvFloat(0.2)
SGLANG_ENABLE_PCG_DSV2_DUAL_STREAM = EnvBool(False) SGLANG_ENABLE_PCG_DSV2_DUAL_STREAM = EnvBool(False)
SGLANG_USE_FUSED_METADATA_COPY = EnvBool(True) SGLANG_USE_FUSED_METADATA_COPY = EnvBool(True)
SGLANG_DSA_USE_FUSED_METADATA_GENERATION = EnvBool(True)
SGLANG_DSA_TOPK_BROADCAST = EnvBool(False) SGLANG_DSA_TOPK_BROADCAST = EnvBool(False)
SGLANG_DISABLE_DSA_INDEXER_FUSION = EnvBool(False) SGLANG_DISABLE_DSA_INDEXER_FUSION = EnvBool(False)
@@ -11,11 +11,20 @@ from typing import TYPE_CHECKING, Optional
import torch import torch
from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsa.utils import compute_dsa_seqlens 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: if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardMode 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 @dataclass
class PrecomputedMetadata: class PrecomputedMetadata:
@@ -116,6 +125,69 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
"""Precompute metadata for normal decode mode.""" """Precompute metadata for normal decode mode."""
max_len = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1] 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 # Convert to int32 and compute cumsum
cache_seqlens = seq_lens.to(torch.int32) cache_seqlens = seq_lens.to(torch.int32)
cu_seqlens_k = compute_cu_seqlens(cache_seqlens) cu_seqlens_k = compute_cu_seqlens(cache_seqlens)
@@ -169,9 +241,84 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
seq_lens_cpu: torch.Tensor, seq_lens_cpu: torch.Tensor,
) -> PrecomputedMetadata: ) -> PrecomputedMetadata:
"""Precompute metadata for target verify mode.""" """Precompute metadata for target verify mode."""
max_seqlen_k = int( max_seqlen_k = self.decode_cuda_graph_metadata[bs].page_table_1.shape[1]
seq_lens_cpu.max().item() + self.speculative_num_draft_tokens 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 with draft tokens
cache_seqlens = (seq_lens + self.speculative_num_draft_tokens).to(torch.int32) 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 page_indices, repeats=self.speculative_num_draft_tokens, dim=0
).contiguous() ).contiguous()
# Generate expanded seqlens # Generate expanded seqlens on device. seq_lens_cpu is optional for DSA
extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * bs # CUDA graph replay, so this fallback must not require a host mirror.
seqlens_int32_cpu = [ extend_seq_lens = torch.full(
self.speculative_num_draft_tokens + kv_len (bs,),
for kv_len in seq_lens_cpu.tolist() self.speculative_num_draft_tokens,
] dtype=torch.int32,
seqlens_expanded = torch.cat( device=self.device,
[ )
torch.arange( seqlens_expanded = seqlens_expand_triton(
kv_len - qo_len + 1, extend_seq_lens,
kv_len + 1, cache_seqlens,
dtype=torch.int32, bs * self.speculative_num_draft_tokens,
device=self.device, self.speculative_num_draft_tokens,
)
for qo_len, kv_len in zip(
extend_seq_lens_cpu,
seqlens_int32_cpu,
strict=True,
)
]
) )
# Compute DSA seqlens # Compute DSA seqlens
+219 -100
View File
@@ -117,6 +117,9 @@ global_workspace_buffer = None
# Control whether to use fused metadata copy kernel for cuda graph replay (default: enabled) # 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 # 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_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) @dataclass(frozen=True)
@@ -588,6 +591,19 @@ class DeepseekSparseAttnBackend(
return _to_2d_context_lens(seqlens_expanded, batch_size) return _to_2d_context_lens(seqlens_expanded, batch_size)
return _to_2d_context_lens(cache_seqlens_int32, 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: def _get_fused_topk_page_table(self, topk_indices: torch.Tensor) -> torch.Tensor:
if ( if (
self.dsa_topk_backend.is_sgl_kernel() self.dsa_topk_backend.is_sgl_kernel()
@@ -1220,58 +1236,133 @@ class DeepseekSparseAttnBackend(
# Normal Decode # Normal Decode
metadata: DSAMetadata = self.decode_cuda_graph_metadata[bs] 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(): if forward_mode.is_decode_or_idle():
# Normal Decode # Normal Decode
max_len = metadata.page_table_1.shape[1] max_len = metadata.page_table_1.shape[1]
cache_seqlens = seq_lens.to(torch.int32) if _USE_FUSED_METADATA_GENERATION and is_cuda():
metadata.cache_seqlens_int32.copy_(cache_seqlens) from sglang.srt.layers.attention.triton_ops.dsa_metadata import (
metadata.cu_seqlens_k[1:].copy_( fused_dsa_decode_metadata,
torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32) )
)
page_indices = self.req_to_token[req_pool_indices, :max_len] fused_dsa_decode_metadata(
metadata.page_table_1[:, :max_len].copy_(page_indices) seq_lens=seq_lens,
dsa_cache_seqlens = compute_dsa_seqlens( req_pool_indices=req_pool_indices,
cache_seqlens, dsa_index_topk=self.dsa_index_topk req_to_token=self.req_to_token,
) cache_seqlens=metadata.cache_seqlens_int32,
metadata.dsa_cache_seqlens_int32.copy_(dsa_cache_seqlens) cu_seqlens_k=metadata.cu_seqlens_k,
seqlens_expanded = cache_seqlens 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(): elif forward_mode.is_target_verify():
max_seqlen_k = metadata.page_table_1.shape[1] max_seqlen_k = metadata.page_table_1.shape[1]
cache_seqlens = (seq_lens + self.speculative_num_draft_tokens).to( if _USE_FUSED_METADATA_GENERATION and is_cuda():
torch.int32 from sglang.srt.layers.attention.triton_ops.dsa_metadata import (
) fused_dsa_target_verify_metadata,
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 (num_draft_tokens) on-device; paged_mqa_ctx_lens_2d = None
# torch.tensor(list, device=cuda) does a pageable H2D copy that if (
# blocks the host on the whole queued stream. self.speculative_num_draft_tokens >= 2
extend_seq_lens = torch.full( and is_sm100_supported()
(bs,), and metadata.paged_mqa_ctx_lens_2d is not None
self.speculative_num_draft_tokens, and metadata.paged_mqa_ctx_lens_2d.dim() == 2
dtype=torch.int32, and metadata.paged_mqa_ctx_lens_2d.size(0) == bs
device=self.device, and metadata.paged_mqa_ctx_lens_2d.size(1)
) == self.speculative_num_draft_tokens
seqlens_expanded = seqlens_expand_triton( ):
extend_seq_lens, paged_mqa_ctx_lens_2d = metadata.paged_mqa_ctx_lens_2d
cache_seqlens,
self.speculative_num_draft_tokens * bs, fused_dsa_target_verify_metadata(
self.speculative_num_draft_tokens, seq_lens=seq_lens,
) req_pool_indices=req_pool_indices,
metadata.dsa_seqlens_expanded.copy_(seqlens_expanded) req_to_token=self.req_to_token,
dsa_cache_seqlens = compute_dsa_seqlens( cache_seqlens=metadata.cache_seqlens_int32,
seqlens_expanded, self.dsa_index_topk cu_seqlens_k=metadata.cu_seqlens_k,
) page_table_1=metadata.page_table_1,
metadata.dsa_cache_seqlens_int32.copy_(dsa_cache_seqlens) 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(): elif forward_mode.is_draft_extend_v2():
# V2 draft-extend processes the full padded tree width # V2 draft-extend processes the full padded tree width
# (speculative_num_draft_tokens) per req -- a static shape, like # (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 # the per-req accept length is handled downstream by output
# selection, not by reshaping the page table here. # selection, not by reshaping the page table here.
max_seqlen_k = metadata.page_table_1.shape[1] max_seqlen_k = metadata.page_table_1.shape[1]
cache_seqlens = seq_lens.to(torch.int32) total_extend_len = self.speculative_num_draft_tokens * bs
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)
# See target-verify note: fill on-device to avoid the blocking # See target-verify note: fill on-device to avoid the blocking
# pageable H2D from torch.tensor(list, device=cuda). # pageable H2D from torch.tensor(list, device=cuda).
@@ -1300,17 +1381,65 @@ class DeepseekSparseAttnBackend(
dtype=torch.int32, dtype=torch.int32,
device=self.device, device=self.device,
) )
seqlens_expanded = seqlens_expand_triton(
extend_seq_lens, if _USE_FUSED_METADATA_GENERATION and is_cuda():
cache_seqlens, from sglang.srt.layers.attention.triton_ops.dsa_metadata import (
self.speculative_num_draft_tokens * bs, fused_dsa_draft_extend_metadata,
self.speculative_num_draft_tokens, )
)
metadata.dsa_seqlens_expanded.copy_(seqlens_expanded) fused_dsa_draft_extend_metadata(
dsa_cache_seqlens = compute_dsa_seqlens( seq_lens=seq_lens,
seqlens_expanded, self.dsa_index_topk extend_seq_lens=extend_seq_lens,
) req_pool_indices=req_pool_indices,
metadata.dsa_cache_seqlens_int32.copy_(dsa_cache_seqlens) 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. # Update DeepGEMM paged MQA schedule metadata outside the captured graph.
if is_cuda() and ( if is_cuda() and (
@@ -1322,26 +1451,22 @@ class DeepseekSparseAttnBackend(
schedule_seqlens_expanded = metadata.dsa_seqlens_expanded schedule_seqlens_expanded = metadata.dsa_seqlens_expanded
else: else:
schedule_seqlens_expanded = seqlens_expanded schedule_seqlens_expanded = seqlens_expanded
seqlens_32_2d = self._build_paged_mqa_schedule_2d_ctx_lens( if target_verify_ctx_lens_written:
forward_mode, seqlens_32_2d = metadata.paged_mqa_ctx_lens_2d
metadata.cache_seqlens_int32, else:
schedule_seqlens_expanded, seqlens_32_2d = self._build_paged_mqa_schedule_2d_ctx_lens(
bs, forward_mode,
) metadata.cache_seqlens_int32,
new_schedule = deep_gemm.get_paged_mqa_logits_metadata( schedule_seqlens_expanded,
seqlens_32_2d, 64, deep_gemm.get_num_sms() bs,
)
if metadata.paged_mqa_schedule_metadata is None:
object.__setattr__(
metadata, "paged_mqa_schedule_metadata", new_schedule
) )
else: self._refresh_paged_mqa_schedule_metadata(metadata, seqlens_32_2d)
metadata.paged_mqa_schedule_metadata.copy_(new_schedule)
# `copy_` preserves the buffer's data_ptr that the captured graph captured. # `copy_` preserves the buffer's data_ptr that the captured graph captured.
if metadata.paged_mqa_ctx_lens_2d is None: if not target_verify_ctx_lens_written:
object.__setattr__(metadata, "paged_mqa_ctx_lens_2d", seqlens_32_2d) if metadata.paged_mqa_ctx_lens_2d is None:
else: object.__setattr__(metadata, "paged_mqa_ctx_lens_2d", seqlens_32_2d)
metadata.paged_mqa_ctx_lens_2d.copy_(seqlens_32_2d) else:
metadata.paged_mqa_ctx_lens_2d.copy_(seqlens_32_2d)
seqlens_expanded_size = seqlens_expanded.shape[0] seqlens_expanded_size = seqlens_expanded.shape[0]
assert ( assert (
metadata.dsa_cache_seqlens_int32 is not None metadata.dsa_cache_seqlens_int32 is not None
@@ -1349,17 +1474,19 @@ class DeepseekSparseAttnBackend(
and self.dsa_index_topk is not None and self.dsa_index_topk is not None
) )
metadata.dsa_cu_seqlens_k[1 : 1 + seqlens_expanded_size].copy_( if not used_fused_metadata_generation:
torch.cumsum(dsa_cache_seqlens, dim=0, dtype=torch.int32) 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 # NOTE(dark): (dsa-) cu_seqlens_q is always arange, no need to copy
assert self.real_page_size == metadata.page_size assert self.real_page_size == metadata.page_size
if self.real_page_size > 1: if self.real_page_size > 1:
real_table = self._transform_table_1_to_real(page_indices) if not used_fused_metadata_generation:
new_rows = real_table.shape[0] real_table = self._transform_table_1_to_real(page_indices)
new_cols = real_table.shape[1] new_rows = real_table.shape[0]
metadata.real_page_table[:new_rows, :new_cols].copy_(real_table) new_cols = real_table.shape[1]
metadata.real_page_table[:new_rows, :new_cols].copy_(real_table)
else: else:
assert metadata.real_page_table is metadata.page_table_1 assert metadata.real_page_table is metadata.page_table_1
@@ -1528,15 +1655,7 @@ class DeepseekSparseAttnBackend(
metadata.dsa_seqlens_expanded, metadata.dsa_seqlens_expanded,
bs, bs,
) )
new_schedule = deep_gemm.get_paged_mqa_logits_metadata( self._refresh_paged_mqa_schedule_metadata(metadata, seqlens_32_2d)
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: if metadata.paged_mqa_ctx_lens_2d is None:
object.__setattr__(metadata, "paged_mqa_ctx_lens_2d", seqlens_32_2d) object.__setattr__(metadata, "paged_mqa_ctx_lens_2d", seqlens_32_2d)
else: else:
@@ -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,
)
@@ -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()