[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_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)
@@ -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
+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)
# 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:
@@ -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,
)