diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 0b1a152f2..c293bc681 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -2,12 +2,13 @@ from __future__ import annotations import logging import math -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Optional import msgspec import torch +from sglang.srt.configs.hybrid_arch import hybrid_gdn_config, mambaish_config from sglang.srt.configs.model_config import ( ModelConfig, get_dsa_index_head_dim, @@ -102,7 +103,6 @@ MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP = 2 MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP_LAZY = 1 MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP = 1 - if TYPE_CHECKING: from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.mem_cache.unified_memory_pool import ( @@ -112,6 +112,9 @@ if TYPE_CHECKING: from sglang.srt.model_executor.model_runner_components.layer_setup import ( ModelLayerInfo, ) + from sglang.srt.model_executor.model_runner_components.spec_aux_hidden_state import ( + SpecAuxHiddenStateConfig, + ) from sglang.srt.model_executor.pool_configurator import ( MemoryPoolConfig, ) @@ -136,44 +139,49 @@ class _InitializedPools(msgspec.Struct, frozen=True, kw_only=True): unified_memory_pool: Optional[UnifiedKVPool] = None -@dataclass(frozen=True, slots=True, kw_only=True) +class _PoolSizes(msgspec.Struct, frozen=True, kw_only=True): + max_total_num_tokens: int + max_running_requests: int + full_max_total_num_tokens: Optional[int] + swa_max_total_num_tokens: Optional[int] + c4_max_total_num_tokens: int + c128_max_total_num_tokens: int + c4_state_pool_size: int + c128_state_pool_size: int + c4_state_dtype: Optional[torch.dtype] + c128_state_dtype: Optional[torch.dtype] + + +@dataclass(slots=True, kw_only=True) class KVCacheConfigurator: device: str gpu_id: int ps: ParallelState + pp_group: Any model_config: ModelConfig server_args: ServerArgs kv_cache_dtype: torch.dtype + model_dtype: torch.dtype page_size: int + sliding_window_size: Optional[int] spec_algorithm: SpeculativeAlgorithm is_draft_worker: bool post_capture_kv_active: bool - dflash_draft_num_layers: Optional[int] + spec_aux_config: SpecAuxHiddenStateConfig is_hybrid_swa: bool is_hybrid_swa_compress: bool use_mla_backend: bool - mambaish_config: Optional[Any] - hybrid_gdn_config: Optional[Any] - # PP slice - start_layer: int - end_layer: int - num_effective_layers: int + layer_info: ModelLayerInfo forward_stream: Any req_to_token_pool: Optional[ReqToTokenPool] token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] memory_pool_config: Optional[MemoryPoolConfig] + mambaish_config: Optional[Any] = field(init=False) + hybrid_gdn_config: Optional[Any] = field(init=False) - @property - def layer_info(self) -> ModelLayerInfo: - from sglang.srt.model_executor.model_runner_components.layer_setup import ( - ModelLayerInfo, - ) - - return ModelLayerInfo( - start_layer=self.start_layer, - end_layer=self.end_layer, - num_effective_layers=self.num_effective_layers, - ) + def __post_init__(self) -> None: + self.mambaish_config = mambaish_config(self.model_config) + self.hybrid_gdn_config = hybrid_gdn_config(self.model_config) def configure(self, *, pre_model_load_memory: int) -> KVCacheConfigResult: """Apply a resolved MemoryPoolConfig and initialize pools.""" @@ -185,6 +193,32 @@ class KVCacheConfigurator: else: config = self._resolve_memory_pool_config(pre_model_load_memory) + sizes = self._derive_pool_sizes(config=config) + + pools = self._init_pools( + sizes=sizes, + req_to_token_pool=self.req_to_token_pool, + token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, + ) + + logger.info( + f"Memory pool end. " + f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB" + ) + + return KVCacheConfigResult( + max_total_num_tokens=sizes.max_total_num_tokens, + max_running_requests=sizes.max_running_requests, + full_max_total_num_tokens=sizes.full_max_total_num_tokens, + swa_max_total_num_tokens=sizes.swa_max_total_num_tokens, + req_to_token_pool=pools.req_to_token_pool, + token_to_kv_pool=pools.token_to_kv_pool, + token_to_kv_pool_allocator=pools.token_to_kv_pool_allocator, + memory_pool_config=config, + unified_memory_pool=pools.unified_memory_pool, + ) + + def _derive_pool_sizes(self, *, config: MemoryPoolConfig) -> _PoolSizes: max_total_num_tokens = config.max_total_num_tokens max_running_requests = config.max_running_requests full_max_total_num_tokens = None @@ -214,7 +248,7 @@ class KVCacheConfigurator: if is_deepseek_v4(self.model_config.hf_config): c4_state_dtype, c128_state_dtype = _get_dsv4_compress_state_dtypes() - pools = self._init_pools( + return _PoolSizes( max_total_num_tokens=max_total_num_tokens, max_running_requests=max_running_requests, full_max_total_num_tokens=full_max_total_num_tokens, @@ -225,40 +259,12 @@ class KVCacheConfigurator: c128_state_pool_size=c128_state_pool_size, c4_state_dtype=c4_state_dtype, c128_state_dtype=c128_state_dtype, - req_to_token_pool=self.req_to_token_pool, - token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, - ) - - logger.info( - f"Memory pool end. " - f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB" - ) - - return KVCacheConfigResult( - max_total_num_tokens=max_total_num_tokens, - max_running_requests=max_running_requests, - full_max_total_num_tokens=full_max_total_num_tokens, - swa_max_total_num_tokens=swa_max_total_num_tokens, - req_to_token_pool=pools.req_to_token_pool, - token_to_kv_pool=pools.token_to_kv_pool, - token_to_kv_pool_allocator=pools.token_to_kv_pool_allocator, - memory_pool_config=config, - unified_memory_pool=pools.unified_memory_pool, ) def _init_pools( self, *, - max_total_num_tokens: int, - max_running_requests: int, - full_max_total_num_tokens: Optional[int], - swa_max_total_num_tokens: Optional[int], - c4_max_total_num_tokens: int, - c128_max_total_num_tokens: int, - c4_state_pool_size: int, - c128_state_pool_size: int, - c4_state_dtype: Optional[torch.dtype], - c128_state_dtype: Optional[torch.dtype], + sizes: _PoolSizes, req_to_token_pool: Optional[ReqToTokenPool], token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator], ) -> _InitializedPools: @@ -275,14 +281,14 @@ class KVCacheConfigurator: ): if self.mambaish_config is not None: bundle = self._init_unified_mamba_pools( - max_num_reqs=max_running_requests, - max_total_num_tokens=max_total_num_tokens, + max_num_reqs=sizes.max_running_requests, + max_total_num_tokens=sizes.max_total_num_tokens, ) elif self.is_hybrid_swa and not is_deepseek_v4(self.model_config.hf_config): bundle = self._init_unified_swa_pools( - max_num_reqs=max_running_requests, - full_max_total_num_tokens=full_max_total_num_tokens, - swa_max_total_num_tokens=swa_max_total_num_tokens, + max_num_reqs=sizes.max_running_requests, + full_max_total_num_tokens=sizes.full_max_total_num_tokens, + swa_max_total_num_tokens=sizes.swa_max_total_num_tokens, ) else: # Fail loud, not silently fall through to the normal pools (which would @@ -300,98 +306,12 @@ class KVCacheConfigurator: token_to_kv_pool_allocator=bundle.token_to_kv_pool_allocator, unified_memory_pool=bundle.unified_memory_pool, ) - max_num_reqs = max_running_requests # Initialize req_to_token_pool if req_to_token_pool is None: - max_spec_draft_tokens = self.server_args.max_speculative_num_draft_tokens - extra_max_context_len = get_req_to_token_extra_context_len(self.server_args) - - if self.server_args.disaggregation_mode == "decode": - from sglang.srt.disaggregation.decode import ( - DecodeReqToTokenPool, - HybridMambaDecodeReqToTokenPool, - ) - - # Extra slots for pre-allocated requests - pre_alloc_size = self.server_args.disaggregation_decode_extra_slots - if config := self.mambaish_config: - req_to_token_pool = HybridMambaDecodeReqToTokenPool( - size=max_num_reqs, - max_context_len=self.model_config.context_len - + extra_max_context_len, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - cache_params=config.mamba2_cache_params, - mamba_layer_ids=( - [ - i - for i in config.mamba2_cache_params.layers - if self.start_layer <= i < self.end_layer - ] - ), - speculative_num_draft_tokens=max_spec_draft_tokens, - speculative_eagle_topk=self.server_args.speculative_eagle_topk, - enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), - pre_alloc_size=pre_alloc_size, - enable_overlap_schedule=not self.server_args.disable_overlap_schedule, - mamba_size=self.server_args.max_mamba_cache_size, - start_layer=self.start_layer, - ) - else: - req_to_token_pool = DecodeReqToTokenPool( - size=max_num_reqs, - max_context_len=self.model_config.context_len - + extra_max_context_len, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - pre_alloc_size=pre_alloc_size, - ) - elif config := self.mambaish_config: - req_to_token_pool = HybridReqToTokenPool( - size=max_num_reqs, - mamba_size=self.server_args.max_mamba_cache_size, - mamba_spec_state_size=max_num_reqs, - max_context_len=self.model_config.context_len - + extra_max_context_len, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - cache_params=config.mamba2_cache_params, - mamba_layer_ids=( - [ - i - for i in config.mamba2_cache_params.layers - if self.start_layer <= i < self.end_layer - ] - ), - enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), - enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(), - speculative_num_draft_tokens=max_spec_draft_tokens, - speculative_eagle_topk=self.server_args.speculative_eagle_topk, - enable_overlap_schedule=not self.server_args.disable_overlap_schedule, - start_layer=self.start_layer, - enable_linear_replayssm=self.server_args.enable_linear_replayssm, - linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len, - mamba_envelope_layout=self.server_args.enable_page_major_kv_layout, - ) - else: - # DSV4 on NPU needs an extended ReqToTokenPool holding per-req - # swa/c4/c128/c{4,128}_state tables; others stay on the stock one. - req_to_token_pool_cls = ReqToTokenPool - if _is_npu and is_deepseek_v4(self.model_config.hf_config): - from sglang.srt.hardware_backend.npu.dsv4.dsv4_req_to_token_pool import ( - DSV4NPUReqToTokenPool, - ) - - req_to_token_pool_cls = DSV4NPUReqToTokenPool - - req_to_token_pool = req_to_token_pool_cls( - size=max_num_reqs, - max_context_len=self.model_config.context_len - + extra_max_context_len, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - ) + req_to_token_pool = self._build_req_to_token_pool( + max_num_reqs=sizes.max_running_requests + ) else: # Draft worker shares req_to_token_pool with the target worker. assert self.is_draft_worker @@ -404,574 +324,20 @@ class KVCacheConfigurator: is_dsa_model, is_dsv4_model, current_platform ) - # Page-granularity envelope layout for the MHA-shaped (full / SWA) pools, - # selected by swapping in the PageMajorMHATokenToKVPool subclass. The - # default keeps upstream's per-layer layout. The Mamba state pool is routed - # separately via `mamba_envelope_layout` on the req-to-token pool above. - enable_page_major = self.server_args.enable_page_major_kv_layout - mha_pool_class = ( - PageMajorMHATokenToKVPool if enable_page_major else MHATokenToKVPool + token_to_kv_pool = self._build_token_to_kv_pool( + sizes=sizes, + is_dsa_model=is_dsa_model, + is_dsv4_model=is_dsv4_model, + req_to_token_pool=req_to_token_pool, ) - if is_dsv4_model: - swa_page_size = self.server_args.page_size - if not _is_npu: - assert swa_page_size == 256, "In paged swa mode, page_size must be 256." - - if self.is_draft_worker: - from sglang.srt.models.deepseek_v4_nextn import ( - COMPRESS_RATIO_NEXTN_LAYER, - ) - - compression_ratios = [ - COMPRESS_RATIO_NEXTN_LAYER - ] * self.num_effective_layers - else: - compression_ratios = self.model_config.compress_ratios - - # NPU + DSV4 → paged-state subclass: the fused compressor kernel - # needs cache_mode=1 (paged); Atlas A3 rejects cache_mode=2 (ring), - # so the CUDA ring-buffer state path can't be shared. CUDA keeps - # DeepSeekV4TokenToKVPool unchanged; NPU recomputes state sizes below. - if _is_npu: - from sglang.srt.hardware_backend.npu.dsv4.dsv4_memory_pool import ( - DSV4NPUTokenToKVPool, - npu_state_pool_size, - ) - - pool_cls = DSV4NPUTokenToKVPool - # Recompute state pool sizes for the NPU paged formula (CUDA's - # ring sizes are dropped here). Tail-only allocation keeps the - # per-req-budget formula sufficient at any prefill length: long - # prompts allocate only ``tail+128`` (c4) / ``tail`` (c128) - # slots (tail = seq_len % 128), and decode is drained by - # sliding eviction in ``ScheduleBatch._evict_swa``. - c4_state_pool_size = npu_state_pool_size( - ratio=4, - page_size=self.server_args.page_size, - max_num_reqs=max_running_requests, - ) - c128_state_pool_size = npu_state_pool_size( - ratio=128, - page_size=self.server_args.page_size, - max_num_reqs=max_running_requests, - ) - else: - pool_cls = DeepSeekV4TokenToKVPool - c4_state_pool_size = c4_state_pool_size - c128_state_pool_size = c128_state_pool_size - - token_to_kv_pool = pool_cls( - max_num_reqs=max_running_requests, - # SWA ring is indexed by req_pool_idx; PD decode inflates req_to_token - # past max_running_requests (pre-alloc), so size to the real capacity. - num_req_slots=req_to_token_pool.req_to_token.shape[0], - swa_size=swa_max_total_num_tokens, - c4_size=c4_max_total_num_tokens, - c128_size=c128_max_total_num_tokens, - c4_state_pool_size=c4_state_pool_size, - c128_state_pool_size=c128_state_pool_size, - page_size=self.server_args.page_size, - swa_page_size=swa_page_size, - sliding_window=self.model_config.window_size, - dtype=self.kv_cache_dtype, - c4_state_dtype=c4_state_dtype, - c128_state_dtype=c128_state_dtype, - qk_nope_head_dim=self.model_config.qk_nope_head_dim, - qk_rope_head_dim=self.model_config.qk_rope_head_dim, - indexer_head_dim=self.model_config.index_head_dim, - layer_num=self.num_effective_layers, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - compression_ratios=compression_ratios, - start_layer=self.start_layer, - end_layer=self.end_layer, - enable_hisparse=self.server_args.enable_hisparse, - online_mtp_max_draft_tokens=( - self.server_args.max_speculative_num_draft_tokens or 0 - ), - ) - elif current_platform.is_out_of_tree() and not self.mambaish_config: - if self.use_mla_backend and is_dsa_model: - PoolCls = current_platform.get_dsa_kv_pool_cls() - token_to_kv_pool = PoolCls( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - kv_lora_rank=self.model_config.kv_lora_rank, - qk_rope_head_dim=self.model_config.qk_rope_head_dim, - layer_num=self.num_effective_layers, - device=self.device, - kv_cache_dim=calculate_mla_kv_cache_dim( - model_config=self.model_config, - kv_cache_dtype=self.kv_cache_dtype, - server_args=self.server_args, - ), - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config), - ) - elif self.use_mla_backend: - PoolCls = current_platform.get_mla_kv_pool_cls() - token_to_kv_pool = PoolCls( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - kv_lora_rank=self.model_config.kv_lora_rank, - qk_rope_head_dim=self.model_config.qk_rope_head_dim, - index_head_dim=( - self.model_config.index_head_dim if is_dsa_model else None - ), - layer_num=self.num_effective_layers, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - ) - else: - PoolCls = current_platform.get_mha_kv_pool_cls() - token_to_kv_pool = PoolCls( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - head_num=self.model_config.get_num_kv_heads( - get_parallel().attn_tp_size - ), - head_dim=self.model_config.head_dim, - layer_num=self.num_effective_layers, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - ) - elif ( - self.server_args.attention_backend == "ascend" and not self.mambaish_config - ): - if self.is_hybrid_swa: - from sglang.srt.hardware_backend.npu.memory_pool_npu import ( - NPUMHATokenToKVPool, - ) - - kwargs = {} - if self.is_hybrid_swa_compress: - kwargs = { - "swa_head_num": max( - 1, - self.model_config.hf_text_config.swa_num_key_value_heads - // get_parallel().attn_tp_size, - ), - "swa_head_dim": self.model_config.swa_head_dim, - "swa_v_head_dim": self.model_config.swa_v_head_dim, - "v_head_dim": self.model_config.v_head_dim, - } - token_to_kv_pool = SWAKVPool( - size=full_max_total_num_tokens, - size_swa=swa_max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - post_capture_active=self.post_capture_kv_active, - head_num=self.model_config.get_num_kv_heads( - get_parallel().attn_tp_size - ), - head_dim=self.model_config.head_dim, - swa_attention_layer_ids=self.model_config.swa_attention_layer_ids, - full_attention_layer_ids=self.model_config.full_attention_layer_ids, - device=self.device, - token_to_kv_pool_class=NPUMHATokenToKVPool, - **kwargs, - ) - elif self.use_mla_backend: - from sglang.srt.hardware_backend.npu.memory_pool_npu import ( - NPUMLATokenToKVPool, - ) - - token_to_kv_pool = NPUMLATokenToKVPool( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - kv_lora_rank=self.model_config.kv_lora_rank, - qk_rope_head_dim=self.model_config.qk_rope_head_dim, - index_head_dim=( - self.model_config.index_head_dim if is_dsa_model else None - ), - layer_num=self.num_effective_layers, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - ) - else: - from sglang.srt.hardware_backend.npu.memory_pool_npu import ( - NPUMHATokenToKVPool, - ) - - token_to_kv_pool = NPUMHATokenToKVPool( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - head_num=self.model_config.get_num_kv_heads( - get_parallel().attn_tp_size - ), - head_dim=self.model_config.head_dim, - layer_num=self.num_effective_layers, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - ) - elif self.use_mla_backend and is_dsa_model: - from sglang.srt.layers.cp.utils import get_glm_dsa_cp_layer_shard_info - - ( - dsa_cp_layer_shard_rank, - dsa_cp_layer_shard_size, - ) = get_glm_dsa_cp_layer_shard_info(self) - pool_kwargs = {} - if self.server_args.enable_hisparse: - PoolCls = HiSparseDSATokenToKVPool - from sglang.srt.mem_cache.sparsity import parse_hisparse_config - - pool_kwargs["host_to_device_ratio"] = parse_hisparse_config( - self.server_args - ).host_to_device_ratio - elif dsa_cp_layer_shard_rank is not None: - # DSA cache layer split: shard KV/indexer layers across CP ranks. - from sglang.srt.mem_cache.dsa_cache_layer_split import ( - LayerSplitDSATokenToKVPool, - ) - - PoolCls = LayerSplitDSATokenToKVPool - pool_kwargs["layer_shard_rank"] = dsa_cp_layer_shard_rank - pool_kwargs["layer_shard_size"] = dsa_cp_layer_shard_size - else: - PoolCls = DSATokenToKVPool - token_to_kv_pool = PoolCls( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - kv_lora_rank=self.model_config.kv_lora_rank, - qk_rope_head_dim=self.model_config.qk_rope_head_dim, - layer_num=self.num_effective_layers, - device=self.device, - kv_cache_dim=calculate_mla_kv_cache_dim( - model_config=self.model_config, - kv_cache_dtype=self.kv_cache_dtype, - server_args=self.server_args, - ), - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config), - **pool_kwargs, - ) - elif self.use_mla_backend and not self.mambaish_config: - assert not is_dsa_model - if is_float4_e2m1fn_x2(self.kv_cache_dtype): - token_to_kv_pool = MLATokenToKVPoolFP4( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - kv_lora_rank=self.model_config.kv_lora_rank, - qk_rope_head_dim=self.model_config.qk_rope_head_dim, - layer_num=self.num_effective_layers, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - ) - else: - token_to_kv_pool = MLATokenToKVPool( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - kv_lora_rank=self.model_config.kv_lora_rank, - qk_rope_head_dim=self.model_config.qk_rope_head_dim, - layer_num=self.num_effective_layers, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - ) - else: - if self.is_hybrid_swa: - kwargs = {} - if self.is_hybrid_swa_compress: - kwargs = { - "swa_head_num": max( - 1, - self.model_config.hf_text_config.swa_num_key_value_heads - // get_parallel().attn_tp_size, - ), - "swa_head_dim": self.model_config.swa_head_dim, - "swa_v_head_dim": self.model_config.swa_v_head_dim, - "v_head_dim": self.model_config.v_head_dim, - } - token_to_kv_pool = SWAKVPool( - size=full_max_total_num_tokens, - size_swa=swa_max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - head_num=self.model_config.get_num_kv_heads( - get_parallel().attn_tp_size - ), - head_dim=self.model_config.head_dim, - swa_attention_layer_ids=self.model_config.swa_attention_layer_ids, - full_attention_layer_ids=self.model_config.full_attention_layer_ids, - device=self.device, - enable_kv_cache_copy=( - self.server_args.speculative_algorithm is not None - ), - token_to_kv_pool_class=mha_pool_class, - **kwargs, - ) - elif is_minimax_sparse(self.model_config.hf_config): - _hf_config = self.model_config.hf_config - sparse_cfg = get_minimax_sparse_attention_config(_hf_config) - dense_layer_ids, sparse_layer_ids = get_minimax_sparse_layer_ids( - sparse_cfg - ) - disable_value_sparse_layer_ids = ( - get_minimax_sparse_disable_value_layer_ids(sparse_cfg) - ) - token_to_kv_pool = MiniMaxSparseKVPool( - size=max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - index_dtype=self.model_dtype, - head_num=self.model_config.get_num_kv_heads( - get_parallel().attn_tp_size - ), - head_dim=self.model_config.head_dim, - idx_head_dim=sparse_cfg["sparse_index_dim"], - dense_layer_ids=dense_layer_ids, - sparse_layer_ids=sparse_layer_ids, - disable_value_sparse_layer_ids=disable_value_sparse_layer_ids, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - ) - elif config := self.mambaish_config: - extra_args = {} - if self.use_mla_backend: - extra_args = { - "kv_lora_rank": self.model_config.kv_lora_rank, - "qk_rope_head_dim": self.model_config.qk_rope_head_dim, - } - token_to_kv_pool = HybridLinearKVPool( - page_size=self.server_args.page_size, - size=max_total_num_tokens, - dtype=self.kv_cache_dtype, - head_num=self.model_config.get_num_kv_heads( - get_parallel().attn_tp_size - ), - head_dim=self.model_config.head_dim, - # if draft worker, we only need 1 attention layer's kv pool - full_attention_layer_ids=( - [0] - if self.is_draft_worker - else [ - i - for i in config.full_attention_layer_ids - if self.start_layer <= i < self.end_layer - ] - ), - device=self.device, - mamba_pool=req_to_token_pool.mamba_pool, - enable_memory_saver=self.server_args.enable_memory_saver, - enable_kv_cache_copy=( - self.server_args.speculative_algorithm is not None - ), - use_mla=self.use_mla_backend, - start_layer=self.start_layer, - full_kv_pool_class=mha_pool_class, - post_capture_active=self.post_capture_kv_active, - **extra_args, - ) - else: - if is_float4_e2m1fn_x2(self.kv_cache_dtype): - assert ( - not enable_page_major - ), "page-major KV layout is not supported with fp4 KV cache" - token_to_kv_pool = MHATokenToKVPoolFP4( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - head_num=self.model_config.get_num_kv_heads( - get_parallel().attn_tp_size - ), - head_dim=self.model_config.head_dim, - v_head_dim=self.model_config.v_head_dim, - layer_num=self.num_effective_layers, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - enable_alt_stream=not self.server_args.enable_pdmux, - enable_kv_cache_copy=( - self.server_args.speculative_algorithm is not None - ), - ) - else: - pool_cls = ( - NoOpMHATokenToKVPool - if self.server_args.prefill_only_disable_kv_cache - else mha_pool_class - ) - token_to_kv_pool = pool_cls( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - head_num=self.model_config.get_num_kv_heads( - get_parallel().attn_tp_size - ), - head_dim=self.model_config.head_dim, - v_head_dim=self.model_config.v_head_dim, - layer_num=self.num_effective_layers, - device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, - start_layer=self.start_layer, - end_layer=self.end_layer, - enable_alt_stream=not self.server_args.enable_pdmux, - enable_kv_cache_copy=( - self.server_args.speculative_algorithm is not None - ), - post_capture_active=self.post_capture_kv_active, - ) - - # Initialize token_to_kv_pool_allocator - need_sort = self.server_args.disaggregation_mode in ("decode", "prefill") - if token_to_kv_pool_allocator is None: - if current_platform.is_out_of_tree(): - AllocatorCls = current_platform.get_paged_allocator_cls() - token_to_kv_pool_allocator = AllocatorCls( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - device=self.device, - kvcache=token_to_kv_pool, - need_sort=need_sort, - ) - elif _is_npu and ( - self.server_args.attention_backend == "ascend" - or is_dsv4_model - or self.hybrid_gdn_config is not None - ): - if self.is_hybrid_swa: - # DSV4 on NPU: SWA allocator subclass that also drives the - # c4/c128 allocators, producing a DSV4OutCacheLoc per alloc. - if is_dsv4_model: - from sglang.srt.hardware_backend.npu.dsv4.dsv4_allocator import ( - DSV4NPUTokenToKVPoolAllocator, - ) - - swa_allocator_cls = DSV4NPUTokenToKVPoolAllocator - else: - swa_allocator_cls = SWATokenToKVPoolAllocator - token_to_kv_pool_allocator = swa_allocator_cls( - full_max_total_num_tokens, - swa_max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - device=self.device, - kvcache=token_to_kv_pool, - need_sort=need_sort, - ) - else: - from sglang.srt.hardware_backend.npu.allocator_npu import ( - NPUPagedTokenToKVPoolAllocator, - ) - - token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - device=self.device, - kvcache=token_to_kv_pool, - need_sort=need_sort, - ) - else: - if self.is_hybrid_swa and full_max_total_num_tokens == 0: - token_to_kv_pool_allocator = PureSWATokenToKVPoolAllocator( - swa_max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - device=self.device, - kvcache=token_to_kv_pool, - need_sort=need_sort, - ) - elif self.is_hybrid_swa: - token_to_kv_pool_allocator = SWATokenToKVPoolAllocator( - full_max_total_num_tokens, - swa_max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - device=self.device, - kvcache=token_to_kv_pool, - need_sort=need_sort, - ) - else: - if self.server_args.enable_hisparse: - from sglang.srt.mem_cache.sparsity import ( - parse_hisparse_config, - ) - - hisparse_cfg = parse_hisparse_config(self.server_args) - token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator( - max_total_num_tokens, - page_size=self.server_args.page_size, - dtype=self.kv_cache_dtype, - device=self.device, - kvcache=token_to_kv_pool, - need_sort=need_sort, - host_to_device_ratio=hisparse_cfg.host_to_device_ratio, - ) - elif ( - self.server_args.page_size == 1 - and self.server_args.dcp_size == 1 - ): - token_to_kv_pool_allocator = TokenToKVPoolAllocator( - max_total_num_tokens, - dtype=self.kv_cache_dtype, - device=self.device, - kvcache=token_to_kv_pool, - need_sort=need_sort, - ) - else: - token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator( - max_total_num_tokens * self.server_args.dcp_size, - page_size=self.server_args.page_size - * self.server_args.dcp_size, - dtype=self.kv_cache_dtype, - device=self.device, - kvcache=token_to_kv_pool, - need_sort=need_sort, - ) - - if self.server_args.enable_hisparse and is_dsv4_model: - assert self.is_hybrid_swa, "DeepSeek V4 HiSparse requires SWA mode." - token_to_kv_pool_allocator = DeepSeekV4HiSparseTokenToKVPoolAllocator( - token_to_kv_pool_allocator - ) - - # DSV4-NPU: wire allocator back-ref into req_to_token_pool so its - # free(req) can release c4/c128 pool pages alongside the slot. - if hasattr(req_to_token_pool, "register_dsv4_allocator"): - req_to_token_pool.register_dsv4_allocator(token_to_kv_pool_allocator) - - else: - assert self.is_draft_worker - if self.is_hybrid_swa: - swa_allocator = getattr( - token_to_kv_pool_allocator, - "logical_attn_allocator", - token_to_kv_pool_allocator, - ) - assert isinstance(swa_allocator, SWATokenToKVPoolAllocator) - token_to_kv_pool.register_mapping( - swa_allocator.full_to_swa_index_mapping - ) + token_to_kv_pool_allocator = self._build_token_to_kv_pool_allocator( + sizes=sizes, + token_to_kv_pool=token_to_kv_pool, + is_dsv4_model=is_dsv4_model, + req_to_token_pool=req_to_token_pool, + token_to_kv_pool_allocator=token_to_kv_pool_allocator, + ) # Defensive check: the explicit validation above should reject known # unsupported pool families before allocation. Keep this guard here so @@ -1194,6 +560,844 @@ class KVCacheConfigurator: "attention, no HiSparse, and --kv-cache-dtype != fp4_e2m1." ) + def _build_req_to_token_pool(self, *, max_num_reqs: int) -> ReqToTokenPool: + extra_max_context_len = get_req_to_token_extra_context_len(self.server_args) + + if self.server_args.disaggregation_mode == "decode": + # Extra slots for pre-allocated requests + pre_alloc_size = self.server_args.disaggregation_decode_extra_slots + if self.mambaish_config: + req_to_token_pool = self._build_hybrid_mamba_decode_req_pool( + max_num_reqs=max_num_reqs, + extra_max_context_len=extra_max_context_len, + pre_alloc_size=pre_alloc_size, + ) + else: + req_to_token_pool = self._build_decode_req_pool( + max_num_reqs=max_num_reqs, + extra_max_context_len=extra_max_context_len, + pre_alloc_size=pre_alloc_size, + ) + elif self.mambaish_config: + req_to_token_pool = self._build_hybrid_req_pool( + max_num_reqs=max_num_reqs, + extra_max_context_len=extra_max_context_len, + ) + else: + req_to_token_pool = self._build_default_req_pool( + max_num_reqs=max_num_reqs, + extra_max_context_len=extra_max_context_len, + ) + return req_to_token_pool + + def _build_hybrid_mamba_decode_req_pool( + self, + *, + max_num_reqs: int, + extra_max_context_len: int, + pre_alloc_size: int, + ) -> ReqToTokenPool: + from sglang.srt.disaggregation.decode import ( + HybridMambaDecodeReqToTokenPool, + ) + + req_to_token_pool = HybridMambaDecodeReqToTokenPool( + size=max_num_reqs, + max_context_len=self.model_config.context_len + extra_max_context_len, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + cache_params=self.mambaish_config.mamba2_cache_params, + mamba_layer_ids=( + [ + i + for i in self.mambaish_config.mamba2_cache_params.layers + if self.layer_info.start_layer <= i < self.layer_info.end_layer + ] + ), + speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens, + speculative_eagle_topk=self.server_args.speculative_eagle_topk, + enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), + pre_alloc_size=pre_alloc_size, + enable_overlap_schedule=not self.server_args.disable_overlap_schedule, + mamba_size=self.server_args.max_mamba_cache_size, + start_layer=self.layer_info.start_layer, + ) + return req_to_token_pool + + def _build_decode_req_pool( + self, + *, + max_num_reqs: int, + extra_max_context_len: int, + pre_alloc_size: int, + ) -> ReqToTokenPool: + from sglang.srt.disaggregation.decode import DecodeReqToTokenPool + + req_to_token_pool = DecodeReqToTokenPool( + size=max_num_reqs, + max_context_len=self.model_config.context_len + extra_max_context_len, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + pre_alloc_size=pre_alloc_size, + ) + return req_to_token_pool + + def _build_hybrid_req_pool( + self, + *, + max_num_reqs: int, + extra_max_context_len: int, + ) -> ReqToTokenPool: + req_to_token_pool = HybridReqToTokenPool( + size=max_num_reqs, + mamba_size=self.server_args.max_mamba_cache_size, + mamba_spec_state_size=max_num_reqs, + max_context_len=self.model_config.context_len + extra_max_context_len, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + cache_params=self.mambaish_config.mamba2_cache_params, + mamba_layer_ids=( + [ + i + for i in self.mambaish_config.mamba2_cache_params.layers + if self.layer_info.start_layer <= i < self.layer_info.end_layer + ] + ), + enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), + enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(), + speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens, + speculative_eagle_topk=self.server_args.speculative_eagle_topk, + enable_overlap_schedule=not self.server_args.disable_overlap_schedule, + start_layer=self.layer_info.start_layer, + enable_linear_replayssm=self.server_args.enable_linear_replayssm, + linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len, + mamba_envelope_layout=self.server_args.enable_page_major_kv_layout, + ) + return req_to_token_pool + + def _build_default_req_pool( + self, + *, + max_num_reqs: int, + extra_max_context_len: int, + ) -> ReqToTokenPool: + # DSV4 on NPU needs an extended ReqToTokenPool holding per-req + # swa/c4/c128/c{4,128}_state tables; others stay on the stock one. + req_to_token_pool_cls = ReqToTokenPool + if _is_npu and is_deepseek_v4(self.model_config.hf_config): + from sglang.srt.hardware_backend.npu.dsv4.dsv4_req_to_token_pool import ( + DSV4NPUReqToTokenPool, + ) + + req_to_token_pool_cls = DSV4NPUReqToTokenPool + + req_to_token_pool = req_to_token_pool_cls( + size=max_num_reqs, + max_context_len=self.model_config.context_len + extra_max_context_len, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + ) + return req_to_token_pool + + def _build_token_to_kv_pool( + self, + *, + sizes: _PoolSizes, + is_dsa_model: bool, + is_dsv4_model: bool, + req_to_token_pool: ReqToTokenPool, + ) -> KVCache: + # Page-granularity envelope layout for the MHA-shaped (full / SWA) pools, + # selected by swapping in the PageMajorMHATokenToKVPool subclass. The + # default keeps upstream's per-layer layout. The Mamba state pool is routed + # separately via `mamba_envelope_layout` on the req-to-token pool above. + enable_page_major = self.server_args.enable_page_major_kv_layout + mha_pool_class = ( + PageMajorMHATokenToKVPool if enable_page_major else MHATokenToKVPool + ) + + if is_dsv4_model: + token_to_kv_pool = self._build_dsv4_kv_pool( + max_running_requests=sizes.max_running_requests, + swa_max_total_num_tokens=sizes.swa_max_total_num_tokens, + c4_max_total_num_tokens=sizes.c4_max_total_num_tokens, + c128_max_total_num_tokens=sizes.c128_max_total_num_tokens, + c4_state_pool_size=sizes.c4_state_pool_size, + c128_state_pool_size=sizes.c128_state_pool_size, + c4_state_dtype=sizes.c4_state_dtype, + c128_state_dtype=sizes.c128_state_dtype, + req_to_token_pool=req_to_token_pool, + ) + elif current_platform.is_out_of_tree() and not self.mambaish_config: + if self.use_mla_backend and is_dsa_model: + token_to_kv_pool = self._build_oot_dsa_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + ) + elif self.use_mla_backend: + token_to_kv_pool = self._build_oot_mla_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + is_dsa_model=is_dsa_model, + ) + else: + token_to_kv_pool = self._build_oot_mha_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + ) + elif ( + self.server_args.attention_backend == "ascend" and not self.mambaish_config + ): + if self.is_hybrid_swa: + token_to_kv_pool = self._build_ascend_swa_kv_pool( + full_max_total_num_tokens=sizes.full_max_total_num_tokens, + swa_max_total_num_tokens=sizes.swa_max_total_num_tokens, + ) + elif self.use_mla_backend: + token_to_kv_pool = self._build_ascend_mla_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + is_dsa_model=is_dsa_model, + ) + else: + token_to_kv_pool = self._build_ascend_mha_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + ) + elif self.use_mla_backend and is_dsa_model: + token_to_kv_pool = self._build_dsa_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + ) + elif self.use_mla_backend and not self.mambaish_config: + assert not is_dsa_model + if is_float4_e2m1fn_x2(self.kv_cache_dtype): + token_to_kv_pool = self._build_mla_fp4_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + ) + else: + token_to_kv_pool = self._build_mla_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + ) + else: + if self.is_hybrid_swa: + token_to_kv_pool = self._build_hybrid_swa_kv_pool( + full_max_total_num_tokens=sizes.full_max_total_num_tokens, + swa_max_total_num_tokens=sizes.swa_max_total_num_tokens, + mha_pool_class=mha_pool_class, + ) + elif is_minimax_sparse(self.model_config.hf_config): + token_to_kv_pool = self._build_minimax_sparse_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + ) + elif self.mambaish_config: + token_to_kv_pool = self._build_hybrid_linear_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + req_to_token_pool=req_to_token_pool, + mha_pool_class=mha_pool_class, + ) + else: + if is_float4_e2m1fn_x2(self.kv_cache_dtype): + assert ( + not enable_page_major + ), "page-major KV layout is not supported with fp4 KV cache" + token_to_kv_pool = self._build_mha_fp4_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + ) + else: + token_to_kv_pool = self._build_mha_kv_pool( + max_total_num_tokens=sizes.max_total_num_tokens, + mha_pool_class=mha_pool_class, + ) + return token_to_kv_pool + + def _build_dsv4_kv_pool( + self, + *, + max_running_requests: int, + swa_max_total_num_tokens: Optional[int], + c4_max_total_num_tokens: int, + c128_max_total_num_tokens: int, + c4_state_pool_size: int, + c128_state_pool_size: int, + c4_state_dtype: Optional[torch.dtype], + c128_state_dtype: Optional[torch.dtype], + req_to_token_pool: ReqToTokenPool, + ) -> KVCache: + swa_page_size = self.server_args.page_size + if not _is_npu: + assert swa_page_size == 256, "In paged swa mode, page_size must be 256." + + if self.is_draft_worker: + from sglang.srt.models.deepseek_v4_nextn import ( + COMPRESS_RATIO_NEXTN_LAYER, + ) + + compression_ratios = [ + COMPRESS_RATIO_NEXTN_LAYER + ] * self.layer_info.num_effective_layers + else: + compression_ratios = self.model_config.compress_ratios + + # NPU + DSV4 → paged-state subclass: the fused compressor kernel + # needs cache_mode=1 (paged); Atlas A3 rejects cache_mode=2 (ring), + # so the CUDA ring-buffer state path can't be shared. CUDA keeps + # DeepSeekV4TokenToKVPool unchanged; NPU recomputes state sizes below. + if _is_npu: + from sglang.srt.hardware_backend.npu.dsv4.dsv4_memory_pool import ( + DSV4NPUTokenToKVPool, + npu_state_pool_size, + ) + + pool_cls = DSV4NPUTokenToKVPool + # Recompute state pool sizes for the NPU paged formula (CUDA's + # ring sizes are dropped here). Tail-only allocation keeps the + # per-req-budget formula sufficient at any prefill length: long + # prompts allocate only ``tail+128`` (c4) / ``tail`` (c128) + # slots (tail = seq_len % 128), and decode is drained by + # sliding eviction in ``ScheduleBatch._evict_swa``. + c4_state_pool_size = npu_state_pool_size( + ratio=4, + page_size=self.server_args.page_size, + max_num_reqs=max_running_requests, + ) + c128_state_pool_size = npu_state_pool_size( + ratio=128, + page_size=self.server_args.page_size, + max_num_reqs=max_running_requests, + ) + else: + pool_cls = DeepSeekV4TokenToKVPool + c4_state_pool_size = c4_state_pool_size + c128_state_pool_size = c128_state_pool_size + + token_to_kv_pool = pool_cls( + max_num_reqs=max_running_requests, + # SWA ring is indexed by req_pool_idx; PD decode inflates req_to_token + # past max_running_requests (pre-alloc), so size to the real capacity. + num_req_slots=req_to_token_pool.req_to_token.shape[0], + swa_size=swa_max_total_num_tokens, + c4_size=c4_max_total_num_tokens, + c128_size=c128_max_total_num_tokens, + c4_state_pool_size=c4_state_pool_size, + c128_state_pool_size=c128_state_pool_size, + page_size=self.server_args.page_size, + swa_page_size=swa_page_size, + sliding_window=self.model_config.window_size, + dtype=self.kv_cache_dtype, + c4_state_dtype=c4_state_dtype, + c128_state_dtype=c128_state_dtype, + qk_nope_head_dim=self.model_config.qk_nope_head_dim, + qk_rope_head_dim=self.model_config.qk_rope_head_dim, + indexer_head_dim=self.model_config.index_head_dim, + layer_num=self.layer_info.num_effective_layers, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + compression_ratios=compression_ratios, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + enable_hisparse=self.server_args.enable_hisparse, + online_mtp_max_draft_tokens=( + self.server_args.max_speculative_num_draft_tokens or 0 + ), + ) + return token_to_kv_pool + + def _build_oot_dsa_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: + PoolCls = current_platform.get_dsa_kv_pool_cls() + token_to_kv_pool = PoolCls( + max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + kv_lora_rank=self.model_config.kv_lora_rank, + qk_rope_head_dim=self.model_config.qk_rope_head_dim, + layer_num=self.layer_info.num_effective_layers, + device=self.device, + kv_cache_dim=calculate_mla_kv_cache_dim( + model_config=self.model_config, + kv_cache_dtype=self.kv_cache_dtype, + server_args=self.server_args, + ), + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config), + ) + return token_to_kv_pool + + def _build_oot_mla_kv_pool( + self, *, max_total_num_tokens: int, is_dsa_model: bool + ) -> KVCache: + PoolCls = current_platform.get_mla_kv_pool_cls() + token_to_kv_pool = PoolCls( + max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + kv_lora_rank=self.model_config.kv_lora_rank, + qk_rope_head_dim=self.model_config.qk_rope_head_dim, + index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None), + layer_num=self.layer_info.num_effective_layers, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + ) + return token_to_kv_pool + + def _build_oot_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: + PoolCls = current_platform.get_mha_kv_pool_cls() + token_to_kv_pool = PoolCls( + max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), + head_dim=self.model_config.head_dim, + layer_num=self.layer_info.num_effective_layers, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + ) + return token_to_kv_pool + + def _build_ascend_swa_kv_pool( + self, + *, + full_max_total_num_tokens: Optional[int], + swa_max_total_num_tokens: Optional[int], + ) -> KVCache: + from sglang.srt.hardware_backend.npu.memory_pool_npu import ( + NPUMHATokenToKVPool, + ) + + kwargs = {} + if self.is_hybrid_swa_compress: + kwargs = { + "swa_head_num": max( + 1, + self.model_config.hf_text_config.swa_num_key_value_heads + // get_parallel().attn_tp_size, + ), + "swa_head_dim": self.model_config.swa_head_dim, + "swa_v_head_dim": self.model_config.swa_v_head_dim, + "v_head_dim": self.model_config.v_head_dim, + } + token_to_kv_pool = SWAKVPool( + size=full_max_total_num_tokens, + size_swa=swa_max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + post_capture_active=self.post_capture_kv_active, + head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), + head_dim=self.model_config.head_dim, + swa_attention_layer_ids=self.model_config.swa_attention_layer_ids, + full_attention_layer_ids=self.model_config.full_attention_layer_ids, + device=self.device, + token_to_kv_pool_class=NPUMHATokenToKVPool, + **kwargs, + ) + return token_to_kv_pool + + def _build_ascend_mla_kv_pool( + self, *, max_total_num_tokens: int, is_dsa_model: bool + ) -> KVCache: + from sglang.srt.hardware_backend.npu.memory_pool_npu import ( + NPUMLATokenToKVPool, + ) + + token_to_kv_pool = NPUMLATokenToKVPool( + max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + kv_lora_rank=self.model_config.kv_lora_rank, + qk_rope_head_dim=self.model_config.qk_rope_head_dim, + index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None), + layer_num=self.layer_info.num_effective_layers, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + ) + return token_to_kv_pool + + def _build_ascend_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: + from sglang.srt.hardware_backend.npu.memory_pool_npu import ( + NPUMHATokenToKVPool, + ) + + token_to_kv_pool = NPUMHATokenToKVPool( + max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), + head_dim=self.model_config.head_dim, + layer_num=self.layer_info.num_effective_layers, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + ) + return token_to_kv_pool + + def _build_dsa_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: + from sglang.srt.layers.cp.utils import get_glm_dsa_cp_layer_shard_info + + ( + dsa_cp_layer_shard_rank, + dsa_cp_layer_shard_size, + ) = get_glm_dsa_cp_layer_shard_info(self) + pool_kwargs = {} + if self.server_args.enable_hisparse: + PoolCls = HiSparseDSATokenToKVPool + from sglang.srt.mem_cache.sparsity import parse_hisparse_config + + pool_kwargs["host_to_device_ratio"] = parse_hisparse_config( + self.server_args + ).host_to_device_ratio + elif dsa_cp_layer_shard_rank is not None: + # DSA cache layer split: shard KV/indexer layers across CP ranks. + from sglang.srt.mem_cache.dsa_cache_layer_split import ( + LayerSplitDSATokenToKVPool, + ) + + PoolCls = LayerSplitDSATokenToKVPool + pool_kwargs["layer_shard_rank"] = dsa_cp_layer_shard_rank + pool_kwargs["layer_shard_size"] = dsa_cp_layer_shard_size + else: + PoolCls = DSATokenToKVPool + token_to_kv_pool = PoolCls( + max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + kv_lora_rank=self.model_config.kv_lora_rank, + qk_rope_head_dim=self.model_config.qk_rope_head_dim, + layer_num=self.layer_info.num_effective_layers, + device=self.device, + kv_cache_dim=calculate_mla_kv_cache_dim( + model_config=self.model_config, + kv_cache_dtype=self.kv_cache_dtype, + server_args=self.server_args, + ), + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config), + **pool_kwargs, + ) + return token_to_kv_pool + + def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: + token_to_kv_pool = MLATokenToKVPoolFP4( + max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + kv_lora_rank=self.model_config.kv_lora_rank, + qk_rope_head_dim=self.model_config.qk_rope_head_dim, + layer_num=self.layer_info.num_effective_layers, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + ) + return token_to_kv_pool + + def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: + token_to_kv_pool = MLATokenToKVPool( + max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + kv_lora_rank=self.model_config.kv_lora_rank, + qk_rope_head_dim=self.model_config.qk_rope_head_dim, + layer_num=self.layer_info.num_effective_layers, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + ) + return token_to_kv_pool + + def _build_hybrid_swa_kv_pool( + self, + *, + full_max_total_num_tokens: Optional[int], + swa_max_total_num_tokens: Optional[int], + mha_pool_class: type, + ) -> KVCache: + kwargs = {} + if self.is_hybrid_swa_compress: + kwargs = { + "swa_head_num": max( + 1, + self.model_config.hf_text_config.swa_num_key_value_heads + // get_parallel().attn_tp_size, + ), + "swa_head_dim": self.model_config.swa_head_dim, + "swa_v_head_dim": self.model_config.swa_v_head_dim, + "v_head_dim": self.model_config.v_head_dim, + } + token_to_kv_pool = SWAKVPool( + size=full_max_total_num_tokens, + size_swa=swa_max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), + head_dim=self.model_config.head_dim, + swa_attention_layer_ids=self.model_config.swa_attention_layer_ids, + full_attention_layer_ids=self.model_config.full_attention_layer_ids, + device=self.device, + enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), + token_to_kv_pool_class=mha_pool_class, + **kwargs, + ) + return token_to_kv_pool + + def _build_minimax_sparse_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: + _hf_config = self.model_config.hf_config + sparse_cfg = get_minimax_sparse_attention_config(_hf_config) + dense_layer_ids, sparse_layer_ids = get_minimax_sparse_layer_ids(sparse_cfg) + disable_value_sparse_layer_ids = get_minimax_sparse_disable_value_layer_ids( + sparse_cfg + ) + token_to_kv_pool = MiniMaxSparseKVPool( + size=max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + index_dtype=self.model_dtype, + head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), + head_dim=self.model_config.head_dim, + idx_head_dim=sparse_cfg["sparse_index_dim"], + dense_layer_ids=dense_layer_ids, + sparse_layer_ids=sparse_layer_ids, + disable_value_sparse_layer_ids=disable_value_sparse_layer_ids, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + ) + return token_to_kv_pool + + def _build_hybrid_linear_kv_pool( + self, + *, + max_total_num_tokens: int, + req_to_token_pool: ReqToTokenPool, + mha_pool_class: type, + ) -> KVCache: + extra_args = {} + if self.use_mla_backend: + extra_args = { + "kv_lora_rank": self.model_config.kv_lora_rank, + "qk_rope_head_dim": self.model_config.qk_rope_head_dim, + } + token_to_kv_pool = HybridLinearKVPool( + page_size=self.server_args.page_size, + size=max_total_num_tokens, + dtype=self.kv_cache_dtype, + head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), + head_dim=self.model_config.head_dim, + # if draft worker, we only need 1 attention layer's kv pool + full_attention_layer_ids=( + [0] + if self.is_draft_worker + else [ + i + for i in self.mambaish_config.full_attention_layer_ids + if self.layer_info.start_layer <= i < self.layer_info.end_layer + ] + ), + device=self.device, + mamba_pool=req_to_token_pool.mamba_pool, + enable_memory_saver=self.server_args.enable_memory_saver, + enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), + use_mla=self.use_mla_backend, + start_layer=self.layer_info.start_layer, + full_kv_pool_class=mha_pool_class, + post_capture_active=self.post_capture_kv_active, + **extra_args, + ) + return token_to_kv_pool + + def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: + token_to_kv_pool = MHATokenToKVPoolFP4( + max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), + head_dim=self.model_config.head_dim, + v_head_dim=self.model_config.v_head_dim, + layer_num=self.layer_info.num_effective_layers, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + enable_alt_stream=not self.server_args.enable_pdmux, + enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), + ) + return token_to_kv_pool + + def _build_mha_kv_pool( + self, *, max_total_num_tokens: int, mha_pool_class: type + ) -> KVCache: + pool_cls = ( + NoOpMHATokenToKVPool + if self.server_args.prefill_only_disable_kv_cache + else mha_pool_class + ) + token_to_kv_pool = pool_cls( + max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), + head_dim=self.model_config.head_dim, + v_head_dim=self.model_config.v_head_dim, + layer_num=self.layer_info.num_effective_layers, + device=self.device, + enable_memory_saver=self.server_args.enable_memory_saver, + start_layer=self.layer_info.start_layer, + end_layer=self.layer_info.end_layer, + enable_alt_stream=not self.server_args.enable_pdmux, + enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), + post_capture_active=self.post_capture_kv_active, + ) + return token_to_kv_pool + + def _build_token_to_kv_pool_allocator( + self, + *, + sizes: _PoolSizes, + token_to_kv_pool: KVCache, + is_dsv4_model: bool, + req_to_token_pool: ReqToTokenPool, + token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator], + ) -> BaseTokenToKVPoolAllocator: + # Initialize token_to_kv_pool_allocator + need_sort = self.server_args.disaggregation_mode in ("decode", "prefill") + if token_to_kv_pool_allocator is None: + if current_platform.is_out_of_tree(): + AllocatorCls = current_platform.get_paged_allocator_cls() + token_to_kv_pool_allocator = AllocatorCls( + sizes.max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + device=self.device, + kvcache=token_to_kv_pool, + need_sort=need_sort, + ) + elif _is_npu and ( + self.server_args.attention_backend == "ascend" + or is_dsv4_model + or self.hybrid_gdn_config is not None + ): + if self.is_hybrid_swa: + # DSV4 on NPU: SWA allocator subclass that also drives the + # c4/c128 allocators, producing a DSV4OutCacheLoc per alloc. + if is_dsv4_model: + from sglang.srt.hardware_backend.npu.dsv4.dsv4_allocator import ( + DSV4NPUTokenToKVPoolAllocator, + ) + + swa_allocator_cls = DSV4NPUTokenToKVPoolAllocator + else: + swa_allocator_cls = SWATokenToKVPoolAllocator + token_to_kv_pool_allocator = swa_allocator_cls( + sizes.full_max_total_num_tokens, + sizes.swa_max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + device=self.device, + kvcache=token_to_kv_pool, + need_sort=need_sort, + ) + else: + from sglang.srt.hardware_backend.npu.allocator_npu import ( + NPUPagedTokenToKVPoolAllocator, + ) + + token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator( + sizes.max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + device=self.device, + kvcache=token_to_kv_pool, + need_sort=need_sort, + ) + else: + if self.is_hybrid_swa and sizes.full_max_total_num_tokens == 0: + token_to_kv_pool_allocator = PureSWATokenToKVPoolAllocator( + sizes.swa_max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + device=self.device, + kvcache=token_to_kv_pool, + need_sort=need_sort, + ) + elif self.is_hybrid_swa: + token_to_kv_pool_allocator = SWATokenToKVPoolAllocator( + sizes.full_max_total_num_tokens, + sizes.swa_max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + device=self.device, + kvcache=token_to_kv_pool, + need_sort=need_sort, + ) + else: + if self.server_args.enable_hisparse: + from sglang.srt.mem_cache.sparsity import ( + parse_hisparse_config, + ) + + hisparse_cfg = parse_hisparse_config(self.server_args) + token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator( + sizes.max_total_num_tokens, + page_size=self.server_args.page_size, + dtype=self.kv_cache_dtype, + device=self.device, + kvcache=token_to_kv_pool, + need_sort=need_sort, + host_to_device_ratio=hisparse_cfg.host_to_device_ratio, + ) + elif ( + self.server_args.page_size == 1 + and self.server_args.dcp_size == 1 + ): + token_to_kv_pool_allocator = TokenToKVPoolAllocator( + sizes.max_total_num_tokens, + dtype=self.kv_cache_dtype, + device=self.device, + kvcache=token_to_kv_pool, + need_sort=need_sort, + ) + else: + token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator( + sizes.max_total_num_tokens * self.server_args.dcp_size, + page_size=self.server_args.page_size + * self.server_args.dcp_size, + dtype=self.kv_cache_dtype, + device=self.device, + kvcache=token_to_kv_pool, + need_sort=need_sort, + ) + + if self.server_args.enable_hisparse and is_dsv4_model: + assert self.is_hybrid_swa, "DeepSeek V4 HiSparse requires SWA mode." + token_to_kv_pool_allocator = DeepSeekV4HiSparseTokenToKVPoolAllocator( + token_to_kv_pool_allocator + ) + + # DSV4-NPU: wire allocator back-ref into req_to_token_pool so its + # free(req) can release c4/c128 pool pages alongside the slot. + if hasattr(req_to_token_pool, "register_dsv4_allocator"): + req_to_token_pool.register_dsv4_allocator(token_to_kv_pool_allocator) + + else: + assert self.is_draft_worker + if self.is_hybrid_swa: + swa_allocator = getattr( + token_to_kv_pool_allocator, + "logical_attn_allocator", + token_to_kv_pool_allocator, + ) + assert isinstance(swa_allocator, SWATokenToKVPoolAllocator) + token_to_kv_pool.register_mapping( + swa_allocator.full_to_swa_index_mapping + ) + return token_to_kv_pool_allocator + def _profile_available_bytes(self, pre_model_load_memory: int) -> int: # KV pool budget = currently-free GPU memory minus the non-static runtime # slack (pre_model_load_memory * (1 - mem_fraction_static)). Whatever is @@ -1383,7 +1587,7 @@ class KVCacheConfigurator: if server_args.max_mamba_cache_size is not None: # Use explicitly set max_mamba_cache_size server_args.override( - "kv_cache_configurator.max_mamba_cache_size", + "mamba_pool.per_dp_shard", max_mamba_cache_size=server_args.max_mamba_cache_size // self.ps.attn_dp_size, ) @@ -1406,7 +1610,7 @@ class KVCacheConfigurator: ): # Use explicitly set max_running_requests when radix cache is disabled server_args.override( - "kv_cache_configurator.max_mamba_cache_size", + "mamba_pool.from_max_running_requests", max_mamba_cache_size=server_args.max_running_requests // self.ps.attn_dp_size, ) @@ -1440,7 +1644,7 @@ class KVCacheConfigurator: D = server_args.speculative_num_draft_tokens # Joint solve: main_state + intermediate = mamba_budget server_args.override( - "kv_cache_configurator.max_mamba_cache_size", + "mamba_pool.memory_budget_spec", max_mamba_cache_size=int( mamba_budget_bytes // (per_req * (1 + D / ratio)) ), @@ -1455,7 +1659,7 @@ class KVCacheConfigurator: total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) else: server_args.override( - "kv_cache_configurator.max_mamba_cache_size", + "mamba_pool.memory_budget", max_mamba_cache_size=int(mamba_budget_bytes // per_req), ) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 69e6fffe8..94a9bcb50 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -25,10 +25,6 @@ from typing import Optional, Union import torch -from sglang.srt.configs.hybrid_arch import ( - hybrid_gdn_config, - mambaish_config, -) from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.model_config import ( AttentionArch, @@ -161,9 +157,6 @@ from sglang.srt.model_executor.model_runner_components.weight_exporter import ( from sglang.srt.model_executor.model_runner_components.weight_updater import ( WeightUpdater, ) -from sglang.srt.model_executor.model_runner_kv_cache_mixin import ( - ModelRunnerKVCacheMixin, -) from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.runner import ( EagerRunner, @@ -246,7 +239,7 @@ class ModelRunnerOutput: indexer_topk_output: Optional[TopkCaptureOutput] = None -class ModelRunner(ModelRunnerKVCacheMixin): +class ModelRunner: """ModelRunner runs the forward passes of the models.""" def __init__( @@ -469,24 +462,23 @@ class ModelRunner(ModelRunnerKVCacheMixin): device=self.device, gpu_id=self.gpu_id, ps=self.ps, + pp_group=self.pp_group, model_config=self.model_config, server_args=self.server_args, kv_cache_dtype=self.kv_cache_dtype, + model_dtype=self.dtype, page_size=self.page_size, + sliding_window_size=self.sliding_window_size, spec_algorithm=self.spec_algorithm, is_draft_worker=self.is_draft_worker, post_capture_kv_active=is_post_capture_kv_active( server_args=self.server_args, is_draft_worker=self.is_draft_worker ), - dflash_draft_num_layers=self.spec_aux_config.dflash_draft_num_layers, + spec_aux_config=self.spec_aux_config, is_hybrid_swa=self.is_hybrid_swa, is_hybrid_swa_compress=self.is_hybrid_swa_compress, use_mla_backend=self.use_mla_backend, - mambaish_config=mambaish_config(self.model_config), - hybrid_gdn_config=hybrid_gdn_config(self.model_config), - start_layer=self.layer_info.start_layer, - end_layer=self.layer_info.end_layer, - num_effective_layers=self.layer_info.num_effective_layers, + layer_info=self.layer_info, forward_stream=self.forward_stream, req_to_token_pool=self.req_to_token_pool, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, @@ -656,7 +648,20 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.memory_pool_config = memory_pool_config self.init_kv_cache_configurator() - self.init_memory_pool(self.pre_model_load_memory) + result = self.kv_cache_configurator.configure( + pre_model_load_memory=self.pre_model_load_memory + ) + self.max_total_num_tokens = result.max_total_num_tokens + self.max_running_requests = result.max_running_requests + self.req_to_token_pool = result.req_to_token_pool + self.token_to_kv_pool = result.token_to_kv_pool + self.token_to_kv_pool_allocator = result.token_to_kv_pool_allocator + self.memory_pool_config = result.memory_pool_config + if self.is_hybrid_swa: + self.full_max_total_num_tokens = result.full_max_total_num_tokens + self.swa_max_total_num_tokens = result.swa_max_total_num_tokens + # Keep a reference so the shared byte buffer is not GC'd. + self._unified_memory_pool = result.unified_memory_pool # Must be called AFTER init_memory_pool so the pool object exists for # canary to monkey-patch, and BEFORE init_decode_cuda_graph so warmup 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 deleted file mode 100644 index 554a951dd..000000000 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ /dev/null @@ -1,24 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -if TYPE_CHECKING: - from sglang.srt.model_executor.model_runner import ModelRunner - - -class ModelRunnerKVCacheMixin: - def init_memory_pool(self: ModelRunner, pre_model_load_memory: int): - result = self.kv_cache_configurator.configure( - pre_model_load_memory=pre_model_load_memory - ) - self.max_total_num_tokens = result.max_total_num_tokens - self.max_running_requests = result.max_running_requests - self.req_to_token_pool = result.req_to_token_pool - self.token_to_kv_pool = result.token_to_kv_pool - self.token_to_kv_pool_allocator = result.token_to_kv_pool_allocator - self.memory_pool_config = result.memory_pool_config - if self.is_hybrid_swa: - self.full_max_total_num_tokens = result.full_max_total_num_tokens - self.swa_max_total_num_tokens = result.swa_max_total_num_tokens - # Keep a reference so the shared byte buffer is not GC'd. - self._unified_memory_pool = result.unified_memory_pool diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index b9f5cc8c1..81d0c9eab 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -68,7 +68,7 @@ class MemoryPoolConfig: if TYPE_CHECKING: - from sglang.srt.model_executor.model_runner import ModelRunner + from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator logger = logging.getLogger(__name__) @@ -123,28 +123,28 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): bias = 0 """ - def __init__(self, mr: ModelRunner): + def __init__(self, kvc: KVCacheConfigurator): # Determine effective number of layers for KV cache - if mambaish := mambaish_config(mr.model_config): + if mambaish := mambaish_config(kvc.model_config): effective_layer_ids = [ i for i in mambaish.full_attention_layer_ids - if mr.start_layer <= i < mr.end_layer + if kvc.layer_info.start_layer <= i < kvc.layer_info.end_layer ] num_layers = len(effective_layer_ids) else: - num_layers = mr.num_effective_layers + num_layers = kvc.layer_info.num_effective_layers - self._cell_size = self._compute_cell_size(mr, num_layers) + self._cell_size = self._compute_cell_size(kvc, num_layers) # EAGLE/STANDALONE: scale cell_size to account for draft model KV cache. # Assumes draft and target share the same per-layer KV size (head_dim, # num_kv_heads, dtype), which holds for EAGLE/MTP draft models that # reuse the target architecture's attention config. if ( - mr.spec_algorithm.is_eagle() or mr.spec_algorithm.is_standalone() - ) and not mr.is_draft_worker: - eagle_draft_num_layers = mr.spec_aux_config.eagle_draft_num_layers + kvc.spec_algorithm.is_eagle() or kvc.spec_algorithm.is_standalone() + ) and not kvc.is_draft_worker: + eagle_draft_num_layers = kvc.spec_aux_config.eagle_draft_num_layers if ( eagle_draft_num_layers is not None and int(eagle_draft_num_layers) > 0 @@ -156,12 +156,12 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): ) # DFLASH/DSPARK: scale cell_size to account for draft model KV cache - if mr.spec_algorithm.is_dflash_family() and not mr.is_draft_worker: + if kvc.spec_algorithm.is_dflash_family() and not kvc.is_draft_worker: from sglang.srt.speculative.dflash_utils import ( scale_kv_cell_size_per_token_for_dflash, ) - draft_num_layers = mr.spec_aux_config.dflash_draft_num_layers + draft_num_layers = kvc.spec_aux_config.dflash_draft_num_layers if ( draft_num_layers is not None and int(draft_num_layers) > 0 @@ -173,23 +173,23 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): draft_num_layers=int(draft_num_layers), ) - def _compute_cell_size(self, mr: ModelRunner, num_layers: int) -> int: + def _compute_cell_size(self, kvc: KVCacheConfigurator, num_layers: int) -> int: """Compute per-token KV cache cost in bytes. Subclasses can override.""" # args to config cell size - model_config = mr.model_config - kv_cache_dtype = mr.kv_cache_dtype + model_config = kvc.model_config + kv_cache_dtype = kvc.kv_cache_dtype from sglang.srt.layers.cp.utils import ( get_glm_dsa_layer_split_effective_num_layers, ) effective_num_layers = get_glm_dsa_layer_split_effective_num_layers( - mr, num_layers + kvc, num_layers ) kv_size = torch._utils._element_size(kv_cache_dtype) tp_size = get_parallel().attn_tp_size - if mr.use_mla_backend: + if kvc.use_mla_backend: cell_size = ( (model_config.kv_lora_rank + model_config.qk_rope_head_dim) * effective_num_layers @@ -230,10 +230,14 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): ) local_dense_layer_ids = [ - l for l in dense_layer_ids if mr.start_layer <= l < mr.end_layer + l + for l in dense_layer_ids + if kvc.layer_info.start_layer <= l < kvc.layer_info.end_layer ] local_sparse_layer_ids = [ - l for l in sparse_layer_ids if mr.start_layer <= l < mr.end_layer + l + for l in sparse_layer_ids + if kvc.layer_info.start_layer <= l < kvc.layer_info.end_layer ] num_dense = len(local_dense_layer_ids) num_sparse = len(local_sparse_layer_ids) @@ -245,7 +249,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): kv_heads = model_config.get_num_kv_heads(get_parallel().attn_tp_size) head_dim = model_config.head_dim indexer_head_dim = sparse_cfg["sparse_index_dim"] - indexer_dtype_size = torch._utils._element_size(mr.dtype) + indexer_dtype_size = torch._utils._element_size(kvc.model_dtype) main_pool_bytes = ( (num_dense + num_sparse) * 2 * kv_heads * head_dim * kv_size @@ -298,9 +302,9 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator): Does NOT inherit DefaultPoolConfigurator — different coeff model. """ - def __init__(self, mr: ModelRunner): - model_config = mr.model_config - kv_cache_dtype = mr.kv_cache_dtype + def __init__(self, kvc: KVCacheConfigurator): + model_config = kvc.model_config + kv_cache_dtype = kvc.kv_cache_dtype kv_size = torch._utils._element_size(kv_cache_dtype) tp_size = get_parallel().attn_tp_size @@ -310,7 +314,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator): self._swa_layers_num > 0 ), "Hybrid SWA model must have at least one SWA layer" - self._swa_full_tokens_ratio = mr.server_args.swa_full_tokens_ratio + self._swa_full_tokens_ratio = kvc.server_args.swa_full_tokens_ratio # Full layer per-token memory (bytes) self._full_per_token = ( @@ -330,9 +334,9 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator): # full-attn layers; budget into the full term. self._draft_full_layers_num = 0 if ( - mr.spec_algorithm.is_eagle() or mr.spec_algorithm.is_standalone() - ) and not mr.is_draft_worker: - draft_layers = mr.spec_aux_config.eagle_draft_num_layers + kvc.spec_algorithm.is_eagle() or kvc.spec_algorithm.is_standalone() + ) and not kvc.is_draft_worker: + draft_layers = kvc.spec_aux_config.eagle_draft_num_layers if draft_layers is not None and int(draft_layers) > 0: self._draft_full_layers_num = int(draft_layers) @@ -417,13 +421,13 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator): both pools by swa_full_tokens_ratio. """ - def __init__(self, mr: ModelRunner): - super().__init__(mr) + def __init__(self, kvc: KVCacheConfigurator): + super().__init__(kvc) assert self._full_layers_num > 0 - sa = mr.server_args - page_size = mr.page_size - window = mr.sliding_window_size + sa = kvc.server_args + page_size = kvc.page_size + window = kvc.sliding_window_size draft_tokens = sa.speculative_num_draft_tokens or 1 eviction_interval = max(1, envs.SGLANG_SWA_EVICTION_INTERVAL.get()) @@ -443,7 +447,7 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator): decode_alloc = 2 * get_alloc_len_per_decode(sa) per_request = trailing_tokens + decode_alloc - num_reqs = sa.max_running_requests // mr.dp_size + num_reqs = sa.max_running_requests // kvc.ps.attn_dp_size if sa.disaggregation_mode == "decode": self._swa_cap = ( per_request * num_reqs @@ -458,18 +462,18 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator): ) @staticmethod - def is_applicable(mr: ModelRunner) -> bool: + def is_applicable(kvc: KVCacheConfigurator) -> bool: """True when SWAChunkCache can be sized from explicit max requests.""" - sa = mr.server_args + sa = kvc.server_args if sa.max_running_requests is None: return False if not sa.disable_radix_cache: return False if sa.chunked_prefill_size is None: return False - if mr.sliding_window_size is None: + if kvc.sliding_window_size is None: return False - return len(mr.model_config.full_attention_layer_ids) > 0 + return len(kvc.model_config.full_attention_layer_ids) > 0 def calculate_pool_sizes( self, available_bytes: int, page_size: int @@ -528,40 +532,42 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): decode reserves a draft worker, mirroring dflash's cell_size scaling); bias = 0. """ - def __init__(self, mr: ModelRunner): - cfg = mr.model_config + def __init__(self, kvc: KVCacheConfigurator): + cfg = kvc.model_config self.qk_nope_head_dim = cfg.qk_nope_head_dim self.qk_rope_head_dim = cfg.qk_rope_head_dim self.indexer_head_dim = cfg.index_head_dim - self.context_len = mr.model_config.context_len + self.context_len = kvc.model_config.context_len # PP-local slice; matches DeepSeekV4TokenToKVPool's stage_ratios. - self.compression_ratios = cfg.compress_ratios[mr.start_layer : mr.end_layer] - if mr.pp_size > 1: + self.compression_ratios = cfg.compress_ratios[ + kvc.layer_info.start_layer : kvc.layer_info.end_layer + ] + if kvc.ps.pp_size > 1: logger.info( - f"DSV4 pool PP slice: rank={mr.pp_group.rank_in_group} " - f"layers=[{mr.start_layer},{mr.end_layer}) " + f"DSV4 pool PP slice: rank={kvc.pp_group.rank_in_group} " + f"layers=[{kvc.layer_info.start_layer},{kvc.layer_info.end_layer}) " f"local={len(self.compression_ratios)}/{len(cfg.compress_ratios)}" ) self.swa_page_size = cfg.window_size - self.swa_ratio = mr.server_args.swa_full_tokens_ratio - self.is_speculative = mr.server_args.speculative_algorithm is not None + self.swa_ratio = kvc.server_args.swa_full_tokens_ratio + self.is_speculative = kvc.server_args.speculative_algorithm is not None self.online_c128_mtp_max_draft_tokens = ( - mr.server_args.max_speculative_num_draft_tokens or 0 + kvc.server_args.max_speculative_num_draft_tokens or 0 ) self.requested_max_running_requests_per_worker = ( - mr.server_args.max_running_requests // mr.dp_size - if mr.server_args.max_running_requests is not None + kvc.server_args.max_running_requests // kvc.ps.attn_dp_size + if kvc.server_args.max_running_requests is not None else None ) - self.disaggregation_mode = mr.server_args.disaggregation_mode + self.disaggregation_mode = kvc.server_args.disaggregation_mode self.disaggregation_decode_extra_slots = ( - mr.server_args.disaggregation_decode_extra_slots or 0 + kvc.server_args.disaggregation_decode_extra_slots or 0 ) - if mr.server_args.enable_hisparse: + if kvc.server_args.enable_hisparse: from sglang.srt.mem_cache.sparsity import parse_hisparse_config self.c4_shrink_factor = parse_hisparse_config( - mr.server_args + kvc.server_args ).host_to_device_ratio else: self.c4_shrink_factor = 1 @@ -593,9 +599,9 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): if envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get(): allow_experimental_online_c128_mtp = ( envs.SGLANG_EXPERIMENTAL_ONLINE_C128_MTP.get() - and mr.spec_algorithm.is_eagle() + and kvc.spec_algorithm.is_eagle() ) - assert mr.spec_algorithm.is_none() or allow_experimental_online_c128_mtp, ( + assert kvc.spec_algorithm.is_none() or allow_experimental_online_c128_mtp, ( "SGLANG_OPT_USE_ONLINE_COMPRESS does not support speculative decode " "(MTP) yet, except the experimental EAGLE topk=1 path gated by " "SGLANG_EXPERIMENTAL_ONLINE_C128_MTP=1" @@ -782,14 +788,14 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): def create_memory_pool_configurator( - mr: ModelRunner, + kvc: KVCacheConfigurator, ) -> MemoryPoolConfigurator: """Factory: select the right configurator for the model architecture.""" - if is_deepseek_v4(mr.model_config.hf_config) and mr.is_hybrid_swa: - return DSV4PoolConfigurator(mr) - if mr.is_hybrid_swa: - if SWAChunkCapPoolConfigurator.is_applicable(mr): - return SWAChunkCapPoolConfigurator(mr) - return HybridSWAPoolConfigurator(mr) + if is_deepseek_v4(kvc.model_config.hf_config) and kvc.is_hybrid_swa: + return DSV4PoolConfigurator(kvc) + if kvc.is_hybrid_swa: + if SWAChunkCapPoolConfigurator.is_applicable(kvc): + return SWAChunkCapPoolConfigurator(kvc) + return HybridSWAPoolConfigurator(kvc) # Future: MambaPoolConfigurator - return DefaultPoolConfigurator(mr) + return DefaultPoolConfigurator(kvc) diff --git a/test/registered/unit/model_executor/test_pool_configurator.py b/test/registered/unit/model_executor/test_pool_configurator.py index f8921facd..d341fe2ec 100644 --- a/test/registered/unit/model_executor/test_pool_configurator.py +++ b/test/registered/unit/model_executor/test_pool_configurator.py @@ -10,6 +10,7 @@ import unittest from types import SimpleNamespace from unittest.mock import MagicMock, patch +from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_cpu_ci @@ -95,6 +96,7 @@ def _make_model_runner( mc.get_swa_num_kv_heads = lambda tp_size: swa_num_kv_heads or num_kv_heads mc.hf_config = SimpleNamespace(architectures=["LlamaForCausalLM"]) mc.hf_config.get_text_config = lambda: mc.hf_config + mc.linear_attn_registry_result = None mr.model_config = mc mr.kv_cache_dtype = "fake_bf16" @@ -126,6 +128,15 @@ def _make_model_runner( spec.is_none.return_value = True mr.spec_algorithm = spec + mr.layer_info = SimpleNamespace( + start_layer=0, end_layer=num_layers, num_effective_layers=num_layers + ) + mr.ps = ParallelState.trivial() + mr.pp_group = SimpleNamespace(rank_in_group=0) + mr.spec_aux_config = SimpleNamespace( + eagle_draft_num_layers=None, dflash_draft_num_layers=None + ) + return mr @@ -517,7 +528,7 @@ class TestEagleConfigurator(unittest.TestCase): mr.spec_algorithm.is_eagle.return_value = True mr.spec_algorithm.is_standalone.return_value = False mr.spec_algorithm.is_none.return_value = False - mr.eagle_draft_num_layers = eagle_draft_num_layers + mr.spec_aux_config.eagle_draft_num_layers = eagle_draft_num_layers with mock_cpu_env(): from sglang.srt.model_executor.pool_configurator import (