From 6b2c730bf793984c39f7f07b3c074ca05b059b00 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Mon, 22 Jun 2026 22:39:21 +0800 Subject: [PATCH] [codex] Fix DSA indexer in prefill piecewise CUDA graph (#28644) --- .../srt/layers/attention/dsa/dsa_indexer.py | 35 +++++++++++++++++++ 1 file changed, 35 insertions(+) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index c242a1c99..55f1905e3 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -99,6 +99,34 @@ if TYPE_CHECKING: DUAL_STREAM_TOKEN_THRESHOLD = 1024 if _is_cuda else 0 +def _uses_dsa_attention_backend(forward_batch: ForwardBatch) -> bool: + attn_backend = get_attn_backend() + server_args = get_global_server_args() + prefill_backend, decode_backend = server_args.get_attention_backends() + prefill_backend = ( + getattr(attn_backend, "prefill_attention_backend_str", None) or prefill_backend + ) + decode_backend = ( + getattr(attn_backend, "decode_attention_backend_str", None) or decode_backend + ) + + if forward_batch.forward_mode.is_decode_or_idle(): + backend_name = decode_backend + elif ( + forward_batch.forward_mode.is_target_verify() + or forward_batch.forward_mode.is_draft_extend_v2() + ): + backend_name = ( + decode_backend + if server_args.speculative_attention_mode == "decode" + else prefill_backend + ) + else: + backend_name = prefill_backend + + return backend_name in ("dsa", "nsa") + + if _is_cuda: from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.utils.custom_op import register_custom_op @@ -120,6 +148,10 @@ if _is_cuda: forward_batch = get_tc_piecewise_forward_context().forward_batch indexer = get_tc_piecewise_forward_context().dsa_indexers[layer_id] metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch) + assert metadata is not None, ( + "DSA piecewise CUDA graph requires indexer metadata from the DSA " + "attention backend" + ) # slice off padding from piecewise CUDA graph extend_num_tokens = forward_batch.extend_num_tokens @@ -1575,6 +1607,9 @@ class Indexer(MultiPlatformOp): not enable_dual_stream ), "Internal error: piecewise CUDA graph should not be enabled with dual stream" + if not _uses_dsa_attention_backend(forward_batch): + return None + topk_result = torch.full( (q_fp8.shape[0], self.index_topk), -1,