[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_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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
Reference in New Issue
Block a user