diff --git a/python/sglang/jit_kernel/dsv4/topk.py b/python/sglang/jit_kernel/dsv4/topk.py index 273d2434c..fdaedfd75 100644 --- a/python/sglang/jit_kernel/dsv4/topk.py +++ b/python/sglang/jit_kernel/dsv4/topk.py @@ -67,6 +67,15 @@ _PLAN_METADATA_INTS_PER_BATCH = 2 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() bs = seq_lens.shape[0] metadata = seq_lens.new_empty(bs + 1, _PLAN_METADATA_INTS_PER_BATCH) @@ -83,6 +92,17 @@ def topk_transform_512_v2( metadata: torch.Tensor, out_raw_indices: Optional[torch.Tensor] = 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.topk_transform( scores, diff --git a/python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py b/python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py index c33aad1cd..f9db4b88e 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py @@ -84,22 +84,22 @@ class DSATopKBackend(Enum): row_starts: Optional[torch.Tensor] = None, batch_idx_list: Optional[List[int]] = None, force_unfused_topk: bool = False, - allow_topk_v2: bool = True, ) -> torch.Tensor: if not envs.SGLANG_DSA_FUSE_TOPK.get() or force_unfused_topk: 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, - # which fuses top-k selection and the page-table transform in one launch and - # consumes the indexer's own page_size>=1 table directly, so no page_size=1 - # table is materialized. Shared by DeepSeek-V3.2 and GLM DSA. This is a - # deterministic dispatch on the work shape, not a best-effort attempt: the - # fused-decode CUDA graph drops the page_size=1 table for 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. + # Decode-shaped PAGED top-k (plain decode AND spec verify / draft-extend, + # whose expanded rows match the same shape) routes to the DeepSeek-V4 top-k + # v2 JIT kernel, which fuses top-k selection and the page-table transform in + # one launch and consumes the indexer's own page_size>=1 table directly, so + # no page_size=1 table is materialized. Shared by DeepSeek-V3.2 and GLM DSA. + # This is a deterministic dispatch on the work shape, not a best-effort + # attempt: the fused-decode CUDA graph drops the page_size=1 table for + # 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 ( - allow_topk_v2 - and envs.SGLANG_OPT_USE_TOPK_V2.get() + envs.SGLANG_OPT_USE_TOPK_V2.get() and topk_transform_method == TopkTransformMethod.PAGED and row_starts 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 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. + + ``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.srt.model_executor.forward_context import get_token_to_kv_pool diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 1f367c0b8..1adf96cac 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -251,12 +251,6 @@ class DSAIndexerMetadata(BaseIndexerMetadata): paged_mqa_schedule_metadata: Optional[torch.Tensor] = None paged_mqa_ctx_lens_2d: Optional[torch.Tensor] = None 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: return self.attn_metadata.cache_seqlens_int32 @@ -326,7 +320,6 @@ class DSAIndexerMetadata(BaseIndexerMetadata): row_starts=ks, batch_idx_list=batch_idx_list, 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 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( attn_metadata=self.forward_metadata, 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_ctx_lens_2d=self.forward_metadata.paged_mqa_ctx_lens_2d, force_unfused_topk=force_unfused, - allow_topk_v2=allow_topk_v2, ) def _compute_flashmla_metadata(self, cache_seqlens: torch.Tensor, seq_len_q: int): diff --git a/python/sglang/srt/layers/attention/triton_ops/dsa_metadata.py b/python/sglang/srt/layers/attention/triton_ops/dsa_metadata.py index 633b485fb..caa8ed027 100644 --- a/python/sglang/srt/layers/attention/triton_ops/dsa_metadata.py +++ b/python/sglang/srt/layers/attention/triton_ops/dsa_metadata.py @@ -495,7 +495,14 @@ def _fused_dsa_draft_extend_metadata_kernel( mask=mask_e, other=0, ).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 = tl.maximum(expanded_seq, 0) 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) diff --git a/python/sglang/srt/layers/attention/triton_ops/pad.py b/python/sglang/srt/layers/attention/triton_ops/pad.py index dc3fb1590..bd211f62c 100644 --- a/python/sglang/srt/layers/attention/triton_ops/pad.py +++ b/python/sglang/srt/layers/attention/triton_ops/pad.py @@ -347,7 +347,12 @@ def seqlens_expand_kernel( offs = tl.arange(0, BLOCK) 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) diff --git a/test/registered/kernels/test_dsa_indexer.py b/test/registered/kernels/test_dsa_indexer.py index c82cd8dea..90d4bd9c9 100644 --- a/test/registered/kernels/test_dsa_indexer.py +++ b/test/registered/kernels/test_dsa_indexer.py @@ -640,6 +640,12 @@ class TestDSAIndexer(CustomTestCase): ) ).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( page_size=1, cache_seqlens_int32=seq_lens_expanded.clone(), @@ -654,6 +660,7 @@ class TestDSAIndexer(CustomTestCase): dsa_cu_seqlens_k=dsa_cu_seqlens_k, dsa_extend_seq_lens_list=seq_lens_expanded.cpu().tolist(), dsa_seqlens_expanded=seq_lens_expanded, + topk_v2_plan=plan_topk_v2(seq_lens_expanded), topk_indices_offset=( topk_indices_offset if topk_transform_method == TopkTransformMethod.RAGGED