From f2ecfa9cd7d8d9585511904539ecd359e9396207 Mon Sep 17 00:00:00 2001 From: Woosuk Kwon Date: Mon, 30 Mar 2026 17:12:27 +0000 Subject: [PATCH] topk init Signed-off-by: Woosuk Kwon --- .../deepseek_v3_2_monolithic/decoder_layer.py | 3 +++ .../models/deepseek_v3_2_monolithic/ops.py | 26 ++++++++++++++++++- .../sparse_indexer.py | 1 - 3 files changed, 28 insertions(+), 2 deletions(-) diff --git a/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py b/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py index ca524972aad..19974faf6a5 100644 --- a/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py +++ b/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py @@ -166,6 +166,7 @@ class MonolithicDecoderLayer(nn.Module): index_weights, _ = self.attn.indexer_weights_proj(hidden_states) # Step 2. Q RMS norm + KV RMS norm + KV RoPE + Index K layer norm + RoPE + # + Init topk indices q_c, kv_c = fused_norm_rope( positions, # Q RMS norm @@ -185,6 +186,8 @@ class MonolithicDecoderLayer(nn.Module): self.attn.indexer_k_norm.bias, self.attn.rms_norm_eps, self.attn.indexer_rope_emb.cos_sin_cache, + # Top k indices + self.attn.topk_indices_buffer, ) # Step 3. q_c -> q diff --git a/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py b/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py index 7915ebc3471..7cd50e4e8a9 100644 --- a/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py +++ b/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py @@ -429,6 +429,11 @@ def _fused_norm_rope_kernel( index_k_rope_cos_sin_cache_ptr, index_k_rope_cos_sin_cache_stride, INDEX_K_HALF_ROT_DIM: tl.constexpr, + # Top k indices + topk_indices_ptr, + topk_indices_stride, + TOPK: tl.constexpr, + TOPK_BLOCK_SIZE: tl.constexpr, ): tok_idx = tl.program_id(1) if tl.program_id(0) == 0: @@ -518,6 +523,17 @@ def _fused_norm_rope_kernel( False, ) return + elif tl.program_id(0) == 4: + # Fill top k indices buffer with -1 + for i in range(0, TOPK, TOPK_BLOCK_SIZE): + offset = i + tl.arange(0, TOPK_BLOCK_SIZE) + mask = offset < TOPK + tl.store( + topk_indices_ptr + tok_idx * topk_indices_stride + offset, + -1, + mask=mask, + ) + return def fused_norm_rope( @@ -535,21 +551,24 @@ def fused_norm_rope( index_k_layer_norm_bias: torch.Tensor, index_k_layer_norm_eps: float, index_k_rope_cos_sin_cache: torch.Tensor, + topk_indices_buffer: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: assert positions.ndim == 1 assert q_c.ndim == 2 assert kv_c.ndim == 2 assert k_pe.ndim == 2 assert index_k.ndim == 2 + assert topk_indices_buffer.ndim == 2 num_tokens = positions.shape[0] q_dim = q_c.shape[-1] kv_dim = kv_c.shape[-1] index_k_dim = index_k.shape[-1] + topk = topk_indices_buffer.shape[-1] q_c_out = torch.empty_like(q_c) kv_c_out = torch.empty_like(kv_c) - _fused_norm_rope_kernel[(4, num_tokens)]( + _fused_norm_rope_kernel[(5, num_tokens)]( positions, # Q RMS norm q_c, @@ -585,6 +604,11 @@ def fused_norm_rope( index_k_rope_cos_sin_cache, index_k_rope_cos_sin_cache.stride(0), index_k_rope_cos_sin_cache.shape[-1] // 2, + # Top k indices buffer + topk_indices_buffer, + topk_indices_buffer.stride(0), + topk, + TOPK_BLOCK_SIZE=1024, ) return q_c_out, kv_c_out diff --git a/vllm/model_executor/models/deepseek_v3_2_monolithic/sparse_indexer.py b/vllm/model_executor/models/deepseek_v3_2_monolithic/sparse_indexer.py index 1cf9d1decdc..d4d33600c9b 100644 --- a/vllm/model_executor/models/deepseek_v3_2_monolithic/sparse_indexer.py +++ b/vllm/model_executor/models/deepseek_v3_2_monolithic/sparse_indexer.py @@ -62,7 +62,6 @@ def sparse_attn_indexer( scale_fmt, ) - topk_indices_buffer[: hidden_states.shape[0]] = -1 if has_prefill: prefill_metadata = attn_metadata.prefill assert prefill_metadata is not None