diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index de1bd2b84943..b72dc2a29622 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -189,7 +189,7 @@ def __init__( self.scheduler = scheduler self.scheduler_stage_metrics = scheduler_stage_metrics self.max_total_num_tokens = ( - self.scheduler.tp_worker.model_runner.effective_max_total_num_tokens + self.scheduler.tp_worker.model_runner.effective_logical_max_total_num_tokens ) self.transfer_backend = transfer_backend if envs.SGLANG_DISAGG_STAGING_BUFFER.get(): diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index f74ce7307d02..c8952401744b 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1250,7 +1250,8 @@ def emit_metrics_constants(self) -> None: # TODO: max_running_requests_under_SLO has no setter — dead chain. max_running_requests_under_SLO=None, page_size=self.page_size, - num_pages=self.max_total_num_tokens // self.page_size, + num_pages=self.max_total_num_tokens + // self.token_to_kv_pool_allocator.page_size, context_len=self.model_config.context_len, startup_available_gpu_memory_gb=self.startup_available_gpu_memory_gb, ) @@ -2345,8 +2346,7 @@ def init_pool_stats_observer(self) -> None: enable_hisparse=self.enable_hisparse, full_tokens_per_layer=self.full_tokens_per_layer, swa_tokens_per_layer=self.swa_tokens_per_layer, - max_total_num_tokens=self.max_total_num_tokens - * get_parallel().attn_dcp_size, + max_total_num_tokens=self.max_total_num_tokens, get_last_batch=lambda: self.last_batch, get_running_batch=lambda: self.running_batch, ) @@ -2527,7 +2527,7 @@ def init_req_max_new_tokens(self, req): max_new_tokens = self.token_to_kv_pool_allocator.max_new_tokens_for_memory( input_len, max_new_tokens, - token_capacity=self.max_total_num_tokens * get_parallel().attn_dcp_size, + token_capacity=self.max_total_num_tokens, sliding_window_size=self.sliding_window_size, chunk_size=self.chunked_prefill_size, ) diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 477ae90576fd..c8a04268744f 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -409,9 +409,7 @@ def alloc_memory_pool( assert self.model_runner.max_running_requests > 0, "max_running_request is zero" max_req_len = min( self.model_config.context_len - 1, - self.model_runner.effective_max_total_num_tokens - * get_parallel().attn_dcp_size - - 1, + self.model_runner.effective_logical_max_total_num_tokens - 1, ) assert max_req_len > 0, "Memory pool size is too small" @@ -539,17 +537,13 @@ def register_hisparse_coordinator(self, coordinator): def get_worker_info(self): max_req_len = min( self.model_config.context_len - 1, - self.model_runner.effective_max_total_num_tokens - * get_parallel().attn_dcp_size - - 1, + self.model_runner.effective_logical_max_total_num_tokens - 1, ) max_req_input_len = max_req_len - 5 if self.dllm_algorithm is not None: max_req_input_len -= self.dllm_algorithm.block_size return ( - self.model_runner.req_to_token_pool.schedulable_token_capacity( - self.model_runner.max_total_num_tokens - ), + self.model_runner.logical_max_total_num_tokens, get_schedule().max_prefill_tokens, self.model_runner.max_running_requests, get_schedule().max_queued_requests, diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index c0e89ebcc335..13210daa44d3 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -316,6 +316,17 @@ def hybrid_swa_token_capacity( ) return full_capacity or swa_capacity + def logical_token_capacity(self, *, max_total_num_tokens: int) -> int: + """Request tokens a pool of `max_total_num_tokens` per-rank rows holds.""" + # SWA allocators never widen under DCP. + if self.is_hybrid_swa: + return max_total_num_tokens + # Target rows widen into attn_dcp_size ids; draft sizes already carry + # loc_space_scale. + return ( + max_total_num_tokens * get_parallel().attn_dcp_size // self.loc_space_scale + ) + def _build_fp4_quant_method(self, *, num_layers: int): if not is_float4_e2m1fn_x2(self.kv_cache_dtype): return None diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 387a9a3a8234..841fa6762446 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1390,6 +1390,22 @@ def unload_lora_adapter(self, lora_ref: LoRARef): """Unload a lora adapter that was previously loaded during initialization or dynamic loading.""" return self.lora_manager.unload_lora_adapter(lora_ref) + @property + def logical_max_total_num_tokens(self): + """Request-token capacity in logical tokens, not per-rank DCP rows.""" + return self.req_to_token_pool.schedulable_token_capacity( + self.kv_cache_configurator.logical_token_capacity( + max_total_num_tokens=self.max_total_num_tokens + ) + ) + + @property + def effective_logical_max_total_num_tokens(self): + """Logical request limit, preserving hybrid SWA's separate pool bounds.""" + if self.is_hybrid_swa: + return self.effective_max_total_num_tokens + return self.logical_max_total_num_tokens + @property def effective_max_total_num_tokens(self): """Return the max token pool size considering hybrid swa settings.""" diff --git a/test/registered/unit/managers/test_dcp_logical_capacity.py b/test/registered/unit/managers/test_dcp_logical_capacity.py new file mode 100644 index 000000000000..3d7a92322dae --- /dev/null +++ b/test/registered/unit/managers/test_dcp_logical_capacity.py @@ -0,0 +1,139 @@ +"""Logical DCP capacities must agree across validation, admission and telemetry.""" + +import unittest +from types import SimpleNamespace as NS +from unittest.mock import patch + +import torch + +from sglang.srt.managers.tp_worker import TpModelWorker +from sglang.srt.mem_cache.allocator.unified_mamba import ( + UnifiedMambaTokenToKVPoolAllocator, +) +from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator +from sglang.srt.mem_cache.memory_pool import ReqToTokenPool +from sglang.srt.mem_cache.unified_memory_pool import ( + MambaSubPoolSpec, + MLASubPoolSpec, + UnifiedKVPool, +) +from sglang.srt.model_executor.model_runner import ModelRunner +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase, enter_scope, published_topology + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +PHYSICAL = 690240 +CONTEXT = 1048576 + + +def make_configurator(*, is_hybrid_swa=False, is_draft_worker=False): + configurator = KVCacheConfigurator.__new__(KVCacheConfigurator) + configurator.is_hybrid_swa = is_hybrid_swa + configurator.is_draft_worker = is_draft_worker + return configurator + + +def make_unified_mamba_allocator(*, n_full_tokens, n_mamba_slots): + full = MLASubPoolSpec( + name="full", + layer_num=3, + kv_lora_rank=6, + qk_rope_head_dim=2, + store_dtype=torch.bfloat16, + grow_direction="down", + ) + mamba = MambaSubPoolSpec( + name="mamba", + layer_num=2, + conv_state_shapes=((4, 3),), + conv_dtype=torch.float32, + temporal_state_shape=(2, 2, 2), + temporal_dtype=torch.float32, + grow_direction="up", + ) + pool = UnifiedKVPool( + total_bytes=full.entry_bytes() * n_full_tokens + + mamba.entry_bytes() * n_mamba_slots, + sub_pool_specs=[full, mamba], + device="cpu", + enable_memory_saver=False, + page_size=1, + ) + kvcache = NS( + full_kv_pool=NS(buf=torch.empty(pool.max_slots("full"))), + mamba_pool=NS(buf=torch.empty(pool.max_slots("mamba"))), + ) + return UnifiedMambaTokenToKVPoolAllocator( + unified_buffer=pool, kvcache=kvcache, device="cpu", page_size=1 + ) + + +def make_worker(dcp_size): + kv = NS(size=PHYSICAL, mem_usage=8.89) + runner = ModelRunner.__new__(ModelRunner) + runner.server_args = NS(dcp_size=dcp_size) + runner.kv_cache_configurator = make_configurator() + runner.is_hybrid_swa = False + runner.max_total_num_tokens = PHYSICAL + runner.max_running_requests = 64 + runner.token_to_kv_pool = kv + runner.req_to_token_pool = ReqToTokenPool.__new__(ReqToTokenPool) + runner.req_to_token_pool.size = 96 + runner.req_to_token_pool.max_context_len = CONTEXT + runner.req_to_token_pool._aux_cache = None + runner.forward_stream = None + return NS( + model_runner=runner, + model_config=NS(context_len=CONTEXT), + random_seed=0, + device="cpu", + dllm_algorithm=None, + ) + + +class TestDcpLogicalCapacity(CustomTestCase): + def make_worker(self, dcp_size): + enter_scope(self, published_topology(tp_size=dcp_size, dcp_size=dcp_size)) + return make_worker(dcp_size) + + def setUp(self): + for config in ( + patch( + "sglang.srt.managers.tp_worker.get_schedule", + return_value=NS(max_prefill_tokens=16384, max_queued_requests=None), + ), + ): + config.start() + self.addCleanup(config.stop) + + def test_logical_capacity(self): + for dcp_size in (1, 2, 8): + with self.subTest(dcp_size=dcp_size): + info = TpModelWorker.get_worker_info(self.make_worker(dcp_size)) + self.assertEqual(info[0], PHYSICAL * dcp_size) + self.assertEqual(info[4], min(CONTEXT, PHYSICAL * dcp_size) - 1) + + # Unified Mamba's allocator.size also counts Mamba state bytes, so + # capacity must come from the configured rows. + runner = self.make_worker(8).model_runner + runner.max_total_num_tokens = 64 + runner.token_to_kv_pool_allocator = make_unified_mamba_allocator( + n_full_tokens=64, n_mamba_slots=8 + ) + self.assertGreater(runner.token_to_kv_pool_allocator.size, 64 * 8) + self.assertEqual(runner.logical_max_total_num_tokens, 64 * 8) + + # Draft sizes already carry loc_space_scale; SWA never widens. + runner.kv_cache_configurator = make_configurator(is_draft_worker=True) + self.assertEqual(runner.logical_max_total_num_tokens, 64) + runner.is_hybrid_swa = True + runner.kv_cache_configurator = make_configurator(is_hybrid_swa=True) + runner.full_max_total_num_tokens = 32 + runner.swa_max_total_num_tokens = 16 + self.assertEqual(runner.logical_max_total_num_tokens, 64) + self.assertEqual(runner.effective_logical_max_total_num_tokens, 32) + + +if __name__ == "__main__": + unittest.main()