From 455ab36eebee1eb09769f2779d8d01d1ac6577f3 Mon Sep 17 00:00:00 2001 From: Brayden Zhong Date: Tue, 7 Jul 2026 20:59:13 -0700 Subject: [PATCH] Increase the KV cache pool when using indexShare by 15% (#30310) Co-authored-by: Brayden Zhong --- 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, 44 insertions(+), 11 deletions(-) diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 2a535af1c..cbd1c6ceb 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -3137,6 +3137,7 @@ 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 @@ -3164,6 +3165,13 @@ 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 ( @@ -3180,6 +3188,10 @@ 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: @@ -3188,17 +3200,11 @@ 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 - ( - (index_buf_size + page_size + 1) // self.page_size, - self.page_size - * ( - index_head_dim + index_head_dim // self.quant_block_size * 4 - ), - ), + (0 if self.skip_topk_layers[i] else num_pages, cols), dtype=self.index_k_with_scale_buffer_dtype, device=device, ) - for _ in range(layer_num) + for i in range(layer_num) ] self._finalize_allocation_log(size) @@ -3215,7 +3221,9 @@ class DSATokenToKVPool(MLATokenToKVPool): tgt_loc_flat = tgt_loc.view(-1).long() src_loc_flat = src_loc.view(-1).long() - for index_k in self.index_k_with_scale_buffer: + for i, index_k in enumerate(self.index_k_with_scale_buffer): + if self.skip_topk_layers[i]: + continue index_k[tgt_loc_flat] = index_k[src_loc_flat] def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor: @@ -3306,6 +3314,9 @@ 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] @@ -3328,6 +3339,8 @@ 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] @@ -3346,7 +3359,12 @@ class DSATokenToKVPool(MLATokenToKVPool): self.index_k_with_scale_buffer[i].nbytes for i in range(self.layer_num) ] item_lens = [ - self.index_k_with_scale_buffer[i][0].nbytes for i in range(self.layer_num) + ( + 0 + if self.skip_topk_layers[i] + else 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 3cc968cd1..cb1a16585 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,6 +7,7 @@ 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, @@ -947,6 +948,11 @@ 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 50c37332e..032c20f0b 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -20,6 +20,7 @@ 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, @@ -209,7 +210,15 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): element_size = torch._utils._element_size( DSATokenToKVPool.index_k_with_scale_buffer_dtype ) - cell_size += indexer_size_per_token * num_layers * element_size + 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 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).