diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 7f63610da..347805f63 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -54,7 +54,7 @@ from sglang.srt.weight_sync.tensor_bucket import FlattenedTensorBucket if TYPE_CHECKING: from sglang.srt.managers.cache_controller import LayerDoneCounter from sglang.srt.model_executor.model_runner import ModelRunner - from sglang.srt.model_executor.model_runner_kv_cache_mixin import MemoryPoolConfig + from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig logger = logging.getLogger(__name__) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 6ddb53ff1..bde3c135e 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -136,12 +136,12 @@ from sglang.srt.model_executor.forward_batch_info import ( ) from sglang.srt.model_executor.hook_manager import register_forward_hooks from sglang.srt.model_executor.model_runner_kv_cache_mixin import ( - MemoryPoolConfig, ModelRunnerKVCacheMixin, ) from sglang.srt.model_executor.piecewise_cuda_graph_runner import ( PiecewiseCudaGraphRunner, ) +from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( RemoteInstanceWeightLoaderBackend, 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 b701524fd..9ecc0c414 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 @@ -1,8 +1,7 @@ from __future__ import annotations import logging -from dataclasses import dataclass -from typing import TYPE_CHECKING, Optional, Tuple +from typing import TYPE_CHECKING import torch @@ -39,25 +38,7 @@ from sglang.srt.utils.common import ( if TYPE_CHECKING: from sglang.srt.model_executor.model_runner import ModelRunner - - -@dataclass -class MemoryPoolConfig: - """Resolved memory pool config, shared between target and draft workers.""" - - max_total_num_tokens: int - max_running_requests: int - full_max_total_num_tokens: Optional[int] = None - swa_max_total_num_tokens: Optional[int] = None - - mem_fraction_static: Optional[float] = None - - def __post_init__(self): - if self.max_total_num_tokens <= 0: - msg = "Not enough memory. Please try to increase --mem-fraction-static." - if self.mem_fraction_static is not None: - msg += f" Current value: mem_fraction_static={self.mem_fraction_static}" - raise RuntimeError(msg) + from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig # the ratio of mamba cache pool size to max_running_requests @@ -73,9 +54,7 @@ _is_hip = is_hip() class ModelRunnerKVCacheMixin: - def _profile_available_bytes( - self: ModelRunner, pre_model_load_memory: int - ) -> float: + def _profile_available_bytes(self: ModelRunner, pre_model_load_memory: int) -> int: post_model_load_memory = get_available_gpu_memory( self.device, self.gpu_id, @@ -89,48 +68,7 @@ class ModelRunnerKVCacheMixin: if self.mambaish_config is not None: rest_memory = self.handle_max_mamba_cache(rest_memory) - return rest_memory * (1 << 30) # return in bytes - - def profile_max_num_token(self: ModelRunner, pre_model_load_memory: int): - # Get the number of layers used for KV cache calculation - if self.is_draft_worker: - num_layers = getattr( - self.model_config.hf_config, - "num_nextn_predict_layers", - self.num_effective_layers, - ) - elif mambaish := self.mambaish_config: - effective_layer_ids = [ - i - for i in mambaish.full_attention_layer_ids - if self.start_layer <= i < self.end_layer - ] - num_layers = len(effective_layer_ids) - else: - num_layers = self.num_effective_layers - - from sglang.srt.model_executor.pool_configurator import get_cell_size_per_token - - cell_size = get_cell_size_per_token(self, num_layers) - if self.spec_algorithm.is_dflash() and not self.is_draft_worker: - from sglang.srt.speculative.dflash_utils import ( - scale_kv_cell_size_per_token_for_dflash, - ) - - draft_num_layers = getattr(self, "dflash_draft_num_layers", None) - if ( - draft_num_layers is not None - and int(draft_num_layers) > 0 - and int(num_layers) > 0 - ): - cell_size = scale_kv_cell_size_per_token_for_dflash( - target_cell_size_per_token=cell_size, - target_num_layers=int(num_layers), - draft_num_layers=int(draft_num_layers), - ) - - available_bytes = self._profile_available_bytes(pre_model_load_memory) - return int(available_bytes) // cell_size + return int(rest_memory * (1 << 30)) # return in bytes def handle_max_mamba_cache(self: ModelRunner, total_rest_memory): config = self.mambaish_config @@ -240,20 +178,6 @@ class ModelRunnerKVCacheMixin: return kv_cache_dim - def _resolve_hybrid_swa_tokens( - self: ModelRunner, token_capacity: int - ) -> Tuple[int, int, int]: - """Split token_capacity into full/swa pools. - - Returns (effective_capacity, full_max_total_num_tokens, swa_max_total_num_tokens). - """ - from sglang.srt.model_executor.pool_configurator import ( - resolve_hybrid_swa_tokens, - ) - - assert self.sliding_window_size is not None and self.sliding_window_size > 0 - return resolve_hybrid_swa_tokens(self, token_capacity) - def _calculate_mamba_ratio(self: ModelRunner) -> int: if self.server_args.disable_radix_cache: return 1 @@ -690,7 +614,11 @@ class ModelRunnerKVCacheMixin: ) def _apply_token_constraints(self: ModelRunner, token_capacity: int) -> int: - """Apply external constraints to token capacity: user cap, page alignment, PP sync.""" + """Apply external constraints to token capacity: user cap, PP sync. + + Page alignment is handled by the configurator, not here. + If constraints change the value, the configurator re-runs and re-aligns. + """ user_limit = self.server_args.max_total_tokens # Apply user-specified upper bound @@ -702,10 +630,6 @@ class ModelRunnerKVCacheMixin: ) token_capacity = min(token_capacity, user_limit) - # Align to page boundary - page_size = self.server_args.page_size - token_capacity = token_capacity // page_size * page_size - # Sync across PP ranks (each may have different layer counts) if self.pp_size > 1: tensor = torch.tensor(token_capacity, dtype=torch.int64) @@ -753,23 +677,28 @@ class ModelRunnerKVCacheMixin: self: ModelRunner, pre_model_load_memory: int ) -> MemoryPoolConfig: """Profile GPU memory and resolve all pool parameters into a config.""" - profiled_tokens = self.profile_max_num_token(pre_model_load_memory) - token_capacity = self._apply_token_constraints(profiled_tokens) + from sglang.srt.model_executor.pool_configurator import ( + create_memory_pool_configurator, + ) - full_tokens = None - swa_tokens = None - if self.is_hybrid_swa: - token_capacity, full_tokens, swa_tokens = self._resolve_hybrid_swa_tokens( - token_capacity + available_bytes = self._profile_available_bytes(pre_model_load_memory) + page_size = self.server_args.page_size + + configurator = create_memory_pool_configurator(self) + config = configurator.calculate_pool_sizes(available_bytes, page_size) + + # Apply external constraints (user cap, page alignment, PP sync) + constrained = self._apply_token_constraints(config.max_total_num_tokens) + if constrained != config.max_total_num_tokens: + config = configurator.calculate_pool_sizes_from_max_tokens( + constrained, page_size ) - return MemoryPoolConfig( - max_total_num_tokens=token_capacity, - max_running_requests=self._resolve_max_num_reqs(token_capacity), - full_max_total_num_tokens=full_tokens, - swa_max_total_num_tokens=swa_tokens, - mem_fraction_static=self.server_args.mem_fraction_static, + config.max_running_requests = self._resolve_max_num_reqs( + config.max_total_num_tokens ) + config.mem_fraction_static = self.server_args.mem_fraction_static + return config def init_memory_pool(self: ModelRunner, pre_model_load_memory: int): if not self.spec_algorithm.is_none() and self.is_draft_worker: diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index afb253b38..b6c250c16 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -1,7 +1,21 @@ +"""Memory pool configurators for profiling and sizing KV cache pools. + +Each model architecture has its own configurator that computes pool sizes +from available GPU memory using a unified coeff+bias model: + + available_bytes = max_tokens * coeff + bias + max_tokens = (available_bytes - bias) / coeff + +Two entry points, same core computation: +- calculate_pool_sizes(available_bytes, page_size): profiling path +- calculate_pool_sizes_from_max_tokens(max_tokens, page_size): constraint path +""" + from __future__ import annotations import logging -from typing import TYPE_CHECKING +from dataclasses import dataclass +from typing import TYPE_CHECKING, Optional import torch @@ -10,163 +24,265 @@ from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool from sglang.srt.utils.common import is_float4_e2m1fn_x2 + +@dataclass +class MemoryPoolConfig: + """Resolved memory pool config, shared between target and draft workers.""" + + max_total_num_tokens: int + max_running_requests: Optional[int] = None + full_max_total_num_tokens: Optional[int] = None + swa_max_total_num_tokens: Optional[int] = None + + mem_fraction_static: Optional[float] = None + + def __post_init__(self): + if self.max_total_num_tokens <= 0: + msg = "Not enough memory. Please try to increase --mem-fraction-static." + if self.mem_fraction_static is not None: + msg += f" Current value: mem_fraction_static={self.mem_fraction_static}" + raise RuntimeError(msg) + + if TYPE_CHECKING: from sglang.srt.model_executor.model_runner import ModelRunner logger = logging.getLogger(__name__) -def get_cell_size_per_token(mr: ModelRunner, num_layers: int) -> int: - # args to config cell size - model_config = mr.model_config - kv_cache_dtype = mr.kv_cache_dtype - use_mla_backend = mr.use_mla_backend +class MemoryPoolConfigurator: + """Base class for memory pool configurators. - kv_size = torch._utils._element_size(kv_cache_dtype) - if use_mla_backend: - cell_size = ( - (model_config.kv_lora_rank + model_config.qk_rope_head_dim) - * num_layers - * kv_size - ) - if is_float4_e2m1fn_x2(kv_cache_dtype): - # kv_scale_buffer - scale_block_size = 16 - cell_size = (cell_size // 2) + ( - ( - (model_config.kv_lora_rank + model_config.qk_rope_head_dim) - // scale_block_size + Subclasses compute pool sizes for their architecture via coeff+bias model. + Both entry points return MemoryPoolConfig (with max_running_requests=None, + to be filled by the consumer). + """ + + def calculate_pool_sizes( + self, available_bytes: int, page_size: int + ) -> MemoryPoolConfig: + """Profiling path: compute pool sizes from available bytes.""" + raise NotImplementedError + + def calculate_pool_sizes_from_max_tokens( + self, max_total_num_tokens: int, page_size: int + ) -> MemoryPoolConfig: + """Constraint path: recalculate pool sizes from a constrained max_tokens.""" + raise NotImplementedError + + +class DefaultPoolConfigurator(MemoryPoolConfigurator): + """Configurator for standard models: MHA, MLA, NSA, FP4. + + coeff = cell_size (bytes per token across all layers) + bias = 0 + """ + + def __init__(self, mr: ModelRunner): + # Determine effective number of layers for KV cache + if mambaish := mr.mambaish_config: + effective_layer_ids = [ + i + for i in mambaish.full_attention_layer_ids + if mr.start_layer <= i < mr.end_layer + ] + num_layers = len(effective_layer_ids) + else: + num_layers = mr.num_effective_layers + + self._cell_size = self._compute_cell_size(mr, num_layers) + + # DFLASH: scale cell_size to account for draft model KV cache + if mr.spec_algorithm.is_dflash() and not mr.is_draft_worker: + from sglang.srt.speculative.dflash_utils import ( + scale_kv_cell_size_per_token_for_dflash, + ) + + draft_num_layers = getattr(mr, "dflash_draft_num_layers", None) + if ( + draft_num_layers is not None + and int(draft_num_layers) > 0 + and int(num_layers) > 0 + ): + self._cell_size = scale_kv_cell_size_per_token_for_dflash( + target_cell_size_per_token=self._cell_size, + target_num_layers=int(num_layers), + draft_num_layers=int(draft_num_layers), ) + + def _compute_cell_size(self, mr: ModelRunner, 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 + + kv_size = torch._utils._element_size(kv_cache_dtype) + tp_size = get_attention_tp_size() + + if mr.use_mla_backend: + cell_size = ( + (model_config.kv_lora_rank + model_config.qk_rope_head_dim) * num_layers * kv_size ) + if is_float4_e2m1fn_x2(kv_cache_dtype): + # kv_scale_buffer + scale_block_size = 16 + cell_size = (cell_size // 2) + ( + ( + (model_config.kv_lora_rank + model_config.qk_rope_head_dim) + // scale_block_size + ) + * num_layers + * kv_size + ) - # Add indexer KV cache overhead for NSA models (DeepSeek V3.2) - if is_deepseek_nsa(model_config.hf_config): - index_head_dim = get_nsa_index_head_dim(model_config.hf_config) - indexer_size_per_token = ( - index_head_dim + index_head_dim // NSATokenToKVPool.quant_block_size * 4 - ) - element_size = torch._utils._element_size( - NSATokenToKVPool.index_k_with_scale_buffer_dtype - ) - cell_size += indexer_size_per_token * num_layers * element_size - else: - if model_config.is_hybrid_swa: - full_layers_num = len(model_config.full_attention_layer_ids) - swa_layers_num = len(model_config.swa_attention_layer_ids) - - full_per_token = model_config.get_num_kv_heads(get_attention_tp_size()) * ( - model_config.head_dim + model_config.v_head_dim - ) - - swa_per_token = model_config.get_swa_num_kv_heads( - get_attention_tp_size() - ) * (model_config.swa_head_dim + model_config.swa_v_head_dim) - - cell_size = ( - full_per_token * full_layers_num + swa_per_token * swa_layers_num - ) * kv_size + # Add indexer KV cache overhead for NSA models (DeepSeek V3.2) + if is_deepseek_nsa(model_config.hf_config): + index_head_dim = get_nsa_index_head_dim(model_config.hf_config) + indexer_size_per_token = ( + index_head_dim + + index_head_dim // NSATokenToKVPool.quant_block_size * 4 + ) + element_size = torch._utils._element_size( + NSATokenToKVPool.index_k_with_scale_buffer_dtype + ) + cell_size += indexer_size_per_token * num_layers * element_size else: cell_size = ( - model_config.get_num_kv_heads(get_attention_tp_size()) + model_config.get_num_kv_heads(tp_size) * (model_config.head_dim + model_config.v_head_dim) * num_layers * kv_size ) - if is_float4_e2m1fn_x2(kv_cache_dtype): - # kv_scale_buffer - scale_block_size = 16 + if is_float4_e2m1fn_x2(kv_cache_dtype): + # kv_scale_buffer + scale_block_size = 16 + n = model_config.get_num_kv_heads(tp_size) + k = model_config.head_dim + cell_size = (cell_size // 2) + ( + (n * k * num_layers * 2 * kv_size) // scale_block_size + ) - n = model_config.get_num_kv_heads(get_attention_tp_size()) - k = model_config.head_dim - cell_size = (cell_size // 2) + ( - (n * k * num_layers * 2 * kv_size) // scale_block_size - ) - return cell_size + return cell_size + + def calculate_pool_sizes( + self, available_bytes: int, page_size: int + ) -> MemoryPoolConfig: + max_total_num_tokens = available_bytes // self._cell_size + max_total_num_tokens = max_total_num_tokens // page_size * page_size + return MemoryPoolConfig(max_total_num_tokens=max_total_num_tokens) + + def calculate_pool_sizes_from_max_tokens( + self, max_total_num_tokens: int, page_size: int + ) -> MemoryPoolConfig: + max_total_num_tokens = max_total_num_tokens // page_size * page_size + return MemoryPoolConfig(max_total_num_tokens=max_total_num_tokens) -def resolve_hybrid_swa_tokens( - mr: ModelRunner, token_capacity: int -) -> tuple[int, int, int]: - """Split token_capacity into full/swa pools. +class HybridSWAPoolConfigurator(MemoryPoolConfigurator): + """Configurator for hybrid sliding window attention models (Gemma2, Command-R, MiMo). - Returns (effective_capacity, full_max_total_num_tokens, swa_max_total_num_tokens). + Splits available memory between full attention and SWA pools. + Does NOT inherit DefaultPoolConfigurator — different coeff model. """ - model_config = mr.model_config - page_size = mr.server_args.page_size - swa_full_tokens_ratio = mr.server_args.swa_full_tokens_ratio - full_layers_num = len(model_config.full_attention_layer_ids) - swa_layers_num = len(model_config.swa_attention_layer_ids) - assert swa_layers_num > 0, "Hybrid SWA model must have at least one SWA layer" + def __init__(self, mr: ModelRunner): + model_config = mr.model_config + kv_cache_dtype = mr.kv_cache_dtype + kv_size = torch._utils._element_size(kv_cache_dtype) + tp_size = get_attention_tp_size() - def align_page_size(x: int) -> int: - return (x // page_size) * page_size + self._full_layers_num = len(model_config.full_attention_layer_ids) + self._swa_layers_num = len(model_config.swa_attention_layer_ids) + assert ( + self._swa_layers_num > 0 + ), "Hybrid SWA model must have at least one SWA layer" - if full_layers_num == 0: - # all layers are SWA - swa_tokens = align_page_size(token_capacity) - logger.info( - f"Use sliding window memory pool (all SWA). swa_layer_tokens={swa_tokens}" + self._swa_full_tokens_ratio = mr.server_args.swa_full_tokens_ratio + + # Full layer per-token memory (bytes) + self._full_per_token = ( + model_config.get_num_kv_heads(tp_size) + * (model_config.head_dim + model_config.v_head_dim) + * kv_size ) - return swa_tokens, 0, swa_tokens - # Use unified memory-based allocation for all hybrid SWA models. - # - # Let: - # F = Full layer per-token memory - # S = SWA layer per-token memory (may differ from F) - # r = swa_full_tokens_ratio = swa_tokens / full_tokens - # - # The profile phase computed: - # cell_size = F * n_full + S * n_swa - # token_capacity = rest_memory / cell_size - # => total_memory = token_capacity * (F * n_full + S * n_swa) - # - # We need to solve: - # full_tokens * F * n_full + swa_tokens * S * n_swa = total_memory - # swa_tokens = full_tokens * r - # - # Solution: - # full_tokens = total_memory / (F * n_full + r * S * n_swa) - # = token_capacity * (F * n_full + S * n_swa) / (F * n_full + r * S * n_swa) + # SWA layer per-token memory (bytes) + self._swa_per_token = ( + model_config.get_swa_num_kv_heads(tp_size) + * (model_config.swa_head_dim + model_config.swa_v_head_dim) + * kv_size + ) - kv_size = torch._utils._element_size(mr.kv_cache_dtype) + # Bytes per max_total_num_token. + # For hybrid (full_layers > 0): full_tokens * _cell_size = total memory for both pools. + # For all-SWA (full_layers == 0): swa_tokens * _cell_size = total SWA memory. + if self._full_layers_num == 0: + self._cell_size = self._swa_per_token * self._swa_layers_num + else: + self._cell_size = ( + self._full_per_token * self._full_layers_num + + self._swa_full_tokens_ratio + * self._swa_per_token + * self._swa_layers_num + ) - # Full layer per-token memory - full_per_token = ( - model_config.get_num_kv_heads(get_attention_tp_size()) - * (model_config.head_dim + model_config.v_head_dim) - * kv_size - ) + def _solve_pool_sizes( + self, max_total_num_tokens: int, page_size: int + ) -> MemoryPoolConfig: + """Core computation: split max_total_num_tokens into full/swa pool sizes.""" - # SWA layer per-token memory - swa_per_token = ( - model_config.get_swa_num_kv_heads(get_attention_tp_size()) - * (model_config.swa_head_dim + model_config.swa_v_head_dim) - * kv_size - ) + def align_page_size(x: int) -> int: + return (x // page_size) * page_size - # Total memory available from profile - total_memory = token_capacity * ( - full_per_token * full_layers_num + swa_per_token * swa_layers_num - ) + if self._full_layers_num == 0: + # All layers are SWA — no full pool needed + swa_tokens = align_page_size(max_total_num_tokens) + logger.info( + f"Use sliding window memory pool (all SWA). " + f"swa_layer_tokens={swa_tokens}" + ) + return MemoryPoolConfig( + max_total_num_tokens=swa_tokens, + full_max_total_num_tokens=0, + swa_max_total_num_tokens=swa_tokens, + ) - # Solve the equations - denominator = ( - full_per_token * full_layers_num - + swa_full_tokens_ratio * swa_per_token * swa_layers_num - ) - assert ( - denominator > 0 - ), f"Invalid denominator={denominator} for memory-based allocation. full_per_token={full_per_token}, full_layers_num={full_layers_num}, swa_per_token={swa_per_token}, swa_layers_num={swa_layers_num}, swa_full_tokens_ratio={swa_full_tokens_ratio}" + # full_tokens = max_total_num_tokens (page aligned) + # swa_tokens = full_tokens * ratio (page aligned) + full_tokens = align_page_size(max_total_num_tokens) + swa_tokens = align_page_size(int(full_tokens * self._swa_full_tokens_ratio)) - full_tokens = align_page_size(int(total_memory / denominator)) - swa_tokens = align_page_size(int(full_tokens * swa_full_tokens_ratio)) + logger.info( + f"Use sliding window memory pool. " + f"full_layer_tokens={full_tokens}, swa_layer_tokens={swa_tokens}" + ) - logger.info( - f"Use sliding window memory pool. full_layer_tokens={full_tokens}, swa_layer_tokens={swa_tokens}" - ) - return full_tokens, full_tokens, swa_tokens + return MemoryPoolConfig( + max_total_num_tokens=full_tokens, + full_max_total_num_tokens=full_tokens, + swa_max_total_num_tokens=swa_tokens, + ) + + def calculate_pool_sizes( + self, available_bytes: int, page_size: int + ) -> MemoryPoolConfig: + max_total_num_tokens = int(available_bytes // self._cell_size) + return self._solve_pool_sizes(max_total_num_tokens, page_size) + + def calculate_pool_sizes_from_max_tokens( + self, max_total_num_tokens: int, page_size: int + ) -> MemoryPoolConfig: + return self._solve_pool_sizes(max_total_num_tokens, page_size) + + +def create_memory_pool_configurator( + mr: ModelRunner, +) -> MemoryPoolConfigurator: + """Factory: select the right configurator for the model architecture.""" + if mr.is_hybrid_swa: + return HybridSWAPoolConfigurator(mr) + # Future: MambaPoolConfigurator + return DefaultPoolConfigurator(mr)