Introduce KVCacheConfigurator and migrate KV-cache config logic (#31162)
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user