From fda87173ab62765e8393cdd88a0ea44c61320a53 Mon Sep 17 00:00:00 2001 From: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com> Date: Wed, 8 Jul 2026 14:48:35 +0800 Subject: [PATCH] Revert "Increase the KV cache pool when using indexShare by 15% (#30310)" (#30472) --- python/sglang/srt/mem_cache/memory_pool.py | 38 +++++-------------- .../model_runner_kv_cache_mixin.py | 6 --- .../srt/model_executor/pool_configurator.py | 11 +----- 3 files changed, 11 insertions(+), 44 deletions(-) diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index cbd1c6ceb..2a535af1c 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -3137,7 +3137,6 @@ class DSATokenToKVPool(MLATokenToKVPool): start_layer: Optional[int] = None, end_layer: Optional[int] = None, index_buf_size: Optional[int] = None, - skip_topk_layers: Optional[List[bool]] = None, ): override_dim = ( kv_cache_dim if kv_cache_dim != kv_lora_rank + qk_rope_head_dim else None @@ -3165,13 +3164,6 @@ class DSATokenToKVPool(MLATokenToKVPool): # num head == 1 and head dim == 128 for index_k in DSA assert index_head_dim == 128 - self.skip_topk_layers = ( - list(skip_topk_layers) - if skip_topk_layers is not None - else [False] * layer_num - ) - assert len(self.skip_topk_layers) == layer_num - if _is_hip: if aiter_can_use_preshuffle_paged_mqa(): assert ( @@ -3188,10 +3180,6 @@ class DSATokenToKVPool(MLATokenToKVPool): if self.custom_mem_pool else nullcontext() ): - cols = self.page_size * ( - index_head_dim + index_head_dim // self.quant_block_size * 4 - ) - num_pages = (index_buf_size + page_size + 1) // self.page_size self.index_k_with_scale_buffer = [ torch.zeros( # Layout: @@ -3200,11 +3188,17 @@ class DSATokenToKVPool(MLATokenToKVPool): # data: for page i, # * buf[i, :page_size * head_dim] for fp8 data # * buf[i, page_size * head_dim:].view(float32) for scale - (0 if self.skip_topk_layers[i] else num_pages, cols), + ( + (index_buf_size + page_size + 1) // self.page_size, + self.page_size + * ( + index_head_dim + index_head_dim // self.quant_block_size * 4 + ), + ), dtype=self.index_k_with_scale_buffer_dtype, device=device, ) - for i in range(layer_num) + for _ in range(layer_num) ] self._finalize_allocation_log(size) @@ -3221,9 +3215,7 @@ class DSATokenToKVPool(MLATokenToKVPool): tgt_loc_flat = tgt_loc.view(-1).long() src_loc_flat = src_loc.view(-1).long() - for i, index_k in enumerate(self.index_k_with_scale_buffer): - if self.skip_topk_layers[i]: - continue + for index_k in self.index_k_with_scale_buffer: index_k[tgt_loc_flat] = index_k[src_loc_flat] def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor: @@ -3314,9 +3306,6 @@ class DSATokenToKVPool(MLATokenToKVPool): chunk_size = self.cpu_offloading_chunk_size page_chunk_size = max(1, chunk_size // self.page_size) for layer_id in range(self.layer_num): - if self.skip_topk_layers[layer_id]: - index_k_cpu.append([]) - continue index_k_cpu.append([]) for i in range(0, len(page_indices), page_chunk_size): chunk_page_indices = page_indices[i : i + page_chunk_size] @@ -3339,8 +3328,6 @@ class DSATokenToKVPool(MLATokenToKVPool): chunk_size = self.cpu_offloading_chunk_size page_chunk_size = max(1, chunk_size // self.page_size) for layer_id in range(self.layer_num): - if self.skip_topk_layers[layer_id]: - continue for i in range(0, len(page_indices), page_chunk_size): chunk_page_indices = page_indices[i : i + page_chunk_size] idx_cpu = index_k_cpu[layer_id][i // page_chunk_size] @@ -3359,12 +3346,7 @@ class DSATokenToKVPool(MLATokenToKVPool): self.index_k_with_scale_buffer[i].nbytes for i in range(self.layer_num) ] item_lens = [ - ( - 0 - if self.skip_topk_layers[i] - else self.index_k_with_scale_buffer[i][0].nbytes - ) - for i in range(self.layer_num) + self.index_k_with_scale_buffer[i][0].nbytes for i in range(self.layer_num) ] return data_ptrs, data_lens, item_lens diff --git a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py index 42bf6cc34..7ba98d7b6 100644 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py @@ -7,7 +7,6 @@ from typing import TYPE_CHECKING, Optional import torch from sglang.srt.configs.model_config import ( - dsa_layer_skips_topk, get_dsa_index_head_dim, get_minimax_sparse_attention_config, get_minimax_sparse_disable_value_layer_ids, @@ -958,11 +957,6 @@ class ModelRunnerKVCacheMixin: pool_kwargs["host_to_device_ratio"] = parse_hisparse_config( self.server_args ).host_to_device_ratio - elif not self.is_draft_worker: - pool_kwargs["skip_topk_layers"] = [ - dsa_layer_skips_topk(self.model_config.hf_config, layer_id) - for layer_id in range(self.start_layer, self.end_layer) - ] self.token_to_kv_pool = PoolCls( self.max_total_num_tokens, page_size=self.page_size, diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index 032c20f0b..50c37332e 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -20,7 +20,6 @@ from typing import TYPE_CHECKING, Optional import torch from sglang.srt.configs.model_config import ( - dsa_layer_skips_topk, get_dsa_index_head_dim, get_minimax_sparse_attention_config, get_minimax_sparse_disable_value_layer_ids, @@ -210,15 +209,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): element_size = torch._utils._element_size( DSATokenToKVPool.index_k_with_scale_buffer_dtype ) - if mr.enable_hisparse or mr.is_draft_worker: - num_indexer_layers = num_layers - else: - num_indexer_layers = sum( - 1 - for layer_id in range(mr.start_layer, mr.end_layer) - if not dsa_layer_skips_topk(model_config.hf_config, layer_id) - ) - cell_size += indexer_size_per_token * num_indexer_layers * element_size + cell_size += indexer_size_per_token * num_layers * element_size elif is_minimax_sparse(model_config.hf_config): # Mirrors MiniMaxSparseKVPool: main pool (K+V all layers) + indexer pool # (sparse-only, single-head; kv layers store K+V, k-only layers store K).