[AMD][DSV4] fix: skip compressed-KV metadata on the draft worker in the HIP radix backend (#40205)

This commit is contained in:
amd-danli103
2026-09-19 12:13:49 -07:00
committed by GitHub
parent 9e5a62a767
commit 2305242f51
@@ -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,