[DSv4 Perf] Skip topk and router when not needed, 3.4% E2E TTFT improvement for Decode case (#49486)

Signed-off-by: yewentao256 <zhyanwentao@126.com>
This commit is contained in:
Wentao Ye
2026-07-23 13:08:08 -04:00
committed by GitHub
parent 0e36e3bbd1
commit b0cb1da1bd
+43
View File
@@ -47,6 +47,7 @@ from vllm.model_executor.layers.quantization import QuantizationConfig
from vllm.model_executor.models.utils import extract_layer_index
from vllm.models.deepseek_v4.common.rope import build_deepseek_v4_rope
from vllm.models.deepseek_v4.compressor import DeepseekCompressor
from vllm.triton_utils import tl, triton
from vllm.utils.multi_stream_utils import (
execute_in_parallel,
maybe_execute_in_parallel,
@@ -66,6 +67,25 @@ from vllm.v1.kv_cache_interface import (
logger = init_logger(__name__)
@triton.jit
def _fill_short_context_topk_indices(
output,
positions,
TOP_K: tl.constexpr,
COMPRESS_RATIO: tl.constexpr,
PADDED_TOP_K: tl.constexpr,
):
# small triton kernel that selects every candidate, -1 otherwise
row = tl.program_id(0)
offsets = tl.arange(0, PADDED_TOP_K)
num_compressed = (tl.load(positions + row) + 1) // COMPRESS_RATIO
tl.store(
output + row * TOP_K + offsets,
tl.where(offsets < num_compressed, offsets, -1),
mask=offsets < TOP_K,
)
def _resolve_dsv4_kv_cache_dtype(
use_fp8_ds_mla_layout: bool,
kv_cache_dtype: str,
@@ -787,6 +807,29 @@ class DeepseekV4Indexer(nn.Module):
) -> torch.Tensor:
compressor = self.compressor
attn_metadata = get_forward_context().attn_metadata
if isinstance(attn_metadata, dict):
indexer_metadata = cast(Any, attn_metadata[self.k_cache.prefix])
if indexer_metadata.max_seq_len // self.compress_ratio <= self.topk_tokens:
# candidates num smaller than topk, every candidate is selected
# but we still need to build k cache
compressor(compressed_kv_score, positions, rotary_emb)
assert self.topk_indices_buffer is not None
num_tokens = (
indexer_metadata.num_decode_tokens
+ indexer_metadata.num_prefill_tokens
)
if num_tokens > 0:
_fill_short_context_topk_indices[(num_tokens,)](
self.topk_indices_buffer,
positions,
TOP_K=self.topk_tokens,
COMPRESS_RATIO=self.compress_ratio,
PADDED_TOP_K=triton.next_power_of_2(self.topk_tokens),
num_warps=8,
)
return self.topk_indices_buffer
def wq_b_and_q_quant():
# ReplicatedLinear returns (output, bias); bias is None.
q, _ = self.wq_b(qr)