[DSA] Optimize DSA CUDA graph replay metadata generation (#29499)
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user