[DSA] Re-enable fused top-k v2 for MTP: clamp padded-row seq_lens to >= 0 (#30378)
Co-authored-by: ziyi.xu <ziyi.xu@radixark.ai> Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
ziyi.xu
Claude Fable 5
parent
6875df3378
commit
bbc537035a
@@ -67,6 +67,15 @@ _PLAN_METADATA_INTS_PER_BATCH = 2
|
|||||||
|
|
||||||
|
|
||||||
def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Tensor:
|
def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Tensor:
|
||||||
|
"""Preprocess the per-batch routing plan for :func:`topk_transform_512_v2`.
|
||||||
|
|
||||||
|
IMPORTANT: every entry of ``seq_lens`` must be NON-NEGATIVE. The device
|
||||||
|
kernel reads the int32 buffer as ``uint32_t``, so a negative length (e.g.
|
||||||
|
-4 from a DP-padded / idle-companion row) reinterprets as ~4e9, poisons
|
||||||
|
the plan, and drives the transform kernel into an illegal memory access.
|
||||||
|
Producers of padded rows must clamp their lengths to 0 (0 selects the
|
||||||
|
trivial all-(-1) output path, which is safe).
|
||||||
|
"""
|
||||||
module = _jit_topk_v2_module()
|
module = _jit_topk_v2_module()
|
||||||
bs = seq_lens.shape[0]
|
bs = seq_lens.shape[0]
|
||||||
metadata = seq_lens.new_empty(bs + 1, _PLAN_METADATA_INTS_PER_BATCH)
|
metadata = seq_lens.new_empty(bs + 1, _PLAN_METADATA_INTS_PER_BATCH)
|
||||||
@@ -83,6 +92,17 @@ def topk_transform_512_v2(
|
|||||||
metadata: torch.Tensor,
|
metadata: torch.Tensor,
|
||||||
out_raw_indices: Optional[torch.Tensor] = None,
|
out_raw_indices: Optional[torch.Tensor] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
"""Fused top-k + page-table transform (DeepSeek-V4 top-k v2 kernel).
|
||||||
|
|
||||||
|
IMPORTANT: every entry of ``seq_lens`` must be NON-NEGATIVE, and
|
||||||
|
``metadata`` must come from :func:`plan_topk_v2` over the same ``seq_lens``
|
||||||
|
values. The kernel reads lengths as ``uint32_t``: a negative entry
|
||||||
|
reinterprets as a ~4e9-token sequence, sending the row down the cluster
|
||||||
|
path over garbage scores and crashing with an illegal memory access
|
||||||
|
(GLM 5.2 MTP DP-idle companion rows hit exactly this). A length of 0 is
|
||||||
|
the valid way to express "no tokens": the row takes the trivial path and
|
||||||
|
the output is all -1.
|
||||||
|
"""
|
||||||
module = _jit_topk_v2_module()
|
module = _jit_topk_v2_module()
|
||||||
module.topk_transform(
|
module.topk_transform(
|
||||||
scores,
|
scores,
|
||||||
|
|||||||
@@ -84,22 +84,22 @@ class DSATopKBackend(Enum):
|
|||||||
row_starts: Optional[torch.Tensor] = None,
|
row_starts: Optional[torch.Tensor] = None,
|
||||||
batch_idx_list: Optional[List[int]] = None,
|
batch_idx_list: Optional[List[int]] = None,
|
||||||
force_unfused_topk: bool = False,
|
force_unfused_topk: bool = False,
|
||||||
allow_topk_v2: bool = True,
|
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if not envs.SGLANG_DSA_FUSE_TOPK.get() or force_unfused_topk:
|
if not envs.SGLANG_DSA_FUSE_TOPK.get() or force_unfused_topk:
|
||||||
return self.topk_func(logits, lengths, topk, row_starts=row_starts)
|
return self.topk_func(logits, lengths, topk, row_starts=row_starts)
|
||||||
|
|
||||||
# Decode-shaped PAGED top-k routes to the DeepSeek-V4 top-k v2 JIT kernel,
|
# Decode-shaped PAGED top-k (plain decode AND spec verify / draft-extend,
|
||||||
# which fuses top-k selection and the page-table transform in one launch and
|
# whose expanded rows match the same shape) routes to the DeepSeek-V4 top-k
|
||||||
# consumes the indexer's own page_size>=1 table directly, so no page_size=1
|
# v2 JIT kernel, which fuses top-k selection and the page-table transform in
|
||||||
# table is materialized. Shared by DeepSeek-V3.2 and GLM DSA. This is a
|
# one launch and consumes the indexer's own page_size>=1 table directly, so
|
||||||
# deterministic dispatch on the work shape, not a best-effort attempt: the
|
# no page_size=1 table is materialized. Shared by DeepSeek-V3.2 and GLM DSA.
|
||||||
# fused-decode CUDA graph drops the page_size=1 table for exactly this case
|
# This is a deterministic dispatch on the work shape, not a best-effort
|
||||||
# (see dsa_drop_wide_page_table), so once the shape matches we commit to v2
|
# attempt: the fused-decode CUDA graph drops the page_size=1 table for
|
||||||
# and never silently fall back to the legacy page_size=1 path from here.
|
# exactly this case (see dsa_drop_wide_page_table), so once the shape
|
||||||
|
# matches we commit to v2 and never silently fall back to the legacy
|
||||||
|
# page_size=1 path from here.
|
||||||
if (
|
if (
|
||||||
allow_topk_v2
|
envs.SGLANG_OPT_USE_TOPK_V2.get()
|
||||||
and envs.SGLANG_OPT_USE_TOPK_V2.get()
|
|
||||||
and topk_transform_method == TopkTransformMethod.PAGED
|
and topk_transform_method == TopkTransformMethod.PAGED
|
||||||
and row_starts is None
|
and row_starts is None
|
||||||
and batch_idx_list is None
|
and batch_idx_list is None
|
||||||
@@ -260,6 +260,12 @@ def _topk_transform_v2_paged(
|
|||||||
than fall back to the slow legacy path (which may not even have a page_size=1
|
than fall back to the slow legacy path (which may not even have a page_size=1
|
||||||
table to fall back to) or silently paper over bad input (padding, recomputing
|
table to fall back to) or silently paper over bad input (padding, recomputing
|
||||||
the plan) at the cost of the performance this path exists to deliver.
|
the plan) at the cost of the performance this path exists to deliver.
|
||||||
|
|
||||||
|
``lengths`` entries must be NON-NEGATIVE: the kernel reads them as
|
||||||
|
``uint32_t``, so a negative row length (DP-padded / idle-companion rows)
|
||||||
|
reinterprets as ~4e9 tokens and illegal-addresses. Metadata producers clamp
|
||||||
|
padded rows to 0 (see ``fused_dsa_draft_extend_metadata`` /
|
||||||
|
``seqlens_expand_kernel``); 0 takes the trivial all-(-1) output path.
|
||||||
"""
|
"""
|
||||||
from sglang.jit_kernel.dsv4.topk import topk_transform_512_v2
|
from sglang.jit_kernel.dsv4.topk import topk_transform_512_v2
|
||||||
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
|
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
|
||||||
|
|||||||
@@ -251,12 +251,6 @@ class DSAIndexerMetadata(BaseIndexerMetadata):
|
|||||||
paged_mqa_schedule_metadata: Optional[torch.Tensor] = None
|
paged_mqa_schedule_metadata: Optional[torch.Tensor] = None
|
||||||
paged_mqa_ctx_lens_2d: Optional[torch.Tensor] = None
|
paged_mqa_ctx_lens_2d: Optional[torch.Tensor] = None
|
||||||
force_unfused_topk: bool = False
|
force_unfused_topk: bool = False
|
||||||
# Whether the fused top-k v2 kernel may be used for this forward. Disabled for
|
|
||||||
# spec verify / draft-extend: the v2 small-batch path illegal-addresses under
|
|
||||||
# those multi-query-per-request CUDA graphs (GLM 5.2 MTP), and it is only
|
|
||||||
# e2e-validated for single-query decode. TODO(dsa-topk-v2): re-enable once the
|
|
||||||
# small-batch kernel is fixed; see the crash analysis in the PR.
|
|
||||||
allow_topk_v2: bool = True
|
|
||||||
|
|
||||||
def get_seqlens_int32(self) -> torch.Tensor:
|
def get_seqlens_int32(self) -> torch.Tensor:
|
||||||
return self.attn_metadata.cache_seqlens_int32
|
return self.attn_metadata.cache_seqlens_int32
|
||||||
@@ -326,7 +320,6 @@ class DSAIndexerMetadata(BaseIndexerMetadata):
|
|||||||
row_starts=ks,
|
row_starts=ks,
|
||||||
batch_idx_list=batch_idx_list,
|
batch_idx_list=batch_idx_list,
|
||||||
force_unfused_topk=self.force_unfused_topk,
|
force_unfused_topk=self.force_unfused_topk,
|
||||||
allow_topk_v2=self.allow_topk_v2,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -2849,14 +2842,6 @@ class DeepseekSparseAttnBackend(
|
|||||||
self.hisparse_coordinator is not None
|
self.hisparse_coordinator is not None
|
||||||
and forward_batch.forward_mode.is_decode_or_idle()
|
and forward_batch.forward_mode.is_decode_or_idle()
|
||||||
)
|
)
|
||||||
# TEMP(dsa-topk-v2): the fused v2 small-batch path illegal-addresses under
|
|
||||||
# spec verify / draft-extend CUDA graphs (GLM 5.2 MTP). Restrict v2 to the
|
|
||||||
# single-query decode shape it is e2e-validated on; spec falls back to the
|
|
||||||
# legacy transform (page_table_1 is present for spec, not dropped).
|
|
||||||
allow_topk_v2 = not (
|
|
||||||
forward_batch.forward_mode.is_target_verify()
|
|
||||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
|
||||||
)
|
|
||||||
return DSAIndexerMetadata(
|
return DSAIndexerMetadata(
|
||||||
attn_metadata=self.forward_metadata,
|
attn_metadata=self.forward_metadata,
|
||||||
topk_transform_method=self.get_topk_transform_method(
|
topk_transform_method=self.get_topk_transform_method(
|
||||||
@@ -2866,7 +2851,6 @@ class DeepseekSparseAttnBackend(
|
|||||||
paged_mqa_schedule_metadata=self.forward_metadata.paged_mqa_schedule_metadata,
|
paged_mqa_schedule_metadata=self.forward_metadata.paged_mqa_schedule_metadata,
|
||||||
paged_mqa_ctx_lens_2d=self.forward_metadata.paged_mqa_ctx_lens_2d,
|
paged_mqa_ctx_lens_2d=self.forward_metadata.paged_mqa_ctx_lens_2d,
|
||||||
force_unfused_topk=force_unfused,
|
force_unfused_topk=force_unfused,
|
||||||
allow_topk_v2=allow_topk_v2,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _compute_flashmla_metadata(self, cache_seqlens: torch.Tensor, seq_len_q: int):
|
def _compute_flashmla_metadata(self, cache_seqlens: torch.Tensor, seq_len_q: int):
|
||||||
|
|||||||
@@ -495,7 +495,14 @@ def _fused_dsa_draft_extend_metadata_kernel(
|
|||||||
mask=mask_e,
|
mask=mask_e,
|
||||||
other=0,
|
other=0,
|
||||||
).to(tl.int32)
|
).to(tl.int32)
|
||||||
|
# Clamp to >= 0: DP-padded / idle-companion rows carry the CUDA-graph
|
||||||
|
# seq_len fill value (1), which is smaller than qo_len, so the raw
|
||||||
|
# per-row visible kv length goes negative. Consumers treat these
|
||||||
|
# lengths as unsigned (the top-k v2 kernel reads them as uint32), so a
|
||||||
|
# negative row becomes a ~4e9-token length and an illegal memory
|
||||||
|
# access. 0 keeps padded rows on the trivial all-(-1) output path.
|
||||||
expanded_seq = base_seq - qo_len_for_row + local_off + 1
|
expanded_seq = base_seq - qo_len_for_row + local_off + 1
|
||||||
|
expanded_seq = tl.maximum(expanded_seq, 0)
|
||||||
expanded_seq = tl.where(mask_e, expanded_seq, 0)
|
expanded_seq = tl.where(mask_e, expanded_seq, 0)
|
||||||
dsa_seq = tl.minimum(expanded_seq, dsa_index_topk)
|
dsa_seq = tl.minimum(expanded_seq, dsa_index_topk)
|
||||||
dsa_cu = tl.cumsum(dsa_seq, 0)
|
dsa_cu = tl.cumsum(dsa_seq, 0)
|
||||||
|
|||||||
@@ -347,7 +347,12 @@ def seqlens_expand_kernel(
|
|||||||
offs = tl.arange(0, BLOCK)
|
offs = tl.arange(0, BLOCK)
|
||||||
mask = offs < qo_len
|
mask = offs < qo_len
|
||||||
|
|
||||||
values = start + offs
|
# Clamp to >= 0: rows with kv_len < qo_len (DP-padded / idle-companion
|
||||||
|
# rows whose kv is the CUDA-graph fill value) would otherwise produce
|
||||||
|
# negative lengths, which unsigned consumers (e.g. the top-k v2 kernel,
|
||||||
|
# which reads lengths as uint32) turn into ~4e9-token lengths and an
|
||||||
|
# illegal memory access.
|
||||||
|
values = tl.maximum(start + offs, 0)
|
||||||
tl.store(output_ptr + out_offset + offs, values, mask=mask)
|
tl.store(output_ptr + out_offset + offs, values, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -640,6 +640,12 @@ class TestDSAIndexer(CustomTestCase):
|
|||||||
)
|
)
|
||||||
).contiguous()
|
).contiguous()
|
||||||
|
|
||||||
|
# The fused v2 PAGED dispatch requires the per-forward plan to be
|
||||||
|
# preprocessed alongside the metadata (it asserts rather than silently
|
||||||
|
# recomputing it) -- mirror what init_forward_metadata /
|
||||||
|
# _build_forward_metadata_cuda_graph do.
|
||||||
|
from sglang.jit_kernel.dsv4.topk import plan_topk_v2
|
||||||
|
|
||||||
attn_metadata = DSAMetadata(
|
attn_metadata = DSAMetadata(
|
||||||
page_size=1,
|
page_size=1,
|
||||||
cache_seqlens_int32=seq_lens_expanded.clone(),
|
cache_seqlens_int32=seq_lens_expanded.clone(),
|
||||||
@@ -654,6 +660,7 @@ class TestDSAIndexer(CustomTestCase):
|
|||||||
dsa_cu_seqlens_k=dsa_cu_seqlens_k,
|
dsa_cu_seqlens_k=dsa_cu_seqlens_k,
|
||||||
dsa_extend_seq_lens_list=seq_lens_expanded.cpu().tolist(),
|
dsa_extend_seq_lens_list=seq_lens_expanded.cpu().tolist(),
|
||||||
dsa_seqlens_expanded=seq_lens_expanded,
|
dsa_seqlens_expanded=seq_lens_expanded,
|
||||||
|
topk_v2_plan=plan_topk_v2(seq_lens_expanded),
|
||||||
topk_indices_offset=(
|
topk_indices_offset=(
|
||||||
topk_indices_offset
|
topk_indices_offset
|
||||||
if topk_transform_method == TopkTransformMethod.RAGGED
|
if topk_transform_method == TopkTransformMethod.RAGGED
|
||||||
|
|||||||
Reference in New Issue
Block a user