Introduce KVCacheConfigurator and migrate KV-cache config logic (#31162)

This commit is contained in:
fzyzcjy
2026-07-14 16:01:45 +08:00
committed by GitHub
parent 725920915f
commit d6cf2908ce
6 changed files with 1730 additions and 1475 deletions
@@ -115,7 +115,7 @@ class FlashinferDispatcher(BaseDispatcher):
# The workspace must fit both:
# (a) the fattest prefill batch (bounded by chunked_prefill_size), and
# (b) the largest decode batch (bounded by max_running_requests, which
# _resolve_max_num_reqs caps at 4096 per DP worker).
# resolve_max_num_reqs caps at 4096 per DP worker).
# max_running_requests is not yet resolved at model-construction time,
# so we use 4096 as a floor to cover decode batches and _dummy_run
# (which warms up at batch_size = req_to_token_pool.size).
File diff suppressed because it is too large Load Diff
@@ -25,6 +25,10 @@ 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,
@@ -92,6 +96,9 @@ from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value
from sglang.srt.mem_cache import kv_cache_dtype
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.kv_cache_configurator import (
KVCacheConfigurator,
)
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
from sglang.srt.model_executor.cuda_graph_config import (
@@ -112,6 +119,10 @@ from sglang.srt.model_executor.forward_context import (
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
from sglang.srt.model_executor.hook_manager import register_forward_hooks
from sglang.srt.model_executor.model_runner_components import misc_utils
from sglang.srt.model_executor.model_runner_components.kv_pool_runtime import (
compute_post_capture_kv_resize,
is_post_capture_kv_active,
)
from sglang.srt.model_executor.model_runner_components.layer_setup import (
ModelLayerInfo,
adjust_hybrid_swa_layer_ids,
@@ -453,6 +464,35 @@ class ModelRunner(ModelRunnerKVCacheMixin):
device=self.device,
)
def init_kv_cache_configurator(self):
self.kv_cache_configurator = KVCacheConfigurator(
device=self.device,
gpu_id=self.gpu_id,
ps=self.ps,
model_config=self.model_config,
server_args=self.server_args,
kv_cache_dtype=self.kv_cache_dtype,
page_size=self.page_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,
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,
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,
memory_pool_config=self.memory_pool_config,
)
def init_mindspore_runner(self):
# Init the mindspore runner
# for now, there is only some communication initialization work
@@ -615,6 +655,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
if memory_pool_config is not None:
self.memory_pool_config = memory_pool_config
self.init_kv_cache_configurator()
self.init_memory_pool(self.pre_model_load_memory)
# Must be called AFTER init_memory_pool so the pool object exists for
@@ -657,6 +698,27 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.graph_shared_output = None
def post_capture_resize_kv_pool(self):
resize = compute_post_capture_kv_resize(self)
self.max_total_num_tokens = resize.max_total_num_tokens
if self.is_hybrid_swa:
self.full_max_total_num_tokens = resize.full_max_total_num_tokens
self.swa_max_total_num_tokens = resize.swa_max_total_num_tokens
if self.memory_pool_config is not None:
self.memory_pool_config.max_total_num_tokens = resize.max_total_num_tokens
self.memory_pool_config.full_max_total_num_tokens = (
resize.full_max_total_num_tokens
)
self.memory_pool_config.swa_max_total_num_tokens = (
resize.swa_max_total_num_tokens
)
if resize.capped_max_running_requests is not None:
self.max_running_requests = resize.capped_max_running_requests
if self.memory_pool_config is not None:
self.memory_pool_config.max_running_requests = (
resize.capped_max_running_requests
)
def init_attention_backends(self):
"""Initialize attention backends only (no cuda graph capture)."""
# TODO: Refactor device-specific init branches into platform interface (separate PR).
@@ -0,0 +1,115 @@
from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Optional
import msgspec
import torch
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.distributed import get_world_group
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.platforms import current_platform
from sglang.srt.utils.common import get_available_gpu_memory, get_device_memory_capacity
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
def is_post_capture_kv_active(
*, server_args: ServerArgs, is_draft_worker: bool
) -> bool:
return (
server_args.post_capture_kv_sizing_planned()
and current_platform.is_cuda()
and not is_draft_worker
)
class PostCaptureKVResize(msgspec.Struct, frozen=True, kw_only=True):
max_total_num_tokens: int
full_max_total_num_tokens: Optional[int]
swa_max_total_num_tokens: Optional[int]
capped_max_running_requests: Optional[int]
def compute_post_capture_kv_resize(
model_runner: ModelRunner,
) -> PostCaptureKVResize:
"""Resize the KV pool after capture and return the new sizes for the
orchestrator to assign. Takes the live ModelRunner because it reads
post-capture GPU memory + the pool objects it must resize in place."""
pool = model_runner.token_to_kv_pool
torch.cuda.synchronize()
free_gb = get_available_gpu_memory(
model_runner.device,
model_runner.gpu_id,
distributed=get_world_group().world_size > 1,
cpu_group=get_world_group().cpu_group,
)
headroom_gb = model_runner.pre_model_load_memory * (
1 - model_runner.mem_fraction_static
)
decode_cuda_graph_config = model_runner.server_args.cuda_graph_config.decode
decode_max_bs = int(decode_cuda_graph_config.max_bs or 0)
running_requests = int(model_runner.max_running_requests or decode_max_bs or 1)
eager_decode_gap = (
model_runner.server_args.disaggregation_mode != "prefill"
and decode_cuda_graph_config.backend != Backend.DISABLED
and decode_max_bs < running_requests
)
if eager_decode_gap:
logger.warning(
"Post-capture KV sizing: decode CUDA graph max_bs=%d < "
"max_running_requests=%d; reserving activation headroom",
decode_max_bs,
running_requests,
)
if eager_decode_gap or mambaish_config(model_runner.model_config) is not None:
headroom_gb = max(
headroom_gb,
model_runner.server_args.mamba_pre_capture_reserve_mb(
get_device_memory_capacity(model_runner.device)
)
/ 1024,
)
budget_bytes = (
int(max(0.0, free_gb - headroom_gb) * (1 << 30))
+ pool.post_capture_backed_bytes
)
config = model_runner.kv_cache_configurator.config_from_budget(
budget_bytes, cap_tokens=model_runner.max_total_num_tokens
)
pool.finalize_backing(config)
model_runner.token_to_kv_pool_allocator.resize(config)
capped_max_running_requests = None
if model_runner.max_running_requests is not None:
# Re-calculate max_running_requests for the now smaller pool
capped_reqs = min(
model_runner.max_running_requests,
model_runner.kv_cache_configurator.resolve_max_num_reqs(
config.max_total_num_tokens
),
)
if capped_reqs < model_runner.max_running_requests:
logger.warning(
"Post-capture KV sizing: max_running_requests %d -> %d",
model_runner.max_running_requests,
capped_reqs,
)
capped_max_running_requests = capped_reqs
logger.info(
"Post-capture KV sizing: max_total_num_tokens=%d, free memory=%.2f GB",
config.max_total_num_tokens,
get_available_gpu_memory(model_runner.device, model_runner.gpu_id),
)
return PostCaptureKVResize(
max_total_num_tokens=config.max_total_num_tokens,
full_max_total_num_tokens=config.full_max_total_num_tokens,
swa_max_total_num_tokens=config.swa_max_total_num_tokens,
capped_max_running_requests=capped_max_running_requests,
)
File diff suppressed because it is too large Load Diff
@@ -557,7 +557,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
self.disaggregation_decode_extra_slots = (
mr.server_args.disaggregation_decode_extra_slots or 0
)
if mr.enable_hisparse:
if mr.server_args.enable_hisparse:
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
self.c4_shrink_factor = parse_hisparse_config(