From 8cd7ed3872b663298e5aa63f0188c29259ff8c6d Mon Sep 17 00:00:00 2001 From: huangweizhe1 Date: Wed, 15 Jul 2026 10:55:30 +0800 Subject: [PATCH] bugfix: redirect draft-extend placeholder KV write to padding block. When dp_enabled=true or use_chunked_prefill=true, draft extend always adds a prev_token row. If prev_token_id < 0 (first decode after prefill or all-draft-rejected), the placeholder row overwrites a valid KV cache entry at position-1, permanently degrading speculative acceptance rate. Fix: detect placeholder rows and redirect their new_cache_slots to block 0 (reserved padding block), preserving the correct KV cache content while maintaining DP token-count alignment. Verified on DeepSeek-V3.2-w8a8 (DP=2, num_speculative_tokens=3): - seed1 1req: 41.97% vs 39.86% (nofix), +2.11% - seed3 1req: 16.67% vs 16.67% (identical, short request) - seed4 1req: 57.02% vs 54.70% (nofix), +2.32% - 4req seed1: 48.47% vs 48.33% (nofix), drafted count identical (687) --- xllm/core/runtime/mtp_worker_impl.cpp | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/xllm/core/runtime/mtp_worker_impl.cpp b/xllm/core/runtime/mtp_worker_impl.cpp index 856ea9a7c4..209f67b979 100644 --- a/xllm/core/runtime/mtp_worker_impl.cpp +++ b/xllm/core/runtime/mtp_worker_impl.cpp @@ -1526,11 +1526,16 @@ void MTPWorkerImpl::prepare_draft_extend_inputs( if (use_chunked_prefill) { int32_t prev_token_id = state.prev_token_id; torch::Tensor prev_embedding = state.prev_embedding; - if (prev_token_id < 0) { + const bool prev_is_placeholder = prev_token_id < 0; + if (prev_is_placeholder) { prev_token_id = current_token_id >= 0 ? current_token_id : 0; prev_embedding = torch::Tensor(); } add_row(prev_token_id, /*position_offset=*/-1, prev_embedding); + if (prev_is_placeholder) { + // Redirect to padding block 0 to avoid overwriting correct KV cache. + buf.out_new_cache_slots.back() = 0; + } add_row(state.token_id, /*position_offset=*/0, state.embedding); specBuilder::append_seq_len_by_layout(buf.out_q_seq_lens, 2); const int32_t kv_len = specBuilder::calc_kv_len( @@ -1547,11 +1552,16 @@ void MTPWorkerImpl::prepare_draft_extend_inputs( int32_t prev_token_id = state.prev_token_id; int32_t prev_position_offset = -1; torch::Tensor prev_embedding = state.prev_embedding; - if (prev_token_id < 0) { + const bool prev_is_placeholder = prev_token_id < 0; + if (prev_is_placeholder) { prev_token_id = state.token_id; prev_embedding = torch::Tensor(); } add_row(prev_token_id, prev_position_offset, prev_embedding); + if (prev_is_placeholder) { + // Redirect to padding block 0 to avoid overwriting correct KV cache. + buf.out_new_cache_slots.back() = 0; + } } selected_row_idx.emplace_back(