Extract per-architecture KV-cache pool builders into KVCacheConfigurator (#31163)
This commit is contained in:
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)
|
||||
|
||||
Reference in New Issue
Block a user