From 7bc32041170c6585c1a05be102dc0eeb72ab9d4c Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Fri, 28 Aug 2026 10:21:32 -0700 Subject: [PATCH] config: three cache and pool readers take the bags (#36791) Co-authored-by: Claude Opus 5 --- python/sglang/srt/mem_cache/registry.py | 10 +++++----- python/sglang/srt/mem_cache/unified_radix_cache.py | 6 +++--- .../sglang/srt/model_executor/pool_configurator.py | 4 ++-- .../unit/model_executor/test_pool_configurator.py | 13 +++++++++---- 4 files changed, 19 insertions(+), 14 deletions(-) diff --git a/python/sglang/srt/mem_cache/registry.py b/python/sglang/srt/mem_cache/registry.py index 980136f88..a1371e384 100644 --- a/python/sglang/srt/mem_cache/registry.py +++ b/python/sglang/srt/mem_cache/registry.py @@ -18,7 +18,7 @@ from sglang.srt.environ import envs from sglang.srt.hardware_backend.mlx.runtime import use_mlx from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.cache_init_params import CacheInitParams -from sglang.srt.runtime_context import get_disagg, get_memory +from sglang.srt.runtime_context import get_disagg, get_memory, get_serving if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig @@ -198,7 +198,7 @@ def _create_unified_radix_cache( def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache: """Route to the matching factory to construct Radix Cache.""" - name = ctx.server_args.radix_cache_backend + name = get_memory().radix_cache_backend if name: factory = get_radix_cache_factory(name) if factory is None: @@ -215,7 +215,7 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache: if ( get_memory().enable_hierarchical_cache - and ctx.server_args.hicache_host_memory_mode == "buffer_only" + and get_memory().hicache_host_memory_mode == "buffer_only" ): from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache @@ -225,7 +225,7 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache: f"the unified radix tree; this model selected {type(cache).__name__}." ) - if ctx.server_args.enable_session_radix_cache and not getattr( + if get_memory().enable_session_radix_cache and not getattr( cache, "enable_session_radix_cache", False ): raise ValueError( @@ -237,7 +237,7 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache: hicache_attached = cache.cache_controller is not None streaming_wrapped = False if ( - ctx.server_args.enable_streaming_session + get_serving().enable_streaming_session and not cache.supports_streaming_session() ): from sglang.srt.session.streaming_session import StreamingSession diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index e4d39f76d..ed48bd724 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -82,7 +82,7 @@ from sglang.srt.observability.metrics_collector import ( StorageMetrics, StorageMetricsCollector, ) -from sglang.srt.runtime_context import get_memory +from sglang.srt.runtime_context import get_memory, get_observability from sglang.srt.session.streaming_session import StreamingSession from sglang.srt.utils.common import ceil_align @@ -367,7 +367,7 @@ class UnifiedRadixCache(BasePrefixCache): def init_hicache(self, server_args: ServerArgs, params: CacheInitParams) -> None: """Initialize HiCache infrastructure.""" - self.host_memory_mode = server_args.hicache_host_memory_mode + self.host_memory_mode = get_memory().hicache_host_memory_mode if self.host_memory_mode == "buffer_only": # FULL and FULL+SWA only: Mamba has no state-handoff channel on # the admission-time load-back read path and is not layer-gated. @@ -387,7 +387,7 @@ class UnifiedRadixCache(BasePrefixCache): self.load_cache_event = threading.Event() self.sidecar_pool_specs.clear() - self.extra_metric_labels = server_args.extra_metric_labels + self.extra_metric_labels = get_observability().extra_metric_labels # Parse storage config once, share with assembler and tree storage_backend = get_memory().hicache_storage_backend diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index fb6d4af6c..6edd44817 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -169,7 +169,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): self._zero_kv_max_tokens = ( torch.iinfo(torch.int64).max if has_kv_on_another_pp_stage - else kvc.server_args.max_total_tokens or kvc.model_config.context_len + else get_schedule().max_total_tokens or kvc.model_config.context_len ) # EAGLE/STANDALONE: scale cell_size to account for draft model KV cache. @@ -792,7 +792,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): self.disaggregation_decode_extra_slots = ( get_disagg().disaggregation_decode_extra_slots or 0 ) - if kvc.server_args.enable_hisparse: + if get_memory().enable_hisparse: from sglang.srt.mem_cache.sparsity import parse_hisparse_config self.c4_shrink_factor = parse_hisparse_config( diff --git a/test/registered/unit/model_executor/test_pool_configurator.py b/test/registered/unit/model_executor/test_pool_configurator.py index 2525a676f..3412a18d2 100644 --- a/test/registered/unit/model_executor/test_pool_configurator.py +++ b/test/registered/unit/model_executor/test_pool_configurator.py @@ -12,7 +12,12 @@ from unittest.mock import MagicMock, patch from sglang.srt.configs.model_config import AttentionArch from sglang.srt.distributed.parallel_state_wrapper import ParallelState -from sglang.srt.runtime_context import get_memory, get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_memory, + get_parallel, + get_schedule, + get_server_args, +) from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -311,7 +316,7 @@ class TestHybridSWAConfigurator(CustomTestCase): ) cfg = create_memory_pool_configurator(mr) - config = cfg.calculate_pool_sizes(available_bytes, mr.server_args.page_size) + config = cfg.calculate_pool_sizes(available_bytes, get_schedule().page_size) return mr, cfg, config def test_memory_utilization(self): @@ -413,7 +418,7 @@ class TestHybridSWAConfigurator(CustomTestCase): user_limit = original.full_max_total_num_tokens // 2 with mock_cpu_env(): config = cfg.calculate_pool_sizes_from_max_tokens( - user_limit, mr.server_args.page_size + user_limit, get_schedule().page_size ) used = _actual_memory_used(mr, config) self.assertLessEqual(used, available) @@ -995,7 +1000,7 @@ class TestDflashDraftKvBudget(CustomTestCase): ) cfg = create_memory_pool_configurator(mr) - config = cfg.calculate_pool_sizes(available, mr.server_args.page_size) + config = cfg.calculate_pool_sizes(available, get_schedule().page_size) return config.full_max_total_num_tokens self.assertLess(_tokens(10240), _tokens(None))