From 504570f4250d8ff045263f7c8f8f8c2ee058beaa Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Thu, 9 Jul 2026 16:39:47 -0700 Subject: [PATCH] [Refactor] Make DeepSeek-V4 attention backend tolerate an absent CPU seq_lens mirror (#30695) --- .../layers/attention/deepseek_v4_backend.py | 43 +++++++++++-------- 1 file changed, 26 insertions(+), 17 deletions(-) diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index 1aafb3dfb..3d09014f3 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -457,6 +457,8 @@ class DeepseekV4AttnBackend( ): use_captured_forward_metadata_for_breakable_cuda_graph: bool = True + needs_cpu_seq_lens: bool = False + def __init__( self, model_runner: ModelRunner, @@ -509,6 +511,11 @@ class DeepseekV4AttnBackend( DSV4RawDecodeMetadata, ] = None self.online_c128_mtp = OnlineC128MTPController(self) + # Draft-extend and online-c128 verify metadata are host-planned, so + # spec runs keep the relay publish (the mirror only exists under + # spec-v2; without spec the flag has no consumer either way). + if model_runner.server_args.speculative_algorithm is not None: + self.needs_cpu_seq_lens = True self.sparse_prefill_workspace = SparsePrefillWorkspace(self.device) def _move_to_device(self, x: List[int]) -> torch.Tensor: @@ -706,9 +713,7 @@ class DeepseekV4AttnBackend( if envs.SGLANG_PREP_IN_CUDA_GRAPH.get(): assert out_cache_loc is not None seq_lens_cpu_list = ( - seq_lens.detach().cpu().tolist() - if seq_lens_cpu is None - else seq_lens_cpu.tolist() + seq_lens_cpu.tolist() if seq_lens_cpu is not None else None ) if not hasattr(self, "extend_seq_lens_buffer"): self.extend_seq_lens_buffer = torch.tensor( @@ -984,14 +989,13 @@ class DeepseekV4AttnBackend( ) out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device) - assert seq_lens_cpu is not None seq_lens = seq_lens[:bs] - seq_lens_cpu = seq_lens_cpu[:bs] req_pool_indices = req_pool_indices[:bs] - - actual_max_seq_len = seq_lens_cpu.max().item() chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE - assert actual_max_seq_len <= chosen_max_seq_len + if seq_lens_cpu is not None: + seq_lens_cpu = seq_lens_cpu[:bs] + actual_max_seq_len = seq_lens_cpu.max().item() + assert actual_max_seq_len <= chosen_max_seq_len if bucket == _GraphBucket.DECODE_OR_IDLE: assert out_cache_loc is not None @@ -1060,11 +1064,14 @@ class DeepseekV4AttnBackend( mode="constant", value=0, ) + draft_extend_seq_lens_cpu = ( + seq_lens_cpu.tolist() if seq_lens_cpu is not None else seq_lens.tolist() + ) temp_metadata = self.init_forward_metadata_draft_extend( max_seq_len=chosen_max_seq_len, req_pool_indices=req_pool_indices, seq_lens=seq_lens, - seq_lens_cpu=seq_lens_cpu.tolist(), + seq_lens_cpu=draft_extend_seq_lens_cpu, num_tokens_per_bs=num_tokens_per_bs, out_cache_loc=out_cache_loc, use_prefill_cuda_graph=True, @@ -1112,14 +1119,14 @@ class DeepseekV4AttnBackend( assert self.req_to_token_pool.req_to_token is self.req_to_token assert self.swa_page_size % SWA_WINDOW == 0 and self.page_size % 128 == 0 - assert seq_lens_cpu is not None if max_seq_len_override is None: max_seq_len_override = getattr(forward_batch, "max_seq_len_override", None) - max_seq_len = ( - int(seq_lens_cpu.max().item()) - if max_seq_len_override is None - else max_seq_len_override - ) + if max_seq_len_override is not None: + max_seq_len = max_seq_len_override + elif seq_lens_cpu is not None: + max_seq_len = int(seq_lens_cpu.max().item()) + else: + max_seq_len = int(seq_lens.max().item()) verify_bs = _get_target_verify_bs(forward_batch) online_c128_state_slot_offset = self.online_c128_mtp.prepare_forward( logical_forward_mode, @@ -1159,16 +1166,18 @@ class DeepseekV4AttnBackend( extend_seq_lens = forward_batch.extend_seq_lens assert ( seq_lens is not None - and seq_lens_cpu is not None and extend_seq_lens is not None and extend_seq_lens_cpu is not None ) is_draft = forward_batch.forward_mode.is_draft_extend_v2() + prefill_seq_lens_cpu = ( + seq_lens_cpu.tolist() if seq_lens_cpu is not None else seq_lens.tolist() + ) metadata = self.init_forward_metadata_prefill( max_seq_len=max_seq_len, req_pool_indices=req_pool_indices, seq_lens=seq_lens, - seq_lens_cpu=seq_lens_cpu.tolist(), + seq_lens_cpu=prefill_seq_lens_cpu, out_cache_loc=forward_batch.out_cache_loc, num_tokens=sum(extend_seq_lens_cpu), extend_seq_lens=extend_seq_lens,