From 2305242f514da95eba3c411de0517c4d79ffff94 Mon Sep 17 00:00:00 2001 From: amd-danli103 Date: Sun, 20 Sep 2026 03:13:49 +0800 Subject: [PATCH] [AMD][DSV4] fix: skip compressed-KV metadata on the draft worker in the HIP radix backend (#40205) --- .../deepseek_v4_backend_hip_radix.py | 97 +++++++++++++------ 1 file changed, 67 insertions(+), 30 deletions(-) diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py index 12fbf76e4..373f637b6 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -586,6 +586,10 @@ class DeepseekV4HipRadixBackend( self.is_draft_worker = getattr(model_runner, "is_draft_worker", False) self.is_dspark = model_runner.spec_algorithm.is_dspark() self.is_dspark_draft = self.is_draft_worker and self.is_dspark + # Draft layers are all COMPRESS_RATIO_NEXTN_LAYER (0), so the draft pool + # has neither compressed kv nor an indexer pool. Settled here rather than + # per forward: it cannot change after construction. + self.need_compress = not self.is_draft_worker self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens if self.is_dspark_draft: assert self.speculative_num_draft_tokens is not None @@ -867,7 +871,8 @@ class DeepseekV4HipRadixBackend( num_tokens=num_tokens, extend_seq_lens=extend_seq_lens, extend_seq_lens_cpu=extend_seq_lens_cpu, - need_compress=True, + # guarded inside init_forward_metadata_prefill, not here + need_compress=self.need_compress, use_prefill_cuda_graph=use_prefill_cuda_graph, compress_gpu_plan=ragged_layout is not None, extend_start_loc=extend_start_loc, @@ -898,13 +903,14 @@ class DeepseekV4HipRadixBackend( bs, num_draft_tokens, seq_lens, req_pool_indices ) ) + need_compress = self.need_compress core_attn_metadata = self.make_core_attn_metadata( req_to_token=self.req_to_token, req_pool_indices_repeated=req_pool_indices_repeated, seq_lens_casual=seq_lens_casual, max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, out_loc=out_cache_loc, - need_compress=True, + need_compress=need_compress, ) # extend_seq_lens is uniform here (seq_lens already carries the draft # block, so the minimum above cannot trim it), hence an exact token count. @@ -919,20 +925,27 @@ class DeepseekV4HipRadixBackend( self._attach_unified_kv_decode_streams( core_attn_metadata, req_pool_indices_repeated ) - indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata) - create = functools.partial( - create_paged_compressor_data, - is_prefill=True, - token_to_kv_pool=self.token_to_kv_pool, - req_to_token=self.req_to_token, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - extend_lens=extend_seq_lens, - seq_lens_cpu=None, - extend_lens_cpu=None, - use_prefill_cuda_graph=True, - num_q_tokens=num_draft_tokens * bs, + indexer_metadata = ( + self.init_forward_metadata_indexer(core_attn_metadata) + if need_compress + else None ) + if not need_compress: + create = _create_dummy_paged_compress_data + else: + create = functools.partial( + create_paged_compressor_data, + is_prefill=True, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + extend_lens=extend_seq_lens, + seq_lens_cpu=None, + extend_lens_cpu=None, + use_prefill_cuda_graph=True, + num_q_tokens=num_draft_tokens * bs, + ) return DSV4Metadata( core_attn_metadata, indexer_metadata, @@ -951,26 +964,34 @@ class DeepseekV4HipRadixBackend( # the accepted-prefix lengths unchanged across the captured loop. seq_lens = seq_lens + self.speculative_step_id + 1 + need_compress = self.need_compress core_attn_metadata = self.make_core_attn_metadata( req_to_token=self.req_to_token, req_pool_indices_repeated=req_pool_indices, seq_lens_casual=seq_lens, max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, out_loc=out_cache_loc, - need_compress=True, + need_compress=need_compress, ) self._attach_unified_kv_decode_streams(core_attn_metadata, req_pool_indices) - indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata) - - create = functools.partial( - create_paged_compressor_data, - is_prefill=False, - token_to_kv_pool=self.token_to_kv_pool, - req_to_token=self.req_to_token, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, + indexer_metadata = ( + self.init_forward_metadata_indexer(core_attn_metadata) + if need_compress + else None ) + if not need_compress: + create = _create_dummy_paged_compress_data + else: + create = functools.partial( + create_paged_compressor_data, + is_prefill=False, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + ) + return DSV4Metadata( core_attn_metadata, indexer_metadata, @@ -1465,6 +1486,22 @@ class DeepseekV4HipRadixBackend( state_slot = state_slot[:N] if core.unified is None: core.unified = UnifiedKvMetadata() + swa_len = core.swa_topk_lengths + hca_len = core.c128_topk_lengths_raw + csa_len = core.c4_sparse_topk_lengths_raw + hca_page_indices = core.c128_page_indices + c4_page_indices = core.c4_sparse_page_indices + # A draft pool carries no c4/c128 metadata, but its verify store still + # reads the swa half below, so build these with an empty compressed tail + # rather than skipping the call. + if hca_len is None: + hca_len = torch.zeros_like(swa_len) + if csa_len is None: + csa_len = torch.zeros_like(swa_len) + if hca_page_indices is None: + hca_page_indices = torch.empty( + (N, 0), dtype=torch.int32, device=swa_len.device + ) ( core.unified.swa_indices, core.unified.swa_indptr, @@ -1475,11 +1512,11 @@ class DeepseekV4HipRadixBackend( ) = runtime.build_decode_streams( state_slot=state_slot, positions=core.positions_casual, - swa_len=core.swa_topk_lengths, - hca_len=core.c128_topk_lengths_raw, - csa_len=core.c4_sparse_topk_lengths_raw, - hca_page_indices=core.c128_page_indices, - csa_width=core.c4_sparse_page_indices.shape[1], + swa_len=swa_len, + hca_len=hca_len, + csa_len=csa_len, + hca_page_indices=hca_page_indices, + csa_width=0 if c4_page_indices is None else c4_page_indices.shape[1], win=pool.unified_swa_window, ring_stride=pool.unified_swa_ring_size, swa_pages=pool.unified_swa_pages,