Extract per-architecture KV-cache pool builders into KVCacheConfigurator (#31163)

This commit is contained in:
fzyzcjy
2026-07-14 16:02:09 +08:00
committed by GitHub
parent d6cf2908ce
commit cfd17301a8
5 changed files with 1019 additions and 817 deletions
File diff suppressed because it is too large Load Diff
@@ -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
@@ -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
@@ -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)