[core] Introduce MemoryPoolConfigurator class hierarchy (#22389)

This commit is contained in:
Liangsheng Yin
2026-04-09 15:29:19 +08:00
committed by GitHub
parent b9c316917b
commit de441ac6bb
4 changed files with 270 additions and 225 deletions
+1 -1
View File
@@ -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__)
@@ -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,
@@ -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,24 +677,29 @@ 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)
full_tokens = None
swa_tokens = None
if self.is_hybrid_swa:
token_capacity, full_tokens, swa_tokens = self._resolve_hybrid_swa_tokens(
token_capacity
from sglang.srt.model_executor.pool_configurator import (
create_memory_pool_configurator,
)
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,
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
)
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:
assert (
@@ -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,20 +24,102 @@ 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:
class MemoryPoolConfigurator:
"""Base class for memory pool configurators.
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
use_mla_backend = mr.use_mla_backend
kv_size = torch._utils._element_size(kv_cache_dtype)
if use_mla_backend:
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
@@ -45,31 +141,16 @@ def get_cell_size_per_token(mr: ModelRunner, num_layers: int) -> int:
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
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
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
@@ -78,95 +159,130 @@ def get_cell_size_per_token(mr: ModelRunner, num_layers: int) -> int:
if is_float4_e2m1fn_x2(kv_cache_dtype):
# kv_scale_buffer
scale_block_size = 16
n = model_config.get_num_kv_heads(get_attention_tp_size())
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
)
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 resolve_hybrid_swa_tokens(
mr: ModelRunner, token_capacity: int
) -> tuple[int, int, int]:
"""Split token_capacity into full/swa pools.
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)
Returns (effective_capacity, full_max_total_num_tokens, swa_max_total_num_tokens).
class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
"""Configurator for hybrid sliding window attention models (Gemma2, Command-R, MiMo).
Splits available memory between full attention and SWA pools.
Does NOT inherit DefaultPoolConfigurator — different coeff model.
"""
def __init__(self, mr: ModelRunner):
model_config = mr.model_config
page_size = mr.server_args.page_size
swa_full_tokens_ratio = mr.server_args.swa_full_tokens_ratio
kv_cache_dtype = mr.kv_cache_dtype
kv_size = torch._utils._element_size(kv_cache_dtype)
tp_size = get_attention_tp_size()
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"
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"
def align_page_size(x: int) -> int:
return (x // page_size) * page_size
self._swa_full_tokens_ratio = mr.server_args.swa_full_tokens_ratio
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}"
)
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)
kv_size = torch._utils._element_size(mr.kv_cache_dtype)
# Full layer per-token memory
full_per_token = (
model_config.get_num_kv_heads(get_attention_tp_size())
# 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
)
# SWA layer per-token memory
swa_per_token = (
model_config.get_swa_num_kv_heads(get_attention_tp_size())
# 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
)
# Total memory available from profile
total_memory = token_capacity * (
full_per_token * full_layers_num + swa_per_token * swa_layers_num
# 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
)
# 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}"
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."""
full_tokens = align_page_size(int(total_memory / denominator))
swa_tokens = align_page_size(int(full_tokens * swa_full_tokens_ratio))
def align_page_size(x: int) -> int:
return (x // page_size) * page_size
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,
)
# 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))
logger.info(
f"Use sliding window memory pool. full_layer_tokens={full_tokens}, swa_layer_tokens={swa_tokens}"
f"Use sliding window memory pool. "
f"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)