From a97bc8db321c86baffbdc39b15c8158f8c7e64d3 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Mon, 17 Aug 2026 16:17:26 -0700 Subject: [PATCH] config: the DP/EP topology reads come from the parallel bag (#35025) --- .../srt/disaggregation/encode_server.py | 8 +- python/sglang/srt/entrypoints/engine.py | 14 +- python/sglang/srt/entrypoints/http_server.py | 16 ++- python/sglang/srt/layers/dp_attention.py | 15 ++- python/sglang/srt/lora/lora_manager.py | 10 +- .../srt/managers/data_parallel_controller.py | 31 +++-- python/sglang/srt/managers/load_snapshot.py | 45 ++++--- .../srt/managers/multi_tokenizer_mixin.py | 4 +- python/sglang/srt/managers/prefill_delayer.py | 3 +- python/sglang/srt/managers/scheduler.py | 35 +++-- .../scheduler_components/recv_skipper.py | 14 +- .../srt/managers/tokenizer_control_mixin.py | 29 ++-- .../sglang/srt/managers/tokenizer_manager.py | 6 +- .../sglang/srt/mem_cache/kv_cache_builder.py | 14 +- .../startup_weight_load.py | 5 +- .../srt/ray/data_parallel_controller.py | 19 +-- python/sglang/srt/ray/engine.py | 21 +-- .../dspark_components/dspark_worker_v2.py | 10 +- .../sglang/srt/speculative/eagle_worker_v2.py | 7 +- .../speculative/frozen_kv_mtp_worker_v2.py | 9 +- .../multi_layer_eagle_worker_v2.py | 4 +- .../srt/speculative/standalone_worker_v2.py | 4 +- .../srt/utils/cuda_vmm_transport_utils.py | 11 +- python/sglang/srt/utils/offloader.py | 24 ++-- .../scheduler/test_prefill_delayer.py | 6 + .../entrypoints/test_http_server_warmup.py | 6 + .../managers/test_load_snapshot_backends.py | 126 ++++++++++-------- .../managers/test_scheduler_recv_skipper.py | 36 ++--- .../test_startup_weight_load.py | 12 +- .../multimodal/test_gpu_feature_transport.py | 8 +- .../unit/test_global_config_read_ratchet.py | 15 +++ .../unit/test_publish_precedes_bag_reads.py | 4 - ...test_supplied_instance_exposure_ratchet.py | 65 --------- 33 files changed, 332 insertions(+), 304 deletions(-) diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 5e1ff18e2..7f4d9c2e6 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -3957,7 +3957,7 @@ def launch_server(server_args: ServerArgs): # Publish before the launch path reads configuration; the encoder built # below re-projects the same object. publish(server_args, role="encoder") - if server_args.dp_size > 1: + if get_parallel().dp_size > 1: _launch_server_dp(server_args) return @@ -4018,12 +4018,12 @@ def launch_server(server_args: ServerArgs): def _launch_server_dp(server_args: ServerArgs): global dp_dispatcher - if server_args.dp_size <= 1 or server_args.tp_size != 1: + if get_parallel().dp_size <= 1 or server_args.tp_size != 1: raise ValueError( "Encoder DP mode requires --dp-size > 1 and --tp-size 1; got " - f"dp_size={server_args.dp_size}, tp_size={server_args.tp_size}." + f"dp_size={get_parallel().dp_size}, tp_size={server_args.tp_size}." ) - dp_size = server_args.dp_size + dp_size = get_parallel().dp_size logger.info(f"Launching encoder in DP mode: dp_size={dp_size}") # DP mode: workers (subprocesses) write metrics to the shared multiproc dir; diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 3e98f7208..b53b06208 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -99,7 +99,7 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa from sglang.srt.parser.template_detection import resolve_auto_parsers from sglang.srt.parser.template_manager import TemplateManager from sglang.srt.plugins import load_plugins -from sglang.srt.runtime_context import publish +from sglang.srt.runtime_context import get_parallel, publish from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import ( MultiprocessingSerializer, @@ -336,7 +336,7 @@ class Engine(EngineScoreMixin, EngineBase): routed_dp_rank = data_parallel_rank if routed_dp_rank is not None: - dp_size = self.server_args.dp_size + dp_size = get_parallel().dp_size if dp_size <= 1 and routed_dp_rank == 0: logger.debug( f"routed_dp_rank={routed_dp_rank} is ignored because dp_size={dp_size}" @@ -661,7 +661,7 @@ class Engine(EngineScoreMixin, EngineBase): (``python -m sglang.srt.weight_cache.daemon``) plus ``--weight-cache-mode client``, where the daemon outlives the engine. """ - if server_args.dp_size > 1: + if get_parallel().dp_size > 1: raise ValueError( "Weight cache daemon mode does not support dp_size > 1. " "Please set --dp-size 1 when using --weight-cache-mode daemon." @@ -751,7 +751,7 @@ class Engine(EngineScoreMixin, EngineBase): "--dp-size", "1", "--ep-size", - str(server_args.ep_size), + str(get_parallel().ep_size), "--load-format", server_args.load_format, "--dtype", @@ -863,7 +863,7 @@ class Engine(EngineScoreMixin, EngineBase): """ scheduler_procs = [] use_dp_controller = ( - server_args.dp_size > 1 or server_args.ep_join_mode == "scale" + get_parallel().dp_size > 1 or server_args.ep_join_mode == "scale" ) if not use_dp_controller: @@ -1835,7 +1835,7 @@ def _compute_parallelism_ranks( server_args: ServerArgs, tp_rank: int ) -> Tuple[int, int, int]: """Compute attention-CP, MoE-DP, and MoE-EP ranks for a TP rank.""" - attn_dp_size = server_args.dp_size if server_args.enable_dp_attention else 1 + attn_dp_size = get_parallel().dp_size if get_parallel().enable_dp_attention else 1 # Parallelism hierarchy (outermost to innermost): # - Attention: Global(TP) -> DP -> ATTN_CP -> ATTN_TP (innermost) @@ -1846,6 +1846,6 @@ def _compute_parallelism_ranks( moe_ep_rank = ( tp_rank % (server_args.tp_size // server_args.moe_dp_size) - // (server_args.tp_size // server_args.moe_dp_size // server_args.ep_size) + // (server_args.tp_size // server_args.moe_dp_size // get_parallel().ep_size) ) return attn_cp_rank, moe_dp_rank, moe_ep_rank diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index fc9f5f0cf..6002989ee 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -480,6 +480,7 @@ v1_loads_router.route_class = ORJSONRoute app.include_router(v1_loads_router) from sglang.srt.entrypoints.elastic_ep import router as elastic_ep_router +from sglang.srt.runtime_context import get_parallel elastic_ep_router.route_class = ORJSONRoute app.include_router(elastic_ep_router) @@ -2156,7 +2157,10 @@ async def _send_disaggregation_warmup_requests( headers=headers, ) as session: return await asyncio.gather( - *(send_request(session, dp_rank) for dp_rank in range(server_args.dp_size)) + *( + send_request(session, dp_rank) + for dp_rank in range(get_parallel().dp_size) + ) ) @@ -2214,9 +2218,9 @@ def _execute_server_warmup(server_args: ServerArgs): }, } if server_args.skip_tokenizer_init: - json_data["input_ids"] = [[10, 11, 12] for _ in range(server_args.dp_size)] + json_data["input_ids"] = [[10, 11, 12] for _ in range(get_parallel().dp_size)] # TODO Workaround the bug that embedding errors for list of size 1 - if server_args.dp_size == 1: + if get_parallel().dp_size == 1: json_data["input_ids"] = json_data["input_ids"][0] elif ( is_vlm @@ -2260,9 +2264,9 @@ def _execute_server_warmup(server_args: ServerArgs): "temperature": 0.0, } else: - json_data["text"] = ["The capital city of France is"] * server_args.dp_size + json_data["text"] = ["The capital city of France is"] * get_parallel().dp_size # TODO Workaround the bug that embedding errors for list of size 1 - if server_args.dp_size == 1: + if get_parallel().dp_size == 1: json_data["text"] = json_data["text"][0] # Config debug dumping @@ -2304,7 +2308,7 @@ def _execute_server_warmup(server_args: ServerArgs): if not failed_status_codes: logger.info( "Disaggregation warmup requests completed for all %s DP ranks", - server_args.dp_size, + get_parallel().dp_size, ) logger.info("End of disaggregation warmup") else: diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index 355ed42e5..8205287b0 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -31,7 +31,10 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.runtime_context import ( configured_attn_cp_size, configured_moe_dp_size, + get_device, + get_exec, get_flags, + get_parallel, ) from sglang.srt.utils import get_bool_env_var, is_hip @@ -345,9 +348,9 @@ def initialize_dp_attention( dp.max_len_with_idle = ( getattr(model_config.hf_config, "hybrid_override_pattern", None) is not None ) - enable_dp_attention = server_args.enable_dp_attention - dp_size = server_args.dp_size - attn_cp_size = server_args.attn_cp_size + enable_dp_attention = get_parallel().enable_dp_attention + dp_size = get_parallel().dp_size + attn_cp_size = configured_attn_cp_size() dp.enabled = enable_dp_attention @@ -359,15 +362,15 @@ def initialize_dp_attention( ) _ATTN_DP_SIZE = dp_size if enable_dp_attention else 1 - if server_args.elastic_ep_backend is not None and server_args.max_ep_size: - _ATTN_DP_RANK = tp_rank + server_args.ep_join_rank_offset + if get_exec().moe.elastic_ep_backend is not None and get_parallel().max_ep_size: + _ATTN_DP_RANK = tp_rank + get_parallel().ep_join_rank_offset if server_args.is_ep_scale_joiner: dp.joiner_skip_all_gather = True _DpGatheredBufferWrapper.set_metadata( hidden_size=model_config.hidden_size, dtype=model_config.dtype, - device=torch.device(server_args.device), + device=torch.device(get_device().device), ) diff --git a/python/sglang/srt/lora/lora_manager.py b/python/sglang/srt/lora/lora_manager.py index dc9316806..aae2afffc 100644 --- a/python/sglang/srt/lora/lora_manager.py +++ b/python/sglang/srt/lora/lora_manager.py @@ -46,7 +46,11 @@ from sglang.srt.lora.utils import ( ) from sglang.srt.managers.io_struct import LoRAUpdateOutput from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.runtime_context import get_exec, get_parallel +from sglang.srt.runtime_context import ( + get_exec, + get_lora, + get_parallel, +) from sglang.srt.server_args import ServerArgs from sglang.srt.utils import get_available_gpu_memory, replace_submodule from sglang.srt.utils.hf_transformers_utils import AutoConfig @@ -93,7 +97,7 @@ class LoRAManager: self.pending_lora_load_events = {} self.eviction_policy = server_args.lora_eviction_policy - self.enable_dp_attention: bool = server_args.enable_dp_attention + self.enable_dp_attention: bool = get_parallel().enable_dp_attention self._experts_shared_outer_override: Optional[bool] = ( server_args.experts_shared_outer_loras ) @@ -1031,7 +1035,7 @@ def init_lora_cuda_graph_moe_buffers( from sglang.srt.lora.layers import FusedMoEWithLoRA max_bs = get_exec().graph.cuda_graph_config.decode.max_bs - max_loras = server_args.max_loras_per_batch + max_loras = get_lora().max_loras_per_batch for module in model.modules(): if isinstance(module, FusedMoEWithLoRA): lora_manager.init_cuda_graph_moe_buffers(max_bs, max_loras, dtype, module) diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index dddd0e16c..6a4766fa3 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -49,7 +49,7 @@ from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats from sglang.srt.observability.startup_time import aggregate_scheduler_startup_times from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info -from sglang.srt.runtime_context import get_exec, publish +from sglang.srt.runtime_context import get_exec, get_parallel, publish from sglang.srt.server_args import ( DP_ATTENTION_HANDSHAKE_PORT_DELTA, PortArgs, @@ -147,7 +147,7 @@ class DataParallelController: self.run_scheduler_process_func = run_scheduler_process_func # Init inter-process communication - self.context = zmq.Context(1 + server_args.dp_size) + self.context = zmq.Context(1 + get_parallel().dp_size) if server_args.node_rank == 0: self.recv_from_tokenizer = get_zmq_socket( self.context, zmq.PULL, port_args.scheduler_input_ipc_name, False @@ -167,8 +167,8 @@ class DataParallelController: LoadBalanceMethod.TOTAL_TOKENS, ) - self.launch_dp_size: int = server_args.dp_size - self.max_dp_size: int = server_args.max_ep_size or server_args.dp_size + self.launch_dp_size: int = get_parallel().dp_size + self.max_dp_size: int = server_args.max_ep_size or get_parallel().dp_size assert self.max_dp_size >= self.launch_dp_size, ( f"--max-ep-size ({self.max_dp_size}) must be >= " f"--dp ({self.launch_dp_size})." @@ -178,9 +178,8 @@ class DataParallelController: self.max_dp_size - self.launch_dp_size ) - self.dp_budget = DPBudget(server_args.dp_size) + self.dp_budget = DPBudget(get_parallel().dp_size) self.load_snapshot_reader = create_load_snapshot_reader( - server_args, port_args, caller="DataParallelController", ) @@ -196,14 +195,14 @@ class DataParallelController: self._active_workers: List[int] = list(range(self.launch_dp_size)) self._active_count_cache: int = self.launch_dp_size - if server_args.enable_dp_attention: + if get_parallel().enable_dp_attention: self.launch_dp_attention_schedulers(server_args, port_args) # When local control broadcast is enabled, send control messages to # every DP group leader (attn_tp_rank=0) so each leader broadcasts # within its own attn_tp_group instead of the full tp_group. # Otherwise fall back to the original behaviour: send to only the # first leader, which then broadcasts over the full tp_group. - local_ctrl = server_args.enable_dp_attention_local_control_broadcast + local_ctrl = get_parallel().enable_dp_attention_local_control_broadcast self.control_message_step = 1 if local_ctrl else server_args.tp_size else: self.launch_dp_schedulers(server_args, port_args) @@ -367,7 +366,7 @@ class DataParallelController: threads = [] sockets = [] ready_events = [] - for dp_rank in range(server_args.dp_size): + for dp_rank in range(get_parallel().dp_size): tmp_port_args = PortArgs.init_new(server_args) tmp_port_args.tokenizer_ipc_name = port_args.tokenizer_ipc_name tmp_port_args.detokenizer_ipc_name = port_args.detokenizer_ipc_name @@ -569,7 +568,7 @@ class DataParallelController: bind_count = ( self.max_dp_size if server_args.elastic_ep_backend is not None - else server_args.dp_size + else get_parallel().dp_size ) for slot in range(bind_count): worker_port, worker_socket = get_zmq_socket_on_host( @@ -599,7 +598,7 @@ class DataParallelController: dp_rank: Optional[int], worker_ports: Optional[List[int]] = None, ): - if not server_args.enable_dp_attention: + if not get_parallel().enable_dp_attention: logger.info(f"Launch DP{dp_rank} starting at GPU #{base_gpu_id}.") memory_saver_adapter = TorchMemorySaverAdapter.create( @@ -633,13 +632,13 @@ class DataParallelController: for tp_rank in tp_rank_range: rank_port_args = port_args - if server_args.enable_dp_attention: + if get_parallel().enable_dp_attention: # dp attention has different sharding logic _, _, dp_rank, _ = compute_dp_attention_world_info( - server_args.enable_dp_attention, + get_parallel().enable_dp_attention, tp_rank, server_args.tp_size, - server_args.dp_size, + get_parallel().dp_size, server_args.attn_cp_size, ) # compute zmq ports for this dp rank @@ -669,7 +668,7 @@ class DataParallelController: + (tp_rank % tp_size_per_node) * server_args.gpu_id_step ) attn_dp_size = ( - server_args.dp_size if server_args.enable_dp_attention else 1 + get_parallel().dp_size if get_parallel().enable_dp_attention else 1 ) # Parallelism hierarchy (outermost to innermost): @@ -688,7 +687,7 @@ class DataParallelController: // ( server_args.tp_size // server_args.moe_dp_size - // server_args.ep_size + // get_parallel().ep_size ) ) diff --git a/python/sglang/srt/managers/load_snapshot.py b/python/sglang/srt/managers/load_snapshot.py index 6b147c98c..26e306e4d 100644 --- a/python/sglang/srt/managers/load_snapshot.py +++ b/python/sglang/srt/managers/load_snapshot.py @@ -52,6 +52,7 @@ import msgspec.msgpack import msgspec.structs from sglang.srt.environ import envs +from sglang.srt.runtime_context import get_parallel, get_serving from sglang.srt.utils.network import is_zmq_endpoint_ipv6 logger = logging.getLogger(__name__) @@ -61,7 +62,7 @@ logger = logging.getLogger(__name__) # --------------------------------------------------------------------------- -def should_use_zmq(server_args) -> bool: +def should_use_zmq() -> bool: """Whether to use zmq PUSH/PULL instead of shared memory for load snapshots. Shared memory (mmap) only works within a single node. When schedulers @@ -70,14 +71,14 @@ def should_use_zmq(server_args) -> bool: ``SGLANG_LOAD_SNAPSHOT_USE_ZMQ`` forces zmq mode for testing. """ return ( - server_args.enable_dp_attention and server_args.nnodes > 1 + get_parallel().enable_dp_attention and get_parallel().nnodes > 1 ) or envs.SGLANG_LOAD_SNAPSHOT_USE_ZMQ.get() _LOAD_AWARE_METHODS = frozenset({"total_requests", "total_tokens"}) -def _tokenizer_load_snapshot_owner_caller(server_args) -> str: +def _tokenizer_load_snapshot_owner_caller() -> str: """The caller that plays the tokenizer-side zmq owner role. In multi-tokenizer mode (``tokenizer_worker_num > 1``) there are N @@ -85,12 +86,12 @@ def _tokenizer_load_snapshot_owner_caller(server_args) -> str: same zmq PULL endpoint. Instead, the single ``MultiTokenizerRouter`` process owns the socket (polls zmq -> SHM) and every worker reads SHM. """ - if server_args.tokenizer_worker_num > 1: + if get_serving().tokenizer_worker_num > 1: return "MultiTokenizerRouter" return "TokenizerManager" -def zmq_reader_owner(server_args, caller: str) -> bool: +def zmq_reader_owner(caller: str) -> bool: """Decide which process owns the zmq PULL socket. Exactly one of ``"DataParallelController"``, ``"TokenizerManager"``, or @@ -108,18 +109,25 @@ def zmq_reader_owner(server_args, caller: str) -> bool: load data -> tokenizer-side owner owns it (polls on /v1/loads calls). The tokenizer-side owner is the ``"MultiTokenizerRouter"`` caller in - multi-tokenizer mode, otherwise the ``"TokenizerManager"`` caller. + multi-tokenizer mode, otherwise the ``"TokenizerManager"`` caller. Which of + the two it is only matters to a tokenizer-side caller, and asking costs the + DP controller a `serving` read its role is not audited for -- so the + controller answers from the parallel leaves alone. """ - if not should_use_zmq(server_args): + if not should_use_zmq(): return False - if server_args.node_rank != 0: + if get_parallel().node_rank != 0: return False - tokenizer_owner = _tokenizer_load_snapshot_owner_caller(server_args) - if server_args.dp_size == 1: - return caller == tokenizer_owner - if server_args.load_balance_method.lower() in _LOAD_AWARE_METHODS: - return caller == "DataParallelController" - return caller == tokenizer_owner + if caller == "DataParallelController": + return ( + get_parallel().dp_size > 1 + and get_parallel().load_balance_method.lower() in _LOAD_AWARE_METHODS + ) + if get_parallel().dp_size > 1 and ( + get_parallel().load_balance_method.lower() in _LOAD_AWARE_METHODS + ): + return False + return caller == _tokenizer_load_snapshot_owner_caller() # --------------------------------------------------------------------------- @@ -627,14 +635,13 @@ def _zmq_addr_for(port_args) -> str: def create_load_snapshot_writer( - server_args, port_args, dp_size: int, dp_rank: int, publish_interval: int = 1, ): """Return a SHM or ZMQ writer based on server configuration.""" - if should_use_zmq(server_args): + if should_use_zmq(): return ZmqLoadSnapshotWriter( _zmq_addr_for(port_args), dp_size, dp_rank, publish_interval ) @@ -643,7 +650,7 @@ def create_load_snapshot_writer( ) -def create_load_snapshot_reader(server_args, port_args, caller: str): +def create_load_snapshot_reader(port_args, caller: str): """Create a load snapshot reader. Args: @@ -651,8 +658,8 @@ def create_load_snapshot_reader(server_args, port_args, caller: str): ``"MultiTokenizerRouter"`` -- determines who binds the zmq PULL socket when zmq mode is active. """ - dp_size = server_args.dp_size - if zmq_reader_owner(server_args, caller): + dp_size = get_parallel().dp_size + if zmq_reader_owner(caller): return ZmqShmLoadSnapshotReader( _zmq_addr_for(port_args), shm_path_for(port_args.instance_id), dp_size ) diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 0a2cd2b57..0281ca332 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -467,9 +467,9 @@ class MultiTokenizerRouter: # read SHM only. Drain it event-driven via the socket's fd instead of # polling on a timer. self.load_snapshot_reader = None - if zmq_reader_owner(server_args, "MultiTokenizerRouter"): + if zmq_reader_owner("MultiTokenizerRouter"): self.load_snapshot_reader = create_load_snapshot_reader( - server_args, port_args, caller="MultiTokenizerRouter" + port_args, caller="MultiTokenizerRouter" ) self._loop.call_soon_threadsafe(self._register_load_snapshot_reader) diff --git a/python/sglang/srt/managers/prefill_delayer.py b/python/sglang/srt/managers/prefill_delayer.py index c00726d5b..5c7e968e0 100644 --- a/python/sglang/srt/managers/prefill_delayer.py +++ b/python/sglang/srt/managers/prefill_delayer.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, NamedTuple, Optional import torch from sglang.srt.environ import envs +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import get_bool_env_var if TYPE_CHECKING: @@ -103,7 +104,7 @@ class PrefillDelayer: f"queue_trigger_enabled={self._queue_trigger_enabled}" ) self.dp_size = dp_size - self.enable_dp_attention = server_args.enable_dp_attention + self.enable_dp_attention = get_parallel().enable_dp_attention dp_size_dim = dp_size if self.enable_dp_attention else 1 # Mirror scheduler_dp_attn_mixin's NCCL all-gather path: when the diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 3f1e7e661..e94f7040e 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -29,7 +29,10 @@ from typing import TYPE_CHECKING, Any, Deque, Dict, List, Optional, Set, Tuple, from sglang.srt.runtime_context import ( attention_backends, + configured_attn_cp_size, + configured_moe_dp_size, configured_pp_size, + configured_tp_size, get_device, get_disagg, get_exec, @@ -285,12 +288,7 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.platforms import current_platform from sglang.srt.plugins import load_plugins -from sglang.srt.runtime_context import ( - get_context, - get_device, - get_parallel, - publish, -) +from sglang.srt.runtime_context import get_context, publish from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.server_args import PortArgs, ServerArgs @@ -451,16 +449,16 @@ class Scheduler( self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get() self.max_new_tokens_limit = envs.SGLANG_MAX_NEW_TOKENS_LIMIT.get() self.enable_hisparse = server_args.enable_hisparse - self.enable_dp_attention = server_args.enable_dp_attention + self.enable_dp_attention = get_parallel().enable_dp_attention self.enable_unified_memory = server_args.enable_unified_memory # Distributed rank info attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = ( compute_dp_attention_world_info( - server_args.enable_dp_attention, + get_parallel().enable_dp_attention, tp_rank, server_args.tp_size, - server_args.dp_size, + get_parallel().dp_size, server_args.attn_cp_size, ) ) @@ -470,7 +468,7 @@ class Scheduler( pp_rank=pp_rank, pp_size=server_args.pp_size, dp_rank=dp_rank, - dp_size=server_args.dp_size, + dp_size=get_parallel().dp_size, attn_tp_rank=attn_tp_rank, attn_tp_size=attn_tp_size, attn_cp_rank=attn_cp_rank, @@ -480,7 +478,7 @@ class Scheduler( attn_dp_rank=attn_dp_rank, attn_dp_size=attn_dp_size, moe_ep_rank=moe_ep_rank, - moe_ep_size=server_args.ep_size, + moe_ep_size=get_parallel().ep_size, moe_dp_rank=moe_dp_rank, moe_dp_size=server_args.moe_dp_size, gpu_id=gpu_id, @@ -768,7 +766,6 @@ class Scheduler( dp_rank = self.ps.dp_rank if self.ps.dp_rank is not None else 0 try: self.load_snapshot_writer = create_load_snapshot_writer( - self.server_args, port_args, self.ps.dp_size, dp_rank, @@ -1277,7 +1274,7 @@ class Scheduler( ) # Init recv skipper and input blocker - self.recv_skipper = SchedulerRecvSkipper.maybe_create(self.server_args) + self.recv_skipper = SchedulerRecvSkipper.maybe_create() self.input_blocker = ( SchedulerInputBlocker(noop=self.ps.attn_tp_rank != 0) if get_bool_env_var("SGLANG_ENABLE_COLOCATED_BATCH_GEN") @@ -4967,15 +4964,15 @@ def configure_scheduler_process( prefix = "" if shown_dp is not None: prefix += f" DP{shown_dp}" - if server_args.pp_size > 1: + if configured_pp_size() > 1: prefix += f" PP{pp_rank}" - if server_args.attn_cp_size > 1: + if configured_attn_cp_size() > 1: prefix += f" ATTN_CP{attn_cp_rank}" - if server_args.moe_dp_size > 1: + if configured_moe_dp_size() > 1: prefix += f" MOE_DP{moe_dp_rank}" - if server_args.tp_size > 1: + if configured_tp_size() > 1: prefix += f" TP{shown_tp}" - if server_args.ep_size > 1: + if get_parallel().ep_size > 1: prefix += f" EP{shown_moe_ep}" # Config the process @@ -4989,7 +4986,7 @@ def configure_scheduler_process( # Set cpu affinity to this gpu process if envs.SGLANG_SET_CPU_AFFINITY.get(): set_gpu_proc_affinity( - server_args.pp_size, server_args.tp_size, server_args.nnodes, gpu_id + configured_pp_size(), configured_tp_size(), get_parallel().nnodes, gpu_id ) if not envs.SGLANG_NUMA_BIND_V2.get(): numa_node = get_numa_node_if_available(server_args, gpu_id) diff --git a/python/sglang/srt/managers/scheduler_components/recv_skipper.py b/python/sglang/srt/managers/scheduler_components/recv_skipper.py index 946e50247..364723b6c 100644 --- a/python/sglang/srt/managers/scheduler_components/recv_skipper.py +++ b/python/sglang/srt/managers/scheduler_components/recv_skipper.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, List, Optional from sglang.srt.environ import envs from sglang.srt.model_executor.forward_batch_info import ForwardMode -from sglang.srt.server_args import ServerArgs +from sglang.srt.runtime_context import get_parallel, get_schedule if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import ScheduleBatch @@ -12,10 +12,10 @@ if TYPE_CHECKING: class SchedulerRecvSkipper: @staticmethod - def maybe_create(server_args: ServerArgs): - if server_args.scheduler_recv_interval <= 1: + def maybe_create(): + if get_schedule().scheduler_recv_interval <= 1: return None - return SchedulerRecvSkipper(server_args) + return SchedulerRecvSkipper() @staticmethod def derive_forward_mode(gathered_modes: List[int]) -> Optional[ForwardMode]: @@ -33,10 +33,10 @@ class SchedulerRecvSkipper: return ForwardMode.TARGET_VERIFY return ForwardMode.DECODE - def __init__(self, server_args: ServerArgs): - self._use_synced_mode = server_args.enable_dp_attention + def __init__(self): + self._use_synced_mode = get_parallel().enable_dp_attention self._counter = 0 - self._threshold = server_args.scheduler_recv_interval + self._threshold = get_schedule().scheduler_recv_interval # All can be tuned if needed self._default_weight = envs.SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_DEFAULT.get() self._weight_of_forward_mode = { diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index c724a53d3..73465a191 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -74,6 +74,7 @@ from sglang.srt.managers.io_struct import ( UpdateWeightsFromTensorReqOutput, ) from sglang.srt.managers.load_snapshot import LoadSnapshot +from sglang.srt.runtime_context import get_parallel from sglang.srt.server_args import LoRARef, ServerArgs from sglang.srt.utils import ( get_bool_env_var, @@ -157,7 +158,7 @@ class TokenizerControlMixin: mode = spec[2] if len(spec) > 2 else "queueing" comm = FanOutCommunicator( self._dispatch_to_scheduler, - server_args.dp_size, + get_parallel().dp_size, mode, ) setattr(self, f"{name}_communicator", comm) @@ -166,8 +167,8 @@ class TokenizerControlMixin: def update_control_communicator_fan_out(self: TokenizerManager, worker_count: int): primary_group_control = ( - self.server_args.enable_dp_attention - and not self.server_args.enable_dp_attention_local_control_broadcast + get_parallel().enable_dp_attention + and not get_parallel().enable_dp_attention_local_control_broadcast ) if primary_group_control: control_fan_out = ( @@ -420,7 +421,7 @@ class TokenizerControlMixin: ) -> Tuple[bool, str]: self.auto_create_handle_loop() assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for update weights from distributed" results = await self.init_weights_update_group_communicator(obj) @@ -433,7 +434,7 @@ class TokenizerControlMixin: ) -> Tuple[bool, str]: self.auto_create_handle_loop() assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for destroy parameter update group" results = await self.destroy_weights_update_group_communicator(obj) @@ -446,7 +447,7 @@ class TokenizerControlMixin: ) -> Tuple[bool, str]: self.auto_create_handle_loop() assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for update weights from distributed" if obj.abort_all_requests: @@ -479,7 +480,7 @@ class TokenizerControlMixin: self.auto_create_handle_loop() # TODO: support DP assert ( - self.server_args.dp_size == 1 + get_parallel().dp_size == 1 ), "dp_size must be 1 for init_weights_send_group_for_remote_instance" result = ( await self.init_weights_send_group_for_remote_instance_communicator(obj) @@ -494,7 +495,7 @@ class TokenizerControlMixin: self.auto_create_handle_loop() # TODO: support DP assert ( - self.server_args.dp_size == 1 + get_parallel().dp_size == 1 ), "dp_size must be 1 for send_weights_to_remote_instance" result = (await self.send_weights_to_remote_instance_communicator(obj))[0] return result.success, result.message @@ -506,7 +507,7 @@ class TokenizerControlMixin: ) -> Tuple[bool, str]: self.auto_create_handle_loop() assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for update weights from tensor" if obj.abort_all_requests: @@ -544,7 +545,7 @@ class TokenizerControlMixin: try: # For now, we only support single data parallel instance assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for update weights from IPC" logger.info("Starting IPC weight update") @@ -607,7 +608,7 @@ class TokenizerControlMixin: ) assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for dynamic lora loading" logger.info( "Start load Lora adapter. Lora name=%s, path=%s", @@ -685,7 +686,7 @@ class TokenizerControlMixin: ) assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for dynamic lora loading" logger.info( "Start load Lora adapter from tensors. Lora name=%s", @@ -765,7 +766,7 @@ class TokenizerControlMixin: ), "lora_name must be provided to unload LoRA adapter" assert ( - self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + get_parallel().dp_size == 1 or get_parallel().enable_dp_attention ), "dp_size must be 1 or dp attention must be enabled for dynamic lora loading" logger.info( "Start unload Lora adapter. Lora name=%s", @@ -785,7 +786,7 @@ class TokenizerControlMixin: self.auto_create_handle_loop() results = await self.get_weights_by_name_communicator(obj) all_parameters = [r.parameter for r in results] - if self.server_args.dp_size == 1: + if get_parallel().dp_size == 1: return all_parameters[0] else: return all_parameters diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 46718795c..d95293475 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -121,6 +121,7 @@ from sglang.srt.observability.request_metrics_exporter import ( RequestMetricsExporterManager, ) from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers +from sglang.srt.runtime_context import get_parallel from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import ( PortArgs, @@ -398,7 +399,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): set_global_server_args_for_tokenizer(server_args) self.startup_time: Optional[Dict[str, Any]] = None self._config_updates: List[Tuple[str, Dict[str, Any]]] = [] - self.elastic_worker_count = server_args.dp_size + self.elastic_worker_count = get_parallel().dp_size self.elastic_pending_ep_size = None self.elastic_scale_phase = "idle" self.elastic_last_error = None @@ -551,7 +552,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): self.tokenizer_ipc_name = port_args.tokenizer_ipc_name self.load_snapshot_reader = create_load_snapshot_reader( - self.server_args, port_args, caller="TokenizerManager", ) @@ -1583,7 +1583,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): return batch_size > 0 and ( self.server_args.enable_tokenizer_batch_encode or ( - (not self.server_args.enable_dp_attention) + (not get_parallel().enable_dp_attention) and (not self._batch_has_text(batch_size, requests)) ) ) diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index 3a6554bb5..2d7379780 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -240,10 +240,10 @@ def build_kv_cache( retraction_backup = resolve_decode_retraction_backup(tp_worker=tp_worker) - disable_radix_cache = server_args.disable_radix_cache or ( + disable_radix_cache = get_memory().disable_radix_cache or ( model_config.is_multimodal and uses_transformers_backend ) - if disable_radix_cache and not server_args.disable_radix_cache: + if disable_radix_cache and not get_memory().disable_radix_cache: logger.warning( "Radix cache is disabled for multimodal models with the " "Transformers backend to avoid multimodal prefix-cache mismatches." @@ -253,8 +253,8 @@ def build_kv_cache( # these use specialized memory pools incompatible with the # prefix-match-and-lock allocation path. if ( - server_args.disaggregation_decode_enable_radix_cache - and server_args.disaggregation_mode == "decode" + get_disagg().disaggregation_decode_enable_radix_cache + and get_disagg().disaggregation_mode == "decode" ): if is_hybrid_swa: raise ValueError( @@ -285,15 +285,15 @@ def build_kv_cache( ), is_eagle=spec_algorithm.is_eagle(), tp_cache_group=( - attn_tp_cpu_group if server_args.enable_dp_attention else tp_cpu_group + attn_tp_cpu_group if get_parallel().enable_dp_attention else tp_cpu_group ), attn_cp_cache_group=attn_cp_cpu_group, attn_tp_cache_group=attn_tp_cpu_group, pp_cache_group=pp_group.cpu_group, - eviction_policy=server_args.radix_eviction_policy, + eviction_policy=get_memory().radix_eviction_policy, enable_metrics=enable_metrics, enable_kv_cache_events=enable_kv_cache_events, - enable_session_radix_cache=server_args.enable_session_radix_cache, + enable_session_radix_cache=get_memory().enable_session_radix_cache, enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(), enable_mamba_extra_buffer_lazy=server_args.enable_mamba_extra_buffer_lazy(), pp_rank=ps.pp_rank, diff --git a/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py b/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py index 1939d347c..3b800fc4f 100644 --- a/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py +++ b/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py @@ -20,6 +20,7 @@ from sglang.srt.model_loader.weight_utils import ( CheckpointFilePrefetchHandle, ) from sglang.srt.platforms import current_platform +from sglang.srt.runtime_context import get_parallel if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig @@ -112,8 +113,8 @@ class StartupWeightLoadOptions: attn_cp_size=server_args.attn_cp_size, dcp_size=server_args.dcp_size, pp_size=server_args.pp_size, - dp_size=server_args.dp_size, - ep_size=server_args.ep_size, + dp_size=get_parallel().dp_size, + ep_size=get_parallel().ep_size, cpu_offload_gb=server_args.cpu_offload_gb, offload_group_size=server_args.offload_group_size, enable_memory_saver=server_args.enable_memory_saver, diff --git a/python/sglang/srt/ray/data_parallel_controller.py b/python/sglang/srt/ray/data_parallel_controller.py index 478571ea0..138f0e6fd 100644 --- a/python/sglang/srt/ray/data_parallel_controller.py +++ b/python/sglang/srt/ray/data_parallel_controller.py @@ -31,6 +31,7 @@ from sglang.srt.ray.engine import ( _get_bundle_node_ip, _resolve_bundle_indices, ) +from sglang.srt.runtime_context import get_parallel from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils.network import bind_port, get_zmq_socket, get_zmq_socket_on_host @@ -74,7 +75,7 @@ class RayDataParallelController(DataParallelController): sockets = [] dp_port_args_list = [] - for dp_rank in range(server_args.dp_size): + for dp_rank in range(get_parallel().dp_size): tmp_port_args = PortArgs.init_new(server_args) tmp_port_args.tokenizer_ipc_name = port_args.tokenizer_ipc_name tmp_port_args.detokenizer_ipc_name = port_args.detokenizer_ipc_name @@ -98,7 +99,7 @@ class RayDataParallelController(DataParallelController): sock.close() # Create actors for each DP rank sequentially - for dp_rank in range(server_args.dp_size): + for dp_rank in range(get_parallel().dp_size): self._launch_ray_tp_group(server_args, dp_port_args_list[dp_rank], dp_rank) def launch_dp_attention_schedulers( @@ -109,7 +110,7 @@ class RayDataParallelController(DataParallelController): # rank-0 node IP instead of tcp://* to avoid exposing unauthenticated # ZMQ sockets (CVE-2026-3060). worker_ports = [] - for dp_rank in range(server_args.dp_size): + for dp_rank in range(get_parallel().dp_size): worker_port, worker_socket = get_zmq_socket_on_host( self.context, zmq.PUSH, host=self.rank0_node_ip ) @@ -154,12 +155,12 @@ class RayDataParallelController(DataParallelController): tp_rank % tp_per_node ) - if server_args.enable_dp_attention: + if get_parallel().enable_dp_attention: _, _, actual_dp_rank, _ = compute_dp_attention_world_info( - server_args.enable_dp_attention, + get_parallel().enable_dp_attention, tp_rank, server_args.tp_size, - server_args.dp_size, + get_parallel().dp_size, server_args.attn_cp_size, ) rank_port_args = PortArgs.init_new( @@ -225,12 +226,12 @@ class RayDataParallelController(DataParallelController): bundle_idx = bundle_indices[global_rank] - if server_args.enable_dp_attention: + if get_parallel().enable_dp_attention: _, _, actual_dp_rank, _ = compute_dp_attention_world_info( - server_args.enable_dp_attention, + get_parallel().enable_dp_attention, tp_rank, server_args.tp_size, - server_args.dp_size, + get_parallel().dp_size, server_args.attn_cp_size, ) rank_port_args = PortArgs.init_new( diff --git a/python/sglang/srt/ray/engine.py b/python/sglang/srt/ray/engine.py index 95977c79c..f50523efa 100644 --- a/python/sglang/srt/ray/engine.py +++ b/python/sglang/srt/ray/engine.py @@ -32,6 +32,7 @@ from sglang.srt.entrypoints.engine import ( ) from sglang.srt.environ import envs from sglang.srt.ray.scheduler_actor import SchedulerActor +from sglang.srt.runtime_context import get_parallel from sglang.srt.server_args import PortArgs, ServerArgs logger = logging.getLogger(__name__) @@ -107,9 +108,9 @@ def _compute_world_size(server_args: ServerArgs) -> int: Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size. """ - if server_args.enable_dp_attention: + if get_parallel().enable_dp_attention: return server_args.tp_size * server_args.pp_size - return server_args.dp_size * server_args.tp_size * server_args.pp_size + return get_parallel().dp_size * server_args.tp_size * server_args.pp_size def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]: @@ -267,11 +268,11 @@ class RayEngine(Engine): placement_group as create_placement_group, ) - if server_args.enable_dp_attention: + if get_parallel().enable_dp_attention: total_gpus = server_args.tp_size * server_args.pp_size else: total_gpus = ( - server_args.dp_size * server_args.tp_size * server_args.pp_size + get_parallel().dp_size * server_args.tp_size * server_args.pp_size ) nnodes = server_args.nnodes @@ -313,7 +314,7 @@ class RayEngine(Engine): rank0_bundle_idx = int(indices_str.split(",")[0]) if indices_str else 0 rank0_node_ip = _get_bundle_node_ip(pg, rank0_bundle_idx) - if server_args.dp_size == 1: + if get_parallel().dp_size == 1: dist_init_addr = f"{rank0_node_ip}:{port_args.nccl_port}" logger.info(f"dist_init_addr: {dist_init_addr}") @@ -446,17 +447,19 @@ class RayEngine(Engine): RayDataParallelController, ) - if server_args.enable_dp_attention: + if get_parallel().enable_dp_attention: # DP attention folds DP into TP — total GPUs = tp_size * pp_size total_gpus = server_args.tp_size * server_args.pp_size else: - total_gpus = server_args.dp_size * server_args.tp_size * server_args.pp_size + total_gpus = ( + get_parallel().dp_size * server_args.tp_size * server_args.pp_size + ) gpus_per_node = total_gpus // server_args.nnodes logger.info( f"Ray DP cluster: {server_args.nnodes} nodes, " - f"{gpus_per_node} GPUs/node, dp_size={server_args.dp_size}, " + f"{gpus_per_node} GPUs/node, dp_size={get_parallel().dp_size}, " f"tp_size={server_args.tp_size}, pp_size={server_args.pp_size}, " - f"enable_dp_attention={server_args.enable_dp_attention}" + f"enable_dp_attention={get_parallel().enable_dp_attention}" ) # Set dist_init_addr on server_args so PortArgs.init_new() can compute diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index ee346be66..666cfd8c0 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -96,14 +96,14 @@ class DSparkWorkerV2(BaseSpecWorker): self._draft_is_moe = draft_is_deepseek_v4(server_args=server_args) self._draft_dp_context_enabled = ( - server_args.enable_dp_attention and not self._draft_is_moe + get_parallel().enable_dp_attention and not self._draft_is_moe ) self._is_pd_prefill = server_args.disaggregation_mode == "prefill" self._decode_graph_allowed = ( - not server_args.disable_cuda_graph and not self._is_pd_prefill + not get_exec().graph.disable_cuda_graph and not self._is_pd_prefill ) if ( - server_args.enable_dp_attention + get_parallel().enable_dp_attention and self._draft_is_moe and ps.attn_tp_size > 1 ): @@ -202,7 +202,7 @@ class DSparkWorkerV2(BaseSpecWorker): verify_num_draft_tokens=self.verify_num_draft_tokens, ) if ( - server_args.enable_dp_attention + get_parallel().enable_dp_attention and not self._draft_is_moe and self._verify_planner.is_compact_mode and self._decode_graph_allowed @@ -228,7 +228,7 @@ class DSparkWorkerV2(BaseSpecWorker): gamma=self.gamma, mask_token_id=self._mask_token_id, draft_block_spec_info=self._draft_block_spec_info, - dp_moe_sync=self._draft_is_moe and server_args.enable_dp_attention, + dp_moe_sync=self._draft_is_moe and get_parallel().enable_dp_attention, ) self._verify_epilogue = None if ( diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 12bda19f1..ad6ceb23e 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -158,7 +158,10 @@ class EagleDraftWorker(EagleDraftWorkerBase): self._rebuild_topk1_chain_buffers() # Load draft model weights only. - if server_args.enable_dp_attention and self.speculative_algorithm.is_eagle3(): + if ( + get_parallel().enable_dp_attention + and self.speculative_algorithm.is_eagle3() + ): ctx = draft_tp_context(get_parallel().attn_tp_group) else: ctx = empty_context() @@ -182,7 +185,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): # Eager draft-extend seed buffer (graph paths use their own static ones). self.dsa_extend_topk_buf: Optional[torch.Tensor] = None self.draft_tp_context = ( - draft_tp_context if server_args.enable_dp_attention else empty_context + draft_tp_context if get_parallel().enable_dp_attention else empty_context ) self.tree_mask_mode = default_tree_mask_mode() diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py index 427a416f2..34a96f0e9 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py @@ -44,7 +44,12 @@ from sglang.srt.model_executor.forward_batch_info import ( ) from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig -from sglang.srt.runtime_context import attention_backends, get_schedule, get_spec +from sglang.srt.runtime_context import ( + attention_backends, + get_parallel, + get_schedule, + get_spec, +) from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase from sglang.srt.speculative.eagle_utils import ( @@ -157,7 +162,7 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker): self.kv_context: Optional[FrozenKVMTPContext] = None self.draft_tp_context = ( - draft_tp_context if server_args.enable_dp_attention else empty_context + draft_tp_context if get_parallel().enable_dp_attention else empty_context ) self.draft_attn_backend = None diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 3af1dbc60..759d39250 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -43,7 +43,7 @@ from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, ForwardBatch, ) -from sglang.srt.runtime_context import get_schedule +from sglang.srt.runtime_context import get_parallel, get_schedule from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase from sglang.srt.speculative.draft_utils import DraftBackendFactory @@ -180,7 +180,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase): "InklingForConditionalGenerationMTP", ] self.draft_tp_context = ( - draft_tp_context if server_args.enable_dp_attention else empty_context + draft_tp_context if get_parallel().enable_dp_attention else empty_context ) self.tree_mask_mode = default_tree_mask_mode() self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device) diff --git a/python/sglang/srt/speculative/standalone_worker_v2.py b/python/sglang/srt/speculative/standalone_worker_v2.py index f0bfd52f5..37d809999 100644 --- a/python/sglang/srt/speculative/standalone_worker_v2.py +++ b/python/sglang/srt/speculative/standalone_worker_v2.py @@ -10,7 +10,7 @@ from sglang.srt.layers.moe.utils import ( speculative_moe_backend_context, ) from sglang.srt.managers.tp_worker import TpModelWorker -from sglang.srt.runtime_context import get_schedule +from sglang.srt.runtime_context import get_parallel, get_schedule from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.adaptive_runtime_state import ( AdaptiveController, @@ -89,7 +89,7 @@ class StandaloneDraftWorker(EagleDraftWorker): # Alias for better readability self.draft_runner = self.draft_worker.model_runner self.draft_tp_context = ( - draft_tp_context if server_args.enable_dp_attention else empty_context + draft_tp_context if get_parallel().enable_dp_attention else empty_context ) self.tree_mask_mode = default_tree_mask_mode() self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device) diff --git a/python/sglang/srt/utils/cuda_vmm_transport_utils.py b/python/sglang/srt/utils/cuda_vmm_transport_utils.py index e4258250c..d0ad2d012 100644 --- a/python/sglang/srt/utils/cuda_vmm_transport_utils.py +++ b/python/sglang/srt/utils/cuda_vmm_transport_utils.py @@ -31,7 +31,10 @@ from sglang.srt.managers.schedule_batch import ( MultimodalDataItem, MultimodalProcessorOutput, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import ( + configured_tp_size, + get_parallel, +) from sglang.srt.utils.cuda_ipc_transport_utils import ( DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY, MM_FEATURE_CACHE_SIZE, @@ -159,9 +162,9 @@ def _contains_tensor_container(value) -> bool: def get_vmm_feature_consumer_count(server_args) -> int: - if server_args.enable_dp_attention: - return server_args.tp_size // server_args.dp_size - return server_args.tp_size + if get_parallel().enable_dp_attention: + return configured_tp_size() // get_parallel().dp_size + return configured_tp_size() class CudaVmmMemoryPool: diff --git a/python/sglang/srt/utils/offloader.py b/python/sglang/srt/utils/offloader.py index 2e72b3c06..56b006a52 100644 --- a/python/sglang/srt/utils/offloader.py +++ b/python/sglang/srt/utils/offloader.py @@ -12,7 +12,11 @@ from sglang.srt.distributed.naive_distributed import ( set_naive_distributed, ) from sglang.srt.layers.parameter import ModelWeightParameter -from sglang.srt.runtime_context import get_parallel, get_stream +from sglang.srt.runtime_context import ( + get_exec, + get_parallel, + get_stream, +) from sglang.srt.server_args import ServerArgs from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available from sglang.srt.utils.host_shared_memory import ( @@ -63,21 +67,21 @@ def set_offloader(instance: BaseOffloader): def create_offloader_from_server_args(server_args: ServerArgs, dp_rank: int): - if server_args.cpu_offload_gb > 0: + if get_exec().offload.cpu_offload_gb > 0: return OffloaderV1( - cpu_offload_max_bytes=int(server_args.cpu_offload_gb * 1024**3) + cpu_offload_max_bytes=int(get_exec().offload.cpu_offload_gb * 1024**3) ) - if server_args.offload_group_size > 0: + if get_exec().offload.offload_group_size > 0: assert ( - server_args.cpu_offload_gb == 0 + get_exec().offload.cpu_offload_gb == 0 ), "V2 offload does not support cpu_offload_gb yet" return OffloaderV2( - group_size=server_args.offload_group_size, - num_in_group=server_args.offload_num_in_group, - prefetch_step=server_args.offload_prefetch_step, - mode=server_args.offload_mode, + group_size=get_exec().offload.offload_group_size, + num_in_group=get_exec().offload.offload_num_in_group, + prefetch_step=get_exec().offload.offload_prefetch_step, + mode=get_exec().offload.offload_mode, dp_rank=dp_rank, - dp_size=server_args.dp_size, + dp_size=get_parallel().dp_size, ) return NoopOffloader() diff --git a/test/registered/scheduler/test_prefill_delayer.py b/test/registered/scheduler/test_prefill_delayer.py index 64ce40841..de0b97d07 100644 --- a/test/registered/scheduler/test_prefill_delayer.py +++ b/test/registered/scheduler/test_prefill_delayer.py @@ -13,6 +13,7 @@ import torch from sglang.benchmark.serving import run_benchmark from sglang.srt.managers.prefill_delayer import PrefillDelayer +from sglang.srt.runtime_context import get_context from sglang.srt.utils import kill_process_tree from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.run_eval import run_eval @@ -76,6 +77,9 @@ def _run_negotiate_test(rank, test_cases): cpu_group = torch.distributed.new_group(backend="gloo") for case in test_cases: + # The DP-attention gate is a published config leaf. + override = get_context().override_server_args(enable_dp_attention=True) + override.install() delayer = PrefillDelayer( dp_size=world_size, attn_tp_size=1, @@ -127,6 +131,8 @@ def _run_negotiate_test(rank, test_cases): result.wait_seconds > 0.0 ), f"Case {case.name} rank {rank}: wait_seconds not surfaced" + override.restore() + _NEGOTIATE_TEST_CASES = [ NegotiateTestCase( diff --git a/test/registered/unit/entrypoints/test_http_server_warmup.py b/test/registered/unit/entrypoints/test_http_server_warmup.py index cf26faeb6..e4391e456 100644 --- a/test/registered/unit/entrypoints/test_http_server_warmup.py +++ b/test/registered/unit/entrypoints/test_http_server_warmup.py @@ -14,6 +14,12 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu") class TestDisaggregationServerWarmup(unittest.IsolatedAsyncioTestCase): async def test_sends_concurrent_scalar_request_to_each_dp_rank(self): + from sglang.srt.runtime_context import get_context + + # The warmup fan-out width comes from the published topology. + override = get_context().override_server_args(dp_size=4) + override.install() + self.addCleanup(override.restore) server_args = SimpleNamespace(dp_size=4) all_started = asyncio.Event() calls = [] diff --git a/test/registered/unit/managers/test_load_snapshot_backends.py b/test/registered/unit/managers/test_load_snapshot_backends.py index e689d4958..3e18f8e37 100644 --- a/test/registered/unit/managers/test_load_snapshot_backends.py +++ b/test/registered/unit/managers/test_load_snapshot_backends.py @@ -5,6 +5,7 @@ import tempfile import time import unittest from types import SimpleNamespace +from unittest import mock from sglang.srt.managers.load_snapshot import ( LoadSnapshot, @@ -18,6 +19,7 @@ from sglang.srt.managers.load_snapshot import ( should_use_zmq, zmq_reader_owner, ) +from sglang.srt.runtime_context import get_context from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel @@ -228,23 +230,17 @@ class TestZmqRoundTrip(CustomTestCase): class TestFactoryFunctions(CustomTestCase): + def _publish(self, **fields): + override = get_context().override_server_args(**fields) + override.install() + self.addCleanup(override.restore) + def test_shm_mode(self): - server_args = SimpleNamespace( - enable_dp_attention=False, - nnodes=1, - dp_size=1, - load_balance_method="round_robin", - node_rank=0, - tokenizer_worker_num=1, - ) + self._publish(enable_dp_attention=False, nnodes=1, dp_size=1) port_args = SimpleNamespace(instance_id="test_shm_factory") - writer = create_load_snapshot_writer( - server_args, port_args, dp_size=1, dp_rank=0 - ) + writer = create_load_snapshot_writer(port_args, dp_size=1, dp_rank=0) self.assertIsInstance(writer, ShmLoadSnapshotWriter) - reader = create_load_snapshot_reader( - server_args, port_args, caller="TokenizerManager" - ) + reader = create_load_snapshot_reader(port_args, caller="TokenizerManager") self.assertIsInstance(reader, ShmLoadSnapshotReader) reader.close() writer.close() @@ -255,24 +251,13 @@ class TestFactoryFunctions(CustomTestCase): os.unlink(path) def test_zmq_mode_via_env(self): - server_args = SimpleNamespace( - enable_dp_attention=False, - nnodes=1, - dp_size=1, - load_balance_method="round_robin", - node_rank=0, - tokenizer_worker_num=1, - ) + self._publish(enable_dp_attention=False, nnodes=1, dp_size=1) port_args = SimpleNamespace(instance_id="test_zmq_factory") os.environ["SGLANG_LOAD_SNAPSHOT_USE_ZMQ"] = "1" try: - writer = create_load_snapshot_writer( - server_args, port_args, dp_size=1, dp_rank=0 - ) + writer = create_load_snapshot_writer(port_args, dp_size=1, dp_rank=0) self.assertIsInstance(writer, ZmqLoadSnapshotWriter) - reader = create_load_snapshot_reader( - server_args, port_args, caller="TokenizerManager" - ) + reader = create_load_snapshot_reader(port_args, caller="TokenizerManager") self.assertIsInstance(reader, ZmqShmLoadSnapshotReader) reader.close() writer.close() @@ -280,8 +265,8 @@ class TestFactoryFunctions(CustomTestCase): del os.environ["SGLANG_LOAD_SNAPSHOT_USE_ZMQ"] def test_should_use_zmq_multinode_dp_attention(self): - args = SimpleNamespace(enable_dp_attention=True, nnodes=2) - self.assertTrue(should_use_zmq(args)) + self._publish(enable_dp_attention=True, nnodes=2, dp_size=2) + self.assertTrue(should_use_zmq()) class TestZmqReaderOwner(CustomTestCase): @@ -289,9 +274,9 @@ class TestZmqReaderOwner(CustomTestCase): CALLERS = ("TokenizerManager", "MultiTokenizerRouter", "DataParallelController") - @staticmethod - def _args(**overrides): - base = dict( + def _owners(self, **overrides): + """Publish a config and return the callers that claim the socket.""" + fields = dict( enable_dp_attention=True, nnodes=2, node_rank=0, @@ -299,54 +284,87 @@ class TestZmqReaderOwner(CustomTestCase): load_balance_method="round_robin", tokenizer_worker_num=1, ) - base.update(overrides) - return SimpleNamespace(**base) - - def _owners(self, args): - return {c for c in self.CALLERS if zmq_reader_owner(args, c)} + fields.update(overrides) + override = get_context().override_server_args(**fields) + override.install() + try: + return {c for c in self.CALLERS if zmq_reader_owner(c)} + finally: + override.restore() def test_zmq_disabled_no_owner(self): - args = self._args(enable_dp_attention=False, nnodes=1) - self.assertEqual(self._owners(args), set()) + self.assertEqual(self._owners(enable_dp_attention=False, nnodes=1), set()) def test_non_zero_node_rank_no_owner(self): - args = self._args(node_rank=1, dp_size=4, tokenizer_worker_num=8) - self.assertEqual(self._owners(args), set()) + self.assertEqual( + self._owners(node_rank=1, dp_size=4, tokenizer_worker_num=8), set() + ) def test_tokenizer_manager_owns_when_dp1(self): - self.assertEqual(self._owners(self._args(dp_size=1)), {"TokenizerManager"}) + self.assertEqual(self._owners(dp_size=1), {"TokenizerManager"}) def test_multi_tokenizer_router_owns_in_multi_tokenizer_dp1(self): - args = self._args(dp_size=1, tokenizer_worker_num=8) - self.assertEqual(self._owners(args), {"MultiTokenizerRouter"}) + self.assertEqual( + self._owners(dp_size=1, tokenizer_worker_num=8), {"MultiTokenizerRouter"} + ) def test_multi_tokenizer_router_owns_in_multi_tokenizer_round_robin(self): - args = self._args(dp_size=4, tokenizer_worker_num=8) - self.assertEqual(self._owners(args), {"MultiTokenizerRouter"}) + self.assertEqual( + self._owners(dp_size=4, tokenizer_worker_num=8), {"MultiTokenizerRouter"} + ) def test_data_parallel_controller_owns_load_aware(self): for method in ("total_tokens", "total_requests"): - args = self._args( - dp_size=4, tokenizer_worker_num=8, load_balance_method=method + self.assertEqual( + self._owners( + dp_size=4, tokenizer_worker_num=8, load_balance_method=method + ), + {"DataParallelController"}, ) - self.assertEqual(self._owners(args), {"DataParallelController"}) def test_tokenizer_manager_owns_dp4_round_robin(self): - args = self._args(dp_size=4, tokenizer_worker_num=1) - self.assertEqual(self._owners(args), {"TokenizerManager"}) + self.assertEqual( + self._owners(dp_size=4, tokenizer_worker_num=1), {"TokenizerManager"} + ) + + def test_the_controller_answers_within_its_audited_namespaces(self): + """The DP controller publishes with a narrowed namespace set. + + Under `SGLANG_ROLE_NAMESPACES=enforce` a read outside that set raises, + and this decision runs during its startup -- so asking which tokenizer + process owns the socket would abort the controller before it has a + reader. + """ + import sglang.srt.runtime_context as rc + + fields = dict( + enable_dp_attention=True, + nnodes=2, + node_rank=0, + dp_size=4, + load_balance_method="total_tokens", + tokenizer_worker_num=8, + ) + override = get_context().override_server_args(**fields) + override.install() + self.addCleanup(override.restore) + with mock.patch.object(rc, "_ROLE_NS_MODE", "enforce"), mock.patch.object( + rc._CONTEXT, "_publish_role", "dp_controller" + ): + self.assertTrue(zmq_reader_owner("DataParallelController")) def test_at_most_one_owner_across_configs(self): for dp_size in (1, 4): for tw in (1, 8): for method in ("round_robin", "total_tokens", "total_requests"): for node_rank in (0, 1): - args = self._args( + owners = self._owners( dp_size=dp_size, tokenizer_worker_num=tw, load_balance_method=method, node_rank=node_rank, ) - self.assertLessEqual(len(self._owners(args)), 1, args) + self.assertLessEqual(len(owners), 1, owners) class TestZmqAddr(CustomTestCase): diff --git a/test/registered/unit/managers/test_scheduler_recv_skipper.py b/test/registered/unit/managers/test_scheduler_recv_skipper.py index b57c6749c..14305017c 100644 --- a/test/registered/unit/managers/test_scheduler_recv_skipper.py +++ b/test/registered/unit/managers/test_scheduler_recv_skipper.py @@ -1,6 +1,7 @@ import unittest from types import SimpleNamespace +from sglang.srt.runtime_context import get_context from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel @@ -14,11 +15,15 @@ from sglang.srt.model_executor.forward_batch_info import ForwardMode # noqa: E4 register_cpu_ci(est_time=2, suite="base-a-test-cpu") -def _server_args(interval, enable_dp_attention=False): - return SimpleNamespace( +def _publish(case, interval, enable_dp_attention=False): + """Publish the config the skipper reads; the double is a published config, + not an injected object.""" + override = get_context().override_server_args( scheduler_recv_interval=interval, enable_dp_attention=enable_dp_attention, ) + override.install() + case.addCleanup(override.restore) def _batch(forward_mode, recv_skipper_forward_mode=None): @@ -31,22 +36,24 @@ def _batch(forward_mode, recv_skipper_forward_mode=None): class TestSchedulerRecvSkipper(CustomTestCase): def test_disabled_at_default_interval(self): # interval <= 1 disables the skipper entirely. - self.assertIsNone(SchedulerRecvSkipper.maybe_create(_server_args(1))) + _publish(self, 1) + self.assertIsNone(SchedulerRecvSkipper.maybe_create()) def test_enabled_under_dp_attention(self): # Regression: the constructor used to assert `not enable_dp_attention`. - skipper = SchedulerRecvSkipper.maybe_create( - _server_args(50, enable_dp_attention=True) - ) + _publish(self, 50, enable_dp_attention=True) + skipper = SchedulerRecvSkipper.maybe_create() self.assertIsNotNone(skipper) def test_no_last_batch_accumulates_slowly(self): - skipper = SchedulerRecvSkipper.maybe_create(_server_args(50)) + _publish(self, 50) + skipper = SchedulerRecvSkipper.maybe_create() self.assertFalse(skipper.handle(None)) def test_decode_accumulates_until_threshold(self): # DECODE weight = 1: recv only every `interval` decode steps. - skipper = SchedulerRecvSkipper.maybe_create(_server_args(3)) + _publish(self, 3) + skipper = SchedulerRecvSkipper.maybe_create() decode = _batch(ForwardMode.DECODE) self.assertFalse(skipper.handle(decode)) # counter 1 self.assertFalse(skipper.handle(decode)) # counter 2 @@ -55,21 +62,20 @@ class TestSchedulerRecvSkipper(CustomTestCase): def test_prefill_triggers_recv_immediately(self): # Non-decode passes use the large default weight: recv right away. - skipper = SchedulerRecvSkipper.maybe_create(_server_args(50)) + _publish(self, 50) + skipper = SchedulerRecvSkipper.maybe_create() self.assertTrue(skipper.handle(_batch(ForwardMode.EXTEND))) def test_dp_uses_synced_mode_not_local(self): # Local EXTEND (weight 1000) must be ignored in favor of the synced # DECODE (weight 1); a recv here would mean the local mode leaked in. - skipper = SchedulerRecvSkipper.maybe_create( - _server_args(50, enable_dp_attention=True) - ) + _publish(self, 50, enable_dp_attention=True) + skipper = SchedulerRecvSkipper.maybe_create() self.assertFalse(skipper.handle(_batch(ForwardMode.EXTEND, ForwardMode.DECODE))) def test_dp_synced_extend_triggers_recv(self): - skipper = SchedulerRecvSkipper.maybe_create( - _server_args(50, enable_dp_attention=True) - ) + _publish(self, 50, enable_dp_attention=True) + skipper = SchedulerRecvSkipper.maybe_create() self.assertTrue(skipper.handle(_batch(ForwardMode.IDLE, ForwardMode.EXTEND))) def test_derive_forward_mode(self): diff --git a/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py b/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py index e119e0898..ba5a924dc 100644 --- a/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py +++ b/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py @@ -28,7 +28,7 @@ from sglang.srt.model_executor.model_runner_components.startup_weight_load impor ) from sglang.srt.model_loader.loader import DefaultModelLoader from sglang.srt.model_loader.weight_utils import initialize_capture_safe_weights -from sglang.srt.runtime_context import get_context +from sglang.srt.runtime_context import get_context, publish, reset_context from sglang.srt.server_args import ServerArgs register_cpu_ci(est_time=5, suite="base-a-test-cpu") @@ -203,10 +203,14 @@ class TestStartupWeightLoadSelector(CustomTestCase): def test_options_accept_current_server_args_schema(self): """Removed server options must not break overlap startup initialization.""" + server_args = ServerArgs( + model_path="dummy", cuda_graph_config=CudaGraphConfig() + ) + # The parallel sizes come from the bags, so the config has to be published. + publish(server_args, role="test") + self.addCleanup(reset_context) options = StartupWeightLoadOptions.from_server_args( - server_args=ServerArgs( - model_path="dummy", cuda_graph_config=CudaGraphConfig() - ), + server_args=server_args, is_draft_worker=False, ) diff --git a/test/registered/unit/multimodal/test_gpu_feature_transport.py b/test/registered/unit/multimodal/test_gpu_feature_transport.py index c2e7e57c4..d104b27fe 100644 --- a/test/registered/unit/multimodal/test_gpu_feature_transport.py +++ b/test/registered/unit/multimodal/test_gpu_feature_transport.py @@ -83,16 +83,22 @@ class TestCudaVmmFeatureTransport(unittest.TestCase): get_model_architecture.assert_not_called() def test_vmm_transport_initializes_pool(self): + from sglang.srt.runtime_context import get_context from sglang.srt.utils import cuda_vmm_transport_utils as vmm server_args = SimpleNamespace( mm_feature_transport="cuda_vmm", tokenizer_worker_num=2, base_gpu_id=3, - enable_dp_attention=False, tp_size=4, nnodes=1, ) + # The consumer count comes from the published topology. + override = get_context().override_server_args( + enable_dp_attention=False, tp_size=4 + ) + override.install() + self.addCleanup(override.restore) pool = object() with ( patch.object(vmm, "get_mm_feature_pool_size_per_worker", return_value=123), diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index d221af1c2..2b718ca7a 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -76,6 +76,21 @@ _CONFIGURED_SIZE_CALL_SITES = { "the one where initialize_model_parallel aliases _MOE_DP to _ATTN_CP, so " "the live sizes are equal there and a live comparison is always false" ), + ("srt/managers/scheduler.py", "configured_tp_size"): ( + "configure_scheduler_process runs before the scheduler's own process " + "groups exist -- configuring the process is what it is for -- so there " + "is nothing live to ask yet" + ), + ("srt/managers/scheduler.py", "configured_moe_dp_size"): ( + "same pre-distributed-init arithmetic in configure_scheduler_process" + ), + ("srt/managers/scheduler.py", "configured_attn_cp_size"): ( + "same pre-distributed-init arithmetic in configure_scheduler_process" + ), + ("srt/utils/cuda_vmm_transport_utils.py", "configured_tp_size"): ( + "the consumer count is configured fan-out arithmetic (tp_size // " + "dp_size), which is what the record answered before" + ), ("srt/model_loader/loader.py", "configured_moe_dp_size"): ( "the same dict already carries the live moe_dp_size under 'dp'; this entry " "is the configured intent" diff --git a/test/registered/unit/test_publish_precedes_bag_reads.py b/test/registered/unit/test_publish_precedes_bag_reads.py index 2bae967b5..1dee9c003 100644 --- a/test/registered/unit/test_publish_precedes_bag_reads.py +++ b/test/registered/unit/test_publish_precedes_bag_reads.py @@ -104,10 +104,6 @@ _UNREAD_ENTRIES: dict = { ("multimodal_gen/test/unit/test_disagg_trace.py", "_srt_trace_server_args"): ( "a trace fixture publishing its own context" ), - ("srt/entrypoints/engine.py", "_launch_subprocesses"): ( - "its subprocess targets and its tokenizer-manager factory arrive as " - "parameters, so the walk resolves none of them" - ), ("srt/managers/detokenizer_manager.py", "run_detokenizer_process"): ( "DetokenizerManager reads the handed instance at this revision" ), diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index 0b77d800d..8426d3940 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -155,7 +155,6 @@ _EXPOSED = { ("disaggregation/encode_receiver.py", "tokenizer_path"), ("disaggregation/encode_server.py", "allowed_media_domains"), ("disaggregation/encode_server.py", "device"), - ("disaggregation/encode_server.py", "dp_size"), ("disaggregation/encode_server.py", "encoder_transfer_backend"), ("disaggregation/encode_server.py", "load_format"), ("disaggregation/encode_server.py", "mm_process_config"), @@ -203,12 +202,9 @@ _EXPOSED = { ("elastic_ep/expert_backup_manager.py", "mooncake_ib_device"), ("entrypoints/engine.py", "attn_cp_size"), ("entrypoints/engine.py", "detokenizer_worker_num"), - ("entrypoints/engine.py", "dp_size"), ("entrypoints/engine.py", "dtype"), - ("entrypoints/engine.py", "enable_dp_attention"), ("entrypoints/engine.py", "enable_symm_mem"), ("entrypoints/engine.py", "ep_join_mode"), - ("entrypoints/engine.py", "ep_size"), ("entrypoints/engine.py", "load_format"), ("entrypoints/engine.py", "model_path"), ("entrypoints/engine.py", "moe_dp_size"), @@ -221,7 +217,6 @@ _EXPOSED = { ), ("entrypoints/engine.py", "tool_call_parser"), ("entrypoints/http_server.py", "disaggregation_mode"), - ("entrypoints/http_server.py", "dp_size"), ("entrypoints/http_server.py", "ep_join_mode"), ("entrypoints/http_server.py", "grpc_port"), ("entrypoints/http_server.py", "model_path"), @@ -247,10 +242,6 @@ _EXPOSED = { ("layers/cp/base.py", "enable_prefill_cp"), ("layers/cp/bcg.py", "cp_strategy"), ("layers/cp/bcg.py", "enable_prefill_cp"), - ("layers/dp_attention.py", "attn_cp_size"), - ("layers/dp_attention.py", "device"), - ("layers/dp_attention.py", "dp_size"), - ("layers/dp_attention.py", "enable_dp_attention"), ("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"), ("layers/moe/kt_ep_wrapper.py", "chunked_prefill_size"), ("layers/moe/utils.py", "deepep_mode"), @@ -259,19 +250,11 @@ _EXPOSED = { ("layers/moe/utils.py", "quantization"), ("layers/moe/utils.py", "speculative_moe_runner_backend"), ("layers/quantization/unquant.py", "enable_deterministic_inference"), - ("lora/lora_manager.py", "enable_dp_attention"), ("lora/lora_manager.py", "enable_lora_overlap_loading"), ("lora/marlin_lora_temp/policy.py", "enable_lora"), ("lora/marlin_lora_temp/policy.py", "lora_paths"), ("managers/data_parallel_controller.py", "attn_cp_size"), ("managers/data_parallel_controller.py", "disaggregation_mode"), - ("managers/data_parallel_controller.py", "dp_size"), - ("managers/data_parallel_controller.py", "enable_dp_attention"), - ( - "managers/data_parallel_controller.py", - "enable_dp_attention_local_control_broadcast", - ), - ("managers/data_parallel_controller.py", "ep_size"), ("managers/data_parallel_controller.py", "load_balance_method"), ("managers/data_parallel_controller.py", "moe_dp_size"), ("managers/data_parallel_controller.py", "pp_size"), @@ -282,23 +265,16 @@ _EXPOSED = { ("managers/disagg_service.py", "disaggregation_bootstrap_port"), ("managers/disagg_service.py", "disaggregation_mode"), ("managers/disagg_service.py", "disaggregation_transfer_backend"), - ("managers/load_snapshot.py", "dp_size"), - ("managers/load_snapshot.py", "enable_dp_attention"), - ("managers/load_snapshot.py", "load_balance_method"), ("managers/overlap_utils.py", "speculative_algorithm"), ("managers/prefill_delayer.py", "disable_overlap_schedule"), - ("managers/prefill_delayer.py", "enable_dp_attention"), ("managers/rust_server.py", "mm_process_config"), ("managers/schedule_batch.py", "disaggregation_mode"), ("managers/scheduler.py", "attn_cp_size"), ("managers/scheduler.py", "disable_overlap_schedule"), ("managers/scheduler.py", "disaggregation_mode"), - ("managers/scheduler.py", "dp_size"), - ("managers/scheduler.py", "enable_dp_attention"), ("managers/scheduler.py", "enable_hierarchical_cache"), ("managers/scheduler.py", "enable_lora"), ("managers/scheduler.py", "enable_lora_overlap_loading"), - ("managers/scheduler.py", "ep_size"), ("managers/scheduler.py", "moe_dp_size"), ("managers/scheduler.py", "pp_size"), ("managers/scheduler.py", "soft_watchdog_timeout"), @@ -307,13 +283,9 @@ _EXPOSED = { "managers/scheduler_components/new_token_ratio_tracker.py", "schedule_conservativeness", ), - ("managers/scheduler_components/recv_skipper.py", "enable_dp_attention"), - ("managers/tokenizer_control_mixin.py", "dp_size"), ("managers/tokenizer_manager.py", "disable_radix_cache"), ("managers/tokenizer_manager.py", "disaggregation_mode"), ("managers/tokenizer_manager.py", "disaggregation_transfer_backend"), - ("managers/tokenizer_manager.py", "dp_size"), - ("managers/tokenizer_manager.py", "enable_dp_attention"), ("managers/tokenizer_manager.py", "enable_lora"), ("managers/tokenizer_manager.py", "enable_tokenizer_batch_encode"), ("managers/tokenizer_manager.py", "encoder_transfer_backend"), @@ -339,9 +311,6 @@ _EXPOSED = { ("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "hicache_io_backend"), ("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "hicache_mem_layout"), ("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "served_model_name"), - ("mem_cache/kv_cache_builder.py", "disable_radix_cache"), - ("mem_cache/kv_cache_builder.py", "disaggregation_mode"), - ("mem_cache/kv_cache_builder.py", "enable_dp_attention"), ("mem_cache/kv_cache_builder.py", "hicache_mem_layout"), ("mem_cache/radix_cache_cpp.py", "enable_hierarchical_cache"), ("model_executor/forward_batch_info.py", "enable_return_hidden_states"), @@ -373,9 +342,7 @@ _EXPOSED = { "custom_weight_loader", ), ("model_executor/model_runner_components/startup_weight_load.py", "device"), - ("model_executor/model_runner_components/startup_weight_load.py", "dp_size"), ("model_executor/model_runner_components/startup_weight_load.py", "enable_lora"), - ("model_executor/model_runner_components/startup_weight_load.py", "ep_size"), ("model_executor/model_runner_components/startup_weight_load.py", "lora_paths"), ("model_executor/model_runner_components/startup_weight_load.py", "pp_size"), ( @@ -400,11 +367,7 @@ _EXPOSED = { ("observability/metrics_collector.py", "served_model_name"), ("parser/template_detection.py", "model_path"), ("ray/data_parallel_controller.py", "attn_cp_size"), - ("ray/data_parallel_controller.py", "dp_size"), - ("ray/data_parallel_controller.py", "enable_dp_attention"), ("ray/data_parallel_controller.py", "pp_size"), - ("ray/engine.py", "dp_size"), - ("ray/engine.py", "enable_dp_attention"), ("ray/engine.py", "pp_size"), ("speculative/adaptive_spec_params.py", "speculative_algorithm"), ("speculative/adaptive_spec_params.py", "speculative_eagle_topk"), @@ -416,29 +379,24 @@ _EXPOSED = { "speculative/dspark_components/dspark_config.py", "speculative_draft_model_revision", ), - ("speculative/dspark_components/dspark_worker_v2.py", "disable_cuda_graph"), ("speculative/dspark_components/dspark_worker_v2.py", "disaggregation_mode"), - ("speculative/dspark_components/dspark_worker_v2.py", "enable_dp_attention"), ( "speculative/dspark_components/dspark_worker_v2.py", "speculative_num_draft_tokens", ), ("speculative/eagle_worker_v2.py", "device"), - ("speculative/eagle_worker_v2.py", "enable_dp_attention"), ("speculative/eagle_worker_v2.py", "speculative_adaptive"), ("speculative/eagle_worker_v2.py", "speculative_algorithm"), ("speculative/eagle_worker_v2.py", "speculative_eagle_topk"), ("speculative/eagle_worker_v2.py", "speculative_num_draft_tokens"), ("speculative/eagle_worker_v2.py", "speculative_num_steps"), ("speculative/frozen_kv_mtp_worker_v2.py", "device"), - ("speculative/frozen_kv_mtp_worker_v2.py", "enable_dp_attention"), ("speculative/frozen_kv_mtp_worker_v2.py", "speculative_adaptive"), ("speculative/frozen_kv_mtp_worker_v2.py", "speculative_algorithm"), ("speculative/frozen_kv_mtp_worker_v2.py", "speculative_eagle_topk"), ("speculative/frozen_kv_mtp_worker_v2.py", "speculative_num_draft_tokens"), ("speculative/frozen_kv_mtp_worker_v2.py", "speculative_num_steps"), ("speculative/multi_layer_eagle_worker_v2.py", "device"), - ("speculative/multi_layer_eagle_worker_v2.py", "enable_dp_attention"), ("speculative/multi_layer_eagle_worker_v2.py", "speculative_algorithm"), ("speculative/multi_layer_eagle_worker_v2.py", "speculative_eagle_topk"), ("speculative/multi_layer_eagle_worker_v2.py", "speculative_num_draft_tokens"), @@ -451,7 +409,6 @@ _EXPOSED = { ("speculative/spec_info.py", "enable_multi_layer_eagle"), ("speculative/spec_registry.py", "disable_overlap_schedule"), ("speculative/standalone_worker_v2.py", "device"), - ("speculative/standalone_worker_v2.py", "enable_dp_attention"), ("speculative/standalone_worker_v2.py", "speculative_algorithm"), ("speculative/standalone_worker_v2.py", "speculative_eagle_topk"), ("speculative/standalone_worker_v2.py", "speculative_num_draft_tokens"), @@ -460,11 +417,8 @@ _EXPOSED = { ("utils/common.py", "speculative_eagle_topk"), ("utils/common.py", "speculative_num_draft_tokens"), ("utils/common.py", "speculative_num_steps"), - ("utils/cuda_vmm_transport_utils.py", "dp_size"), - ("utils/cuda_vmm_transport_utils.py", "enable_dp_attention"), ("utils/cuda_vmm_transport_utils.py", "mm_feature_transport"), ("utils/hf_transformers/processor.py", "image_processor_backend"), - ("utils/offloader.py", "dp_size"), } # Pairs whose resolution write only happens on a CUDA host (capability or @@ -487,7 +441,6 @@ _OVERRIDDEN_AND_READ = { "disaggregation/decode_kvcache_offload_manager.py", "hicache_storage_backend_extra_config", ), - ("disaggregation/encode_server.py", "dp_size"), ("disaggregation/encode_server.py", "load_format"), ("disaggregation/encode_server.py", "model_path"), ( @@ -496,24 +449,13 @@ _OVERRIDDEN_AND_READ = { ), ("dllm/config.py", "model_path"), ("elastic_ep/expert_backup_manager.py", "load_format"), - ("entrypoints/engine.py", "dp_size"), ("entrypoints/engine.py", "dtype"), - ("entrypoints/engine.py", "ep_size"), ("entrypoints/engine.py", "load_format"), ("entrypoints/engine.py", "model_path"), - ("entrypoints/http_server.py", "dp_size"), ("entrypoints/http_server.py", "model_path"), ("kv_canary/api.py", "speculative_num_steps"), ("kv_canary/capacities.py", "speculative_num_draft_tokens"), - ("layers/dp_attention.py", "dp_size"), - ("managers/data_parallel_controller.py", "dp_size"), - ("managers/data_parallel_controller.py", "ep_size"), - ("managers/load_snapshot.py", "dp_size"), - ("managers/scheduler.py", "dp_size"), - ("managers/scheduler.py", "ep_size"), ("managers/scheduler.py", "hicache_storage_backend"), - ("managers/tokenizer_control_mixin.py", "dp_size"), - ("managers/tokenizer_manager.py", "dp_size"), ("managers/tokenizer_manager.py", "model_path"), ("managers/tokenizer_manager.py", "speculative_num_draft_tokens"), ("managers/tp_worker.py", "model_path"), @@ -532,13 +474,8 @@ _OVERRIDDEN_AND_READ = { ("mem_cache/unified_radix_cache.py", "hicache_storage_prefetch_policy"), ("mem_cache/unified_radix_cache.py", "hicache_write_policy"), ("model_executor/model_runner_components/load_model_utils.py", "load_format"), - ("model_executor/model_runner_components/startup_weight_load.py", "dp_size"), - ("model_executor/model_runner_components/startup_weight_load.py", "ep_size"), ("parser/template_detection.py", "model_path"), - ("ray/data_parallel_controller.py", "dp_size"), - ("ray/engine.py", "dp_size"), ("speculative/dflash_worker_v2.py", "speculative_num_draft_tokens"), - ("speculative/dspark_components/dspark_worker_v2.py", "disable_cuda_graph"), ( "speculative/dspark_components/dspark_worker_v2.py", "speculative_num_draft_tokens", @@ -555,8 +492,6 @@ _OVERRIDDEN_AND_READ = { ("speculative/standalone_worker_v2.py", "speculative_num_steps"), ("utils/common.py", "speculative_num_draft_tokens"), ("utils/common.py", "speculative_num_steps"), - ("utils/cuda_vmm_transport_utils.py", "dp_size"), - ("utils/offloader.py", "dp_size"), }