From cba3c5d5ac4beef366872194aef53ace9ea7337c Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Mon, 17 Aug 2026 16:17:53 -0700 Subject: [PATCH] config: the per-instance families read the bags (#35026) --- .../skills/sglang-runtime-context/SKILL.md | 43 +-- .../srt/constrained/base_grammar_backend.py | 27 +- .../srt/disaggregation/encode_receiver.py | 9 +- .../srt/disaggregation/encode_server.py | 81 ++--- python/sglang/srt/entrypoints/engine.py | 39 ++- python/sglang/srt/entrypoints/http_server.py | 49 +-- python/sglang/srt/entrypoints/sidecar.py | 2 + .../srt/kv_canary/token_oracle/install.py | 12 +- .../srt/managers/detokenizer_manager.py | 8 +- python/sglang/srt/managers/scheduler.py | 4 +- .../batch_result_processor.py | 2 +- .../new_token_ratio_tracker.py | 6 +- .../sglang/srt/managers/tokenizer_manager.py | 84 +++--- .../srt/model_executor/cpu_graph_runner.py | 2 +- .../srt/model_executor/forward_batch_info.py | 11 +- .../sglang/srt/model_executor/model_runner.py | 1 - .../cuda_graph_setup.py | 3 +- .../srt/model_executor/runner/base_runner.py | 4 +- python/sglang/srt/models/sarvam_moe.py | 21 +- .../sglang/srt/multimodal/cache/identity.py | 23 +- .../multimodal/processors/base_processor.py | 9 +- .../srt/multimodal/processors/mimo_v2.py | 3 +- .../srt/ray/data_parallel_controller.py | 17 +- python/sglang/srt/ray/engine.py | 18 +- .../mock_model/test_self_unit_install.py | 21 +- .../constrained/test_base_grammar_backend.py | 52 +++- .../test_kimi_k3_encoder_mode.py | 51 ++-- ...st_batch_result_processor_hidden_states.py | 16 +- .../managers/test_hidden_state_server_mode.py | 17 +- .../unit/managers/test_mm_process_config.py | 35 ++- .../test_tokenizer_manager_rid_cleanup.py | 52 ++-- .../test_hidden_state_graph_recapture.py | 47 ++- .../test_prefill_cuda_graph_runner.py | 13 +- test/registered/unit/models/test_kimi_k25.py | 203 +++++++------ .../unit/multimodal/rust/qwen/_fixtures.py | 16 +- .../multimodal/rust/qwen/test_e2e_parity.py | 4 +- .../rust/qwen/test_native_mm_host.py | 2 +- .../multimodal/test_gpu_feature_transport.py | 10 +- .../unit/multimodal/test_preprocess_cache.py | 41 +-- .../unit/server_args/test_server_args.py | 1 + .../unit/test_global_config_read_ratchet.py | 21 ++ ...test_launch_path_reads_configured_sizes.py | 280 ++++++++++++++++++ .../unit/test_publish_precedes_bag_reads.py | 3 - .../test_split_attention_backend_decisions.py | 1 + ...test_supplied_instance_exposure_ratchet.py | 75 +---- 45 files changed, 909 insertions(+), 530 deletions(-) create mode 100644 test/registered/unit/test_launch_path_reads_configured_sizes.py diff --git a/.claude/skills/sglang-runtime-context/SKILL.md b/.claude/skills/sglang-runtime-context/SKILL.md index 730175ede..115d06684 100644 --- a/.claude/skills/sglang-runtime-context/SKILL.md +++ b/.claude/skills/sglang-runtime-context/SKILL.md @@ -11,7 +11,7 @@ One container owns process-static runtime state: `sglang.srt.runtime_context.Run | Tier | Accessor | Holds | Lifecycle | |------|----------|-------|-----------| | raw config seed | `get_server_args()` | the published `ServerArgs` — the startup record, for debugging, dumps and provenance. **Business code does not read fields off it**: the read ratchet pins that at zero, and "Reading config: the seed is off limits" below says what to read instead, which forms the ratchet sees, and what is outside it by construction (a runtime-computed name; a whole-object hand-off) | published at process entry; re-publish is **last-publish-wins** (the tokenizer publish in the launcher process; sequential engine rebuild in one process, e.g. unit tests) and re-projects the bags; read-only | -| resolved config | `get_exec()` `get_memory()` `get_schedule()` `get_model()` `get_spec()` `get_serving()` `get_observability()` `get_disagg()` `get_lora()` `get_mm()` `get_device()` | namespace **config bags** — the single source of truth for resolved config; leaves are real attributes (dynamo-traceable) | projected from `server_args` at `publish`; mutated only via `get_context().override` | +| resolved config | `get_exec()` `get_memory()` `get_schedule()` `get_model()` `get_spec()` `get_serving()` `get_observability()` `get_disagg()` `get_lora()` `get_mm()` `get_device()` | namespace **config bags** — the single source of truth for resolved config; leaves are real attributes (dynamo-traceable). Each is a **module function of no arguments**, and a module binds the name once: `manager.get_disagg()`, `self.get_disagg = get_disagg`, or a same-named import next to the bag one (`from model_loader import get_model`) all import fine and fail only when that path runs. `ruff --select F811` catches the import collision; `RuntimeContext` has no bag-named member and no `__getattr__`, so the member-call shapes are an `AttributeError` at call time — give it a delegating `__getattr__` and they go silent instead | projected from `server_args` at `publish`; mutated only via `get_context().override` | | runtime flags | `get_flags()` | state that is *not* a pure function of config: `capture` (cuda-graph lifecycle), `moe` (ACTIVE backends, swappable), `dp` (DP-attention runtime flags) | materialized at subsystem init; groups offer `override()` for tests | | resources | `get_resources()`, `get_stream(name)`, `get_buffer(name, factory)` | process-level handles: graph pools, EPLB state, EP dispatcher state, named side streams, workspace buffers | lazy; cleared by `reset_context()` | | per-forward | `get_forward()` | forward-scoped flags (multi-stream switch, MoE output buffer, attn-TP inputs, extend-in-batch) | contextvar-backed; `scoped(**kw)` restores on exit; new threads see defaults | @@ -108,13 +108,14 @@ bag to override at all. the scope. When there is a runner in hand, read its stamp; that is a different rule from "read the instance". - **Per-instance boundaries** — the tokenizer-manager family, everything under - `entrypoints/`, and the tokenizer-process multimodal processors still read - `self.server_args` today. The old justification ("several `Engine`s can share - one process, bags are last-publish-wins across them") is **retracted** — owner - ruling (2026-08-15): a process holds at most one live config at a time - (concurrent multi-Engine is unsupported; sequential rebuild stays legal, unit - tests rely on it). These reads are scheduled to become bag reads in the - bag-read series; treat them as pinned debt, not as a boundary to imitate. What + `entrypoints/`, and the tokenizer-process multimodal processors read the bags. + The old justification for keeping them on `self.server_args` ("several + `Engine`s can share one process, bags are last-publish-wins across them") is + **retracted** — owner ruling (2026-08-15): a process holds at most one live + config at a time (concurrent multi-Engine is unsupported; sequential rebuild + stays legal, unit tests rely on it). What still reads the instance in those + files is pinned pair by pair in the exposure ratchet, each with its own + disposition; none of it is a boundary to imitate. What genuinely stays per-instance is what differs per *worker* within one engine: `base_gpu_id` travels as a constructor argument (`MMEncoder(gpu_id=...)`; `BaseMultimodalProcessor._fast_image_processor_device` is the shape to copy). @@ -226,24 +227,28 @@ this). "Reads that legitimately stay on a ServerArgs instance" above for the full set — per-instance boundaries and whole-object passes; there are no per-runner config copies to read any more). The allow-list is `GrammarManager` and `MMEncoder`; -the tokenizer-manager family, `entrypoints/`, and the tokenizer-process -multimodal processors sit beside it only as pinned debt — not for one single -reason: +what sits beside it is residue, not a family — and not for one single reason: -- the tokenizer-manager family and `entrypoints/` are **pinned debt awaiting - conversion to bag reads** (the old multi-Engine justification is retracted — - one process, one live config); the reads still work today because the - instance carries resolved values; -- `GrammarManager` is a handed instance — it is constructed with the config its - owner hands it and never assumes a published namespace; +- the tokenizer-manager family and `entrypoints/` **read the bags**; what is + left of them in the exposure ratchet is a handful of individually-dispositioned + pairs, not a family awaiting conversion. Read the ratchet for the current set + rather than assuming a directory is off-limits; +- `GrammarManager` is a handed instance for its residual `self.server_args` + reads, but backend selection is **not** on the instance any more: + `create_grammar_backend` reads `get_exec().kernel.grammar_backend`, and + `__init__` calls that factory whenever `skip_tokenizer_init` is false. In + production the scheduler process has published; a test that constructs one + without publishing has to keep patching the factory (or publish itself); - `MMEncoder` publishes the very instance it is handed (`publish(server_args, role="encoder")`) and takes its per-worker device as a separate `gpu_id` argument, so its `self.server_args` reads and the bag agree today. They are on this list as a construction-path convention rather than a semantic exception — and the residual is real: a post-publish `override` would not reach them. -Their tests tell you the same thing: they construct the object standalone, so a -bag read turns into "config namespace not published". +Their tests are not one story: a `GrammarManager` built standalone turns the +factory's bag read into "config namespace not published" unless the test patches +it or publishes, while `MMEncoder` publishes in its own `__init__` and so needs +no such arrangement. **Test doubles publish, they do not inject.** A stand-in that carries `server_args=SimpleNamespace(field=...)` stops working the moment production reads diff --git a/python/sglang/srt/constrained/base_grammar_backend.py b/python/sglang/srt/constrained/base_grammar_backend.py index 70a544724..8fe2215fc 100644 --- a/python/sglang/srt/constrained/base_grammar_backend.py +++ b/python/sglang/srt/constrained/base_grammar_backend.py @@ -22,7 +22,12 @@ from typing import Dict, List, NamedTuple, Optional, Tuple import torch from sglang.srt.parser.reasoning_parser import ReasoningParser -from sglang.srt.runtime_context import get_context, get_resources +from sglang.srt.runtime_context import ( + get_context, + get_exec, + get_resources, + get_serving, +) from sglang.srt.server_args import ServerArgs logger = logging.getLogger(__name__) @@ -315,7 +320,7 @@ def create_grammar_backend( eos_token_ids: Optional[set] = None, think_end_ids: Optional[List[int]] = None, ) -> Optional[BaseGrammarBackend]: - name = server_args.grammar_backend + name = get_exec().kernel.grammar_backend # Custom grammar backend has the highest priority if name in GRAMMAR_BACKEND_REGISTRY: @@ -329,7 +334,7 @@ def create_grammar_backend( grammar_backend = OutlinesGrammarBackend( tokenizer, - whitespace_pattern=server_args.constrained_json_whitespace_pattern, + whitespace_pattern=get_serving().constrained_json_whitespace_pattern, ) elif name == "xgrammar": from sglang.srt.constrained.xgrammar_backend import ( @@ -345,10 +350,10 @@ def create_grammar_backend( tokenizer, vocab_size=vocab_size, model_eos_token_ids=eos_list, - any_whitespace=not server_args.constrained_json_disable_any_whitespace, + any_whitespace=not get_serving().constrained_json_disable_any_whitespace, ) except TokenizerNotSupportedError as e: - if server_args.enable_strict_thinking: + if get_serving().enable_strict_thinking: raise ValueError( f"--enable-strict-thinking requires a grammar backend with " f"token filtering support, but XGrammar failed to initialize: " @@ -367,13 +372,13 @@ def create_grammar_backend( grammar_backend = GuidanceBackend( tokenizer=tokenizer, - any_whitespace=not server_args.constrained_json_disable_any_whitespace, - whitespace_pattern=server_args.constrained_json_whitespace_pattern, + any_whitespace=not get_serving().constrained_json_disable_any_whitespace, + whitespace_pattern=get_serving().constrained_json_whitespace_pattern, n_vocab=vocab_size, eos_token_ids=eos_token_ids, ) elif name == "none": - if server_args.enable_strict_thinking: + if get_serving().enable_strict_thinking: raise ValueError( "--enable-strict-thinking requires a grammar backend that supports " "token filtering, but grammar_backend='none' was specified. Use " @@ -384,13 +389,13 @@ def create_grammar_backend( else: raise ValueError(f"Invalid grammar backend: {name}") - if server_args.reasoning_parser and think_end_ids: + if get_serving().reasoning_parser and think_end_ids: from sglang.srt.constrained.reasoner_grammar_backend import ( ReasonerGrammarBackend, ) reasoning_parser = ReasoningParser( - model_type=server_args.reasoning_parser, + model_type=get_serving().reasoning_parser, stream_reasoning=False, tokenizer=tokenizer, ) @@ -399,7 +404,7 @@ def create_grammar_backend( grammar_backend, reasoning_parser, tokenizer, - enable_strict_thinking=server_args.enable_strict_thinking, + enable_strict_thinking=get_serving().enable_strict_thinking, ) return grammar_backend diff --git a/python/sglang/srt/disaggregation/encode_receiver.py b/python/sglang/srt/disaggregation/encode_receiver.py index 0718b1d1c..041353aa3 100644 --- a/python/sglang/srt/disaggregation/encode_receiver.py +++ b/python/sglang/srt/disaggregation/encode_receiver.py @@ -34,6 +34,7 @@ from sglang.srt.managers.io_struct import GenerateReqInput, TokenizedGenerateReq from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors from sglang.srt.managers.schedule_batch import Modality, Req from sglang.srt.multimodal.cache import media_preprocess_kwargs +from sglang.srt.runtime_context import get_disagg, get_exec, get_serving from sglang.srt.server_args import ServerArgs from sglang.srt.utils import ImageData from sglang.srt.utils.common import safe_pickle_loads @@ -1586,7 +1587,7 @@ class MMReceiverBase(ABC): # context alive for the process instead of creating a temporary context # whose destruction also closes its per-request socket. self.scheduler_context = zmq.Context() - self.encoder_transfer_backend = server_args.encoder_transfer_backend + self.encoder_transfer_backend = get_disagg().encoder_transfer_backend # When ``encode_urls`` is shared with an :class:`EncoderBootstrapServer` # (tokenizer manager process), it grows / shrinks in place as encoders # register or unregister; the receiver always sees the current set. @@ -1642,8 +1643,8 @@ class MMReceiverBase(ABC): self.embeddings_engine = init_mooncake_transfer_engine( hostname=self.host, ib_device=( - server_args.disaggregation_ib_device - or server_args.mooncake_ib_device + get_disagg().disaggregation_ib_device + or get_exec().moe.mooncake_ib_device ), ) self.embeddings_buffer = dict() @@ -1689,7 +1690,7 @@ class MMReceiverBase(ABC): extra_kwargs["tokenizer_backend"] = server_args.tokenizer_backend _processor = get_processor( - server_args.tokenizer_path, + get_serving().tokenizer_path, tokenizer_mode=server_args.tokenizer_mode, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 7f4d9c2e6..a1e340d49 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -57,7 +57,7 @@ from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalSta from sglang.srt.model_executor.model_runner_components.load_model_utils import ( maybe_precompile_model_kernels_after_loading, ) -from sglang.srt.model_loader import get_model +from sglang.srt.model_loader import get_model as load_model from sglang.srt.multimodal.cache import parse_content_hash, snapshot_media from sglang.srt.multimodal.encoder_preprocessing import ( EncoderPreprocessOutput, @@ -73,10 +73,15 @@ from sglang.srt.observability.trace import ( trace_set_thread_info, ) from sglang.srt.runtime_context import ( + configured_tp_size, + get_device, get_disagg, get_exec, get_mm, + get_model, + get_observability, get_parallel, + get_serving, publish, ) from sglang.srt.server_args import ( @@ -315,7 +320,7 @@ class MMEncoder: logger.info(f"init MMEncoder {rank}/{server_args.tp_size}") self.server_args = server_args configure_media_url_security( - server_args.allowed_media_domains, + get_mm().allowed_media_domains, server_args.media_url_max_file_size_mb, ) self.rank = rank @@ -329,7 +334,7 @@ class MMEncoder: server_args, ) self.load_config = LoadConfig( - load_format=server_args.load_format, + load_format=get_model().load_format, download_dir=server_args.download_dir, model_loader_extra_config=server_args.model_loader_extra_config, remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip, @@ -340,7 +345,7 @@ class MMEncoder: self.model_config.hf_config, "model_type", "unknown" ).lower() - self.device = server_args.device + self.device = get_device().device self.gpu_id = server_args.base_gpu_id + rank if gpu_id is None else gpu_id self.device_config = DeviceConfig( @@ -354,7 +359,7 @@ class MMEncoder: use_image_processor_gpu and resolve_image_processor_backend(server_args) != "pil" ) - self._build_vision_config(server_args.mm_process_config) + self._build_vision_config(get_mm().mm_process_config) self.model_audio_sr = self._resolve_audio_sr() logger.info(f"Resolved model audio sample rate: {self.model_audio_sr} Hz") @@ -368,7 +373,7 @@ class MMEncoder: initialize_model_parallel(tensor_model_parallel_size=server_args.tp_size) initialize_dp_attention(server_args, self.model_config) - self.model = get_model( + self.model = load_model( model_config=self.model_config, load_config=self.load_config, device_config=self.device_config, @@ -628,7 +633,7 @@ class MMEncoder: ) try: self.image_processor = AutoImageProcessor.from_pretrained( - server_args.tokenizer_path or server_args.model_path, + get_serving().tokenizer_path or get_model().model_path, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, **image_processor_kwargs, @@ -639,7 +644,7 @@ class MMEncoder: try: self.video_processor = AutoVideoProcessor.from_pretrained( - server_args.tokenizer_path or server_args.model_path, + get_serving().tokenizer_path or get_model().model_path, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, ) @@ -650,7 +655,7 @@ class MMEncoder: try: # Note: AutoProcessor is used for audio processor _audio_proc = AutoProcessor.from_pretrained( - server_args.tokenizer_path or server_args.model_path, + get_serving().tokenizer_path or get_model().model_path, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, ) @@ -2100,7 +2105,7 @@ class MMEncoder: _zmq_xfer_start = time.perf_counter() if ( - self.server_args.encoder_transfer_backend == "zmq_to_scheduler" + get_disagg().encoder_transfer_backend == "zmq_to_scheduler" and url is not None ): lock = self.scheduler_send_locks.get(endpoint) @@ -2146,7 +2151,7 @@ class MMEncoder: if encoder_metrics_collector is not None: encoder_metrics_collector.observe_transfer( time.perf_counter() - _zmq_xfer_start, - backend=self.server_args.encoder_transfer_backend, + backend=get_disagg().encoder_transfer_backend, ) return @@ -3007,7 +3012,7 @@ async def _push_embedding_to_prefill(enc: MMEncoder, request: dict) -> None: # No-op for mooncake (its /send is separate). embedding_port=None is # rejected upfront, so ports is always a concrete list here. req_id = request["req_id"] - backend = enc.server_args.encoder_transfer_backend + backend = get_disagg().encoder_transfer_backend if backend == "zmq_to_tokenizer": await enc.send( @@ -3050,7 +3055,7 @@ async def _dp_worker_encode_and_send( modality = Modality.from_str(request["modality"]) time_stats.modality = modality.name.lower() time_stats.set_metrics_collector(encoder_metrics_collector) - backend = enc.server_args.encoder_transfer_backend + backend = get_disagg().encoder_transfer_backend # URL state lives in main process module globals; workers don't see it. if backend == "zmq_to_scheduler" and request.get("embedding_port") is None: @@ -3666,14 +3671,14 @@ async def run_dp_worker( ) global encoder_metrics_collector - if server_args.enable_metrics: + if get_observability().enable_metrics: set_prometheus_multiproc_dir() labels = { - "model_name": server_args.served_model_name, + "model_name": get_serving().served_model_name, "dp_rank": str(dp_rank), } - if server_args.extra_metric_labels: - labels.update(server_args.extra_metric_labels) + if get_observability().extra_metric_labels: + labels.update(get_observability().extra_metric_labels) encoder_metrics_collector = EncoderMetricsCollector(labels) enc.dp_rank = dp_rank @@ -3964,14 +3969,14 @@ def launch_server(server_args: ServerArgs): global encoder, encoder_metrics_collector # Set up prometheus metrics. - if server_args.enable_metrics: + if get_observability().enable_metrics: set_prometheus_multiproc_dir() labels = { - "model_name": server_args.served_model_name, + "model_name": get_serving().served_model_name, "dp_rank": "0", } - if server_args.extra_metric_labels: - labels.update(server_args.extra_metric_labels) + if get_observability().extra_metric_labels: + labels.update(get_observability().extra_metric_labels) encoder_metrics_collector = EncoderMetricsCollector(labels) add_prometheus_middleware(app) @@ -3979,21 +3984,21 @@ def launch_server(server_args: ServerArgs): zmq_ctx = zmq.Context(10) ipc_path_prefix = random_uuid() port_args = PortArgs.init_new(server_args) - if server_args.dist_init_addr: - na = NetworkAddress.parse(server_args.dist_init_addr) + if get_parallel().dist_init_addr: + na = NetworkAddress.parse(get_parallel().dist_init_addr) dist_init_method = na.to_tcp() else: dist_init_method = NetworkAddress( - server_args.host or "127.0.0.1", port_args.nccl_port + get_serving().host or "127.0.0.1", port_args.nccl_port ).to_tcp() - if server_args.enable_trace: + if get_observability().enable_trace: process_tracing_init( - server_args.otlp_traces_endpoint, + get_observability().otlp_traces_endpoint, "sglang", - trace_modules=server_args.trace_modules, + trace_modules=get_observability().trace_modules, ) trace_set_thread_info("Encoder") - for rank in range(1, server_args.tp_size): + for rank in range(1, configured_tp_size()): schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}" send_sockets.append( get_zmq_socket(zmq_ctx, zmq.PUSH, schedule_path, bind=False) @@ -4006,13 +4011,13 @@ def launch_server(server_args: ServerArgs): encoder = MMEncoder(server_args, dist_init_method=dist_init_method) # Register this encoder's URL with prefill server(s) if configured. - if server_args.encoder_register_urls: + if get_disagg().encoder_register_urls: import atexit _register_encoder_url_with_bootstrap(server_args) atexit.register(_unregister_encoder_url_from_bootstrap, server_args) - uvicorn.run(app, host=server_args.host, port=server_args.port) + uvicorn.run(app, host=get_serving().host, port=get_serving().port) def _launch_server_dp(server_args: ServerArgs): @@ -4083,7 +4088,7 @@ def _launch_server_dp(server_args: ServerArgs): proc.start() worker_processes.append(proc) - labels = {"model_name": server_args.served_model_name} + labels = {"model_name": get_serving().served_model_name} if server_args.extra_metric_labels: labels.update(server_args.extra_metric_labels) dp_dispatcher = DPDispatcher( @@ -4188,7 +4193,7 @@ async def handle_encode_request(request: dict): # when multiple decoder TP ranks POST /encode # with the same req_id, only the first triggers the VIT forward; # subsequent callers wait and return the same metadata. - if encoder.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": async with encoder._inflight_encode_lock: if req_id in encoder._inflight_encode_events: event = encoder._inflight_encode_events[req_id] @@ -4274,7 +4279,7 @@ async def handle_encode_request(request: dict): time_stats.set_mm_encode_end_time() if error_msg: - if encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler": + if get_disagg().encoder_transfer_backend == "zmq_to_scheduler": if request["embedding_port"] is None: start_background_send(req_id) else: @@ -4285,7 +4290,7 @@ async def handle_encode_request(request: dict): embedding_port=port, ) # Signal waiters on failure for mooncake - if encoder.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": encoder._inflight_encode_meta.pop(req_id, None) evt = encoder._inflight_encode_events.pop(req_id, None) if evt: @@ -4299,7 +4304,7 @@ async def handle_encode_request(request: dict): status_code=error_code, content={"status": "error", "message": error_msg, "req_id": req_id}, ) - if encoder.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": # Store metadata for duplicate callers and signal them encoder._inflight_encode_meta[req_id] = ( nbytes, @@ -4323,7 +4328,7 @@ async def handle_encode_request(request: dict): modality=modality_str, status="success" ) return ORJSONResponse(content=request) - elif encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler": + elif get_disagg().encoder_transfer_backend == "zmq_to_scheduler": logger.info(f"{request['embedding_port'] = }") if request["embedding_port"] is None: await encoder.send_with_url( @@ -4347,7 +4352,7 @@ async def handle_encode_request(request: dict): modality=modality_str, status="success" ) return ORJSONResponse(content=None) - elif encoder.server_args.encoder_transfer_backend == "zmq_to_tokenizer": + elif get_disagg().encoder_transfer_backend == "zmq_to_tokenizer": await encoder.send( req_id=request["req_id"], prefill_host=request["prefill_host"], @@ -4369,7 +4374,7 @@ async def handle_encode_request(request: dict): logger.error(f"Unexpected error in encoder logic for {req_id}: {error_msg}") rid_to_err_msg[req_id] = error_msg # Ensure inflight waiters are unblocked on unexpected errors - if encoder.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": encoder._inflight_encode_meta.pop(req_id, None) evt = encoder._inflight_encode_events.pop(req_id, None) if evt: diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index b53b06208..43bb046ee 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -99,7 +99,14 @@ 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 get_parallel, publish +from sglang.srt.runtime_context import ( + configured_pp_size, + get_exec, + get_model, + get_parallel, + get_serving, + publish, +) from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import ( MultiprocessingSerializer, @@ -158,9 +165,9 @@ def init_tokenizer_manager( template_manager = TemplateManager() template_manager.initialize_templates( tokenizer_manager=tokenizer_manager, - model_path=server_args.model_path, - chat_template=server_args.chat_template, - completion_template=server_args.completion_template, + model_path=get_model().model_path, + chat_template=get_serving().chat_template, + completion_template=get_serving().completion_template, ) # Resolve any remaining auto parsers using template manager's detection results @@ -681,7 +688,7 @@ class Engine(EngineScoreMixin, EngineBase): pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = ( _calculate_rank_ranges( server_args.nnodes, - server_args.pp_size, + configured_pp_size(), tp_size, server_args.node_rank, ) @@ -702,7 +709,7 @@ class Engine(EngineScoreMixin, EngineBase): daemon_procs = [] logger.info( f"Launching {num_daemons} weight cache daemon(s) on node " - f"{server_args.node_rank} for model={server_args.model_path}, " + f"{server_args.node_rank} for model={get_model().model_path}, " f"pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, " f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, " f"dist_init_method={dist_init_method}" @@ -737,7 +744,7 @@ class Engine(EngineScoreMixin, EngineBase): "-m", "sglang.srt.weight_cache.daemon", "--model-path", - server_args.model_path, + get_model().model_path, "--gpu-id", str(gpu_id), "--tp-size", @@ -745,7 +752,7 @@ class Engine(EngineScoreMixin, EngineBase): "--tp-rank", str(tp_rank), "--pp-size", - str(server_args.pp_size), + str(configured_pp_size()), "--pp-rank", str(pp_rank), "--dp-size", @@ -753,14 +760,14 @@ class Engine(EngineScoreMixin, EngineBase): "--ep-size", str(get_parallel().ep_size), "--load-format", - server_args.load_format, + get_model().load_format, "--dtype", - server_args.dtype, + get_model().dtype, "--dist-init-method", dist_init_method, ] - if server_args.quantization: - cmd += ["--quantization", server_args.quantization] + if get_model().quantization: + cmd += ["--quantization", get_model().quantization] if ( server_args.model_loader_extra_config and server_args.model_loader_extra_config != "{}" @@ -863,7 +870,7 @@ class Engine(EngineScoreMixin, EngineBase): """ scheduler_procs = [] use_dp_controller = ( - get_parallel().dp_size > 1 or server_args.ep_join_mode == "scale" + get_parallel().dp_size > 1 or get_exec().moe.ep_join_mode == "scale" ) if not use_dp_controller: @@ -876,7 +883,7 @@ class Engine(EngineScoreMixin, EngineBase): pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = ( _calculate_rank_ranges( server_args.nnodes, - server_args.pp_size, + configured_pp_size(), server_args.tp_size, server_args.node_rank, ) @@ -983,7 +990,7 @@ class Engine(EngineScoreMixin, EngineBase): processes: List[mp.Process] = [] names: List[str] = [] - if server_args.detokenizer_worker_num <= 1: + if get_serving().detokenizer_worker_num <= 1: proc = mp.Process( target=run_detokenizer_process_func, args=(server_args, port_args), @@ -996,7 +1003,7 @@ class Engine(EngineScoreMixin, EngineBase): router_ipc_name = port_args.detokenizer_ipc_name worker_ipc_names: List[str] = [] try: - for i in range(server_args.detokenizer_worker_num): + for i in range(get_serving().detokenizer_worker_num): worker_ipc = f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}" port_args.detokenizer_ipc_name = worker_ipc proc = mp.Process( diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 6002989ee..8e8ffee20 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -247,7 +247,7 @@ async def init_multi_tokenizer() -> ServerArgs: template_manager = TemplateManager() template_manager.initialize_templates( tokenizer_manager=tokenizer_manager, - model_path=server_args.model_path, + model_path=get_model().model_path, chat_template=server_args.chat_template, completion_template=server_args.completion_template, ) @@ -293,9 +293,9 @@ async def lifespan(fast_api_app: FastAPI): "sglang", trace_modules=server_args.trace_modules, ) - if server_args.disaggregation_mode == "prefill": + if get_disagg().disaggregation_mode == "prefill": thread_label = "Prefill" + thread_label - elif server_args.disaggregation_mode == "decode": + elif get_disagg().disaggregation_mode == "decode": thread_label = "Decode" + thread_label trace_set_thread_info(thread_label) @@ -380,7 +380,7 @@ async def lifespan(fast_api_app: FastAPI): # Execute custom warmups if server_args.warmups is not None: await execute_warmups( - server_args.disaggregation_mode, + get_disagg().disaggregation_mode, server_args.warmups.split(","), _global_state.tokenizer_manager, ) @@ -393,7 +393,7 @@ async def lifespan(fast_api_app: FastAPI): try: if ( getattr(fast_api_app, "is_single_tokenizer_mode", False) - and server_args.grpc_port is not None + and get_serving().grpc_port is not None and not (server_args.smg_grpc_mode or server_args.grpc_mode) ): grpc_handle = _start_native_grpc_server_for_runtime( @@ -401,6 +401,7 @@ async def lifespan(fast_api_app: FastAPI): tokenizer_manager=_global_state.tokenizer_manager, template_manager=_global_state.template_manager, scheduler_info=_global_state.scheduler_info, + grpc_port=get_serving().grpc_port, ) if server_args.sidecar is not None: from sglang.srt.entrypoints.sidecar import start_sidecar @@ -480,7 +481,13 @@ 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 +from sglang.srt.runtime_context import ( + get_disagg, + get_exec, + get_model, + get_parallel, + get_serving, +) elastic_ep_router.route_class = ORJSONRoute app.include_router(elastic_ep_router) @@ -679,10 +686,7 @@ async def health_generate(request: Request) -> Response: sampling_params=sampling_params, log_metrics=False, ) - if ( - _global_state.tokenizer_manager.server_args.disaggregation_mode - != DisaggregationMode.NULL.value - ): + if get_disagg().disaggregation_mode != DisaggregationMode.NULL.value: gri.bootstrap_host = FAKE_BOOTSTRAP_HOST gri.bootstrap_room = 0 else: @@ -2224,7 +2228,7 @@ def _execute_server_warmup(server_args: ServerArgs): json_data["input_ids"] = json_data["input_ids"][0] elif ( is_vlm - and server_args.disaggregation_mode == "null" + and get_disagg().disaggregation_mode == "null" and model_info["is_generation"] ): served_model_name = "" @@ -2234,9 +2238,9 @@ def _execute_server_warmup(server_args: ServerArgs): # _global_state.tokenizer_manager is not initialized in the rust server, # so we need to get the model name from the model_info served_model_name = model_info.get( - "model_path", server_args.served_model_name + "model_path", get_serving().served_model_name ) - served_model_name = served_model_name or server_args.model_path + served_model_name = served_model_name or get_model().model_path # TODO: ChatCompletionRequest does not have bootstrap info required by disaggregation mode, disable image-warmup for now # Only use chat completions format for generation models, not embedding models json_data = { @@ -2280,7 +2284,7 @@ def _execute_server_warmup(server_args: ServerArgs): # Send a warmup request warmup_timeout = envs.SGLANG_WARMUP_TIMEOUT.get() try: - if server_args.disaggregation_mode == "null": + if get_disagg().disaggregation_mode == "null": res = requests.post( url + request_name, json=json_data, @@ -2314,7 +2318,7 @@ def _execute_server_warmup(server_args: ServerArgs): else: logger.info( "Disaggregation warmup failed (mode=%s), status codes: %s", - server_args.disaggregation_mode, + get_disagg().disaggregation_mode, failed_status_codes, ) # In rust-server mode there is no TokenizerManager (readiness is @@ -2368,10 +2372,10 @@ def _wait_and_warmup( logger.debug( "[Elastic EP] Skipping server warmup for elastic joiner " "(ep_join_mode=%s)", - server_args.ep_join_mode, + get_exec().moe.ep_join_mode, ) - if not server_args.skip_server_warmup and not skip_elastic_joiner_warmup: + if not get_serving().skip_server_warmup and not skip_elastic_joiner_warmup: if not execute_warmup_func(server_args): return else: @@ -2383,7 +2387,7 @@ def _wait_and_warmup( logger.info("The server is fired up and ready to roll!") if server_args.delete_ckpt_after_loading: - delete_directory(server_args.model_path) + delete_directory(get_model().model_path) if server_args.debug_tensor_dump_input_file: kill_process_tree(os.getpid()) @@ -2711,6 +2715,7 @@ def _start_native_grpc_server_for_runtime( tokenizer_manager, template_manager, scheduler_info, + grpc_port, ): from sglang.srt.entrypoints.grpc_bridge import RuntimeHandle from sglang.srt.rust_extensions import load_rust_extension @@ -2726,13 +2731,11 @@ def _start_native_grpc_server_for_runtime( grpc_handle = grpc_native.start_server( host=server_args.host, - port=server_args.grpc_port, + port=grpc_port, runtime_handle=runtime_handle, worker_threads=server_args.grpc_worker_threads, ) - logger.info( - f"Native gRPC server started on {server_args.host}:{server_args.grpc_port}" - ) + logger.info(f"Native gRPC server started on {server_args.host}:{grpc_port}") return grpc_handle @@ -2790,7 +2793,7 @@ def launch_server( # and /get_model_info endpoints are static (200 as soon as the server # binds, before any forward pass), so without this the first real request # pays the cold-start cost (observed as a >60s first generation). - if not server_args.skip_server_warmup: + if not get_serving().skip_server_warmup: _execute_server_warmup(server_args) logger.info("The server is fired up and ready to roll!") if launch_callback is not None: diff --git a/python/sglang/srt/entrypoints/sidecar.py b/python/sglang/srt/entrypoints/sidecar.py index ddf58b81e..13dd3d1f6 100644 --- a/python/sglang/srt/entrypoints/sidecar.py +++ b/python/sglang/srt/entrypoints/sidecar.py @@ -37,6 +37,8 @@ def _loopback_host(host: str) -> str: def build_sidecar_endpoint(server_args) -> str: + """Both halves of the endpoint come from the argument: this is a helper + over a config object, callable before anything is published.""" return NetworkAddress( _loopback_host(server_args.host), server_args.grpc_port ).to_url() diff --git a/python/sglang/srt/kv_canary/token_oracle/install.py b/python/sglang/srt/kv_canary/token_oracle/install.py index d88cc010d..ab48f7ee1 100644 --- a/python/sglang/srt/kv_canary/token_oracle/install.py +++ b/python/sglang/srt/kv_canary/token_oracle/install.py @@ -1,21 +1,17 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Optional +from typing import Optional from sglang.srt.kv_canary.token_oracle.oracle import HashOracle from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager from sglang.srt.kv_canary.token_oracle.sampler import install_oracle_sampler - -if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs +from sglang.srt.runtime_context import get_exec -def install_token_oracle_from_env( - *, server_args: ServerArgs, vocab_size: int -) -> Optional[TokenOracleManager]: +def install_token_oracle_from_env(*, vocab_size: int) -> Optional[TokenOracleManager]: # Must be called before create_sampler() so the factory is present when the # Sampler is first constructed. - if server_args.sampling_backend != "token_oracle": + if get_exec().kernel.sampling_backend != "token_oracle": return None oracle = HashOracle(vocab_size=vocab_size) diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index ad92bb817..2a506f88c 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -39,7 +39,7 @@ from sglang.srt.managers.io_struct import ( ) from sglang.srt.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread -from sglang.srt.runtime_context import publish +from sglang.srt.runtime_context import get_device, get_serving, publish from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import configure_logger, freeze_gc, kill_itself_when_parent_died from sglang.srt.utils.hf_transformers_utils import get_tokenizer @@ -128,7 +128,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): self.vocab_size = None else: self.tokenizer = get_tokenizer( - server_args.tokenizer_path, + get_serving().tokenizer_path, tokenizer_mode=server_args.tokenizer_mode, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, @@ -142,11 +142,11 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): def init_running_status(self, server_args: ServerArgs): self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES) self.disable_tokenizer_batch_decode = server_args.disable_tokenizer_batch_decode - self.is_tool_call_parser_gpt_oss = server_args.tool_call_parser == "gpt-oss" + self.is_tool_call_parser_gpt_oss = get_serving().tool_call_parser == "gpt-oss" self.soft_watchdog = Watchdog.create( debug_name="DetokenizerManager", - watchdog_timeout=server_args.soft_watchdog_timeout, + watchdog_timeout=get_device().soft_watchdog_timeout, soft=True, test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(), ) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index e94f7040e..7dc5be833 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1252,9 +1252,7 @@ class Scheduler( and not get_schedule().disable_priority_preemption ) - self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args( - self.server_args - ) + self.new_token_ratio_tracker = NewTokenRatioTracker.from_config() def init_soft_watchdog(self, server_args: ServerArgs): if (x := server_args.soft_watchdog_timeout) is not None: diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index c98becd59..024781c37 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -566,7 +566,7 @@ class SchedulerBatchResultProcessor: return get_required_capture_hidden_mode( max( batch.return_hidden_states_mode, - get_server_return_hidden_states_mode(server_args), + get_server_return_hidden_states_mode(), ), batch.spec_info, ) diff --git a/python/sglang/srt/managers/scheduler_components/new_token_ratio_tracker.py b/python/sglang/srt/managers/scheduler_components/new_token_ratio_tracker.py index 65e899953..8a771a5b0 100644 --- a/python/sglang/srt/managers/scheduler_components/new_token_ratio_tracker.py +++ b/python/sglang/srt/managers/scheduler_components/new_token_ratio_tracker.py @@ -4,7 +4,7 @@ from dataclasses import dataclass from typing import TYPE_CHECKING, Sequence from sglang.srt.environ import envs -from sglang.srt.server_args import ServerArgs +from sglang.srt.runtime_context import get_schedule if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req @@ -18,10 +18,10 @@ class NewTokenRatioTracker: current: float @classmethod - def from_server_args(cls, server_args: ServerArgs) -> NewTokenRatioTracker: + def from_config(cls) -> NewTokenRatioTracker: init = min( envs.SGLANG_INIT_NEW_TOKEN_RATIO.get() - * server_args.schedule_conservativeness, + * get_schedule().schedule_conservativeness, 1.0, ) min_ratio = min( diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index d95293475..33bd40f30 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -121,7 +121,18 @@ 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.runtime_context import ( + get_device, + get_disagg, + get_exec, + get_lora, + get_memory, + get_mm, + get_model, + get_parallel, + get_serving, + get_spec, +) from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import ( PortArgs, @@ -160,12 +171,13 @@ _REQUEST_STATE_WAIT_TIMEOUT = envs.SGLANG_REQUEST_STATE_WAIT_TIMEOUT.get() logger = logging.getLogger(__name__) -def _reject_missing_dispatched_encoder_embedding(server_args, request_obj, mm_inputs): +def _reject_missing_dispatched_encoder_embedding(request_obj, mm_inputs): """Do not silently turn a failed EPD request into local vision work.""" + disagg = get_disagg() if ( mm_inputs is None - and server_args.language_only - and server_args.encoder_transfer_backend == "zmq_to_tokenizer" + and disagg.language_only + and disagg.encoder_transfer_backend == "zmq_to_tokenizer" and request_obj.need_wait_for_mm_inputs ): raise fastapi.HTTPException( @@ -405,11 +417,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): self.elastic_last_error = None self.enable_metrics = server_args.enable_metrics self.incremental_streaming_output = server_args.incremental_streaming_output - self.enable_lora = server_args.enable_lora + self.enable_lora = get_lora().enable_lora self.enable_trace = server_args.enable_trace self.allow_auto_truncate = server_args.allow_auto_truncate self.skip_tokenizer_init = server_args.skip_tokenizer_init - self.preferred_sampling_params = server_args.preferred_sampling_params + self.preferred_sampling_params = get_serving().preferred_sampling_params self.crash_dump_folder = server_args.crash_dump_folder # Init model config @@ -501,7 +513,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): self.tokenizer = None else: self.tokenizer = get_tokenizer( - server_args.tokenizer_path, + get_serving().tokenizer_path, tokenizer_mode=server_args.tokenizer_mode, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, @@ -522,7 +534,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): self.async_dynamic_batch_tokenizer = None def _validate_cuda_vmm_feature_transport_support(self) -> None: - if self.server_args.mm_feature_transport != "cuda_vmm": + if get_mm().mm_feature_transport != "cuda_vmm": return from sglang.srt.model_loader.utils import get_model_architecture @@ -633,7 +645,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # The registry dynamically updates as adapters are loaded / unloaded during runtime. It # serves as the source of truth for available adapters and maps user-friendly LoRA names # to internally used unique LoRA IDs. - self.lora_registry = LoRARegistry(self.server_args.lora_paths) + self.lora_registry = LoRARegistry(get_lora().lora_paths) # Lock to serialize LoRA update operations. # Please note that, unlike `model_update_lock`, this does not block inference, allowing # LoRA updates and inference to overlap. @@ -642,15 +654,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # point to their latest LoRARef objects, so that they can be # dynamically loaded if needed for inference self.lora_ref_cache: Dict[str, LoRARef] = {} - if self.server_args.lora_paths is not None: - for lora_ref in self.server_args.lora_paths: + if get_lora().lora_paths is not None: + for lora_ref in get_lora().lora_paths: self.lora_ref_cache[lora_ref.lora_name] = lora_ref def init_disaggregation(self, *, start_pd_bootstrap_service: bool = True): # PD Disaggregation - self.disaggregation_mode = DisaggregationMode( - self.server_args.disaggregation_mode - ) + self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode) # Keep a reference so the bootstrap server is not garbage-collected. self.bootstrap_server = ( start_disagg_service(self.server_args) @@ -688,7 +698,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Metrics if self.enable_metrics: engine_type = DisaggregationMode.to_engine_type( - self.server_args.disaggregation_mode + get_disagg().disaggregation_mode ) labels = { @@ -721,7 +731,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): configure_gc_warning(self.server_args.gc_warning_threshold_secs) self.soft_watchdog = Watchdog.create( debug_name="TokenizerManager", - watchdog_timeout=self.server_args.soft_watchdog_timeout, + watchdog_timeout=get_device().soft_watchdog_timeout, soft=True, test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(), ) @@ -997,7 +1007,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): isinstance(obj, EmbeddingReqInput) and obj.is_cross_encoder_request ) if obj.input_embeds is not None: - if not self.server_args.disable_radix_cache: + if not get_memory().disable_radix_cache: raise ValueError( "input_embeds is provided while disable_radix_cache is False. " "Please add `--disable-radix-cache` when you launch the server " @@ -1062,7 +1072,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): if ( not self.server_args.language_only - or self.server_args.encoder_transfer_backend == "zmq_to_tokenizer" + or get_disagg().encoder_transfer_backend == "zmq_to_tokenizer" ): if self.server_args.language_only: mm_inputs = await self.mm_receiver.recv_mm_data( @@ -1071,11 +1081,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): prompt=mm_processor_input, need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs, ) - _reject_missing_dispatched_encoder_embedding( - self.server_args, obj, mm_inputs - ) + _reject_missing_dispatched_encoder_embedding(obj, mm_inputs) if mm_inputs is None: - if self.server_args.language_only: + if get_disagg().language_only: logger.warning( "Encoder embedding not available, " "falling back to local mm processing" @@ -1089,7 +1097,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): ) elif ( self.server_args.language_only - and self.server_args.encoder_transfer_backend + and get_disagg().encoder_transfer_backend in ["zmq_to_scheduler", "mooncake"] and not obj.need_wait_for_mm_inputs ): @@ -1256,12 +1264,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): requested_hidden_mode = get_request_return_hidden_states_mode( obj.return_hidden_states ) - server_hidden_mode = get_server_return_hidden_states_mode(self.server_args) + server_hidden_mode = get_server_return_hidden_states_mode() if requested_hidden_mode > server_hidden_mode: if server_hidden_mode.need_capture(): raise ValueError( "The requested return_hidden_states mode exceeds the " - f"server maximum `{self.server_args.return_hidden_states_mode}`. " + f"server maximum `{get_exec().features.return_hidden_states_mode}`. " "Please launch with `--return-hidden-states-mode full` " "to allow return_hidden_states=True." ) @@ -1283,10 +1291,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): def _validate_mm_limits( self, obj: Union[GenerateReqInput, EmbeddingReqInput] ) -> None: - if not self.server_args.limit_mm_data_per_request: + if not get_mm().limit_mm_data_per_request: return - for modality, limit in self.server_args.limit_mm_data_per_request.items(): + for modality, limit in get_mm().limit_mm_data_per_request.items(): data = getattr(obj, f"{modality}_data", None) if data: count = len(data) if isinstance(data, list) else 1 @@ -1399,7 +1407,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): bootstrap_room = obj.bootstrap_room if ( bootstrap_room is None - and self.server_args.disaggregation_transfer_backend == "fake" + and get_disagg().disaggregation_transfer_backend == "fake" ): bootstrap_room = self.fake_bootstrap_room_counter self.fake_bootstrap_room_counter += 1 @@ -1581,7 +1589,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): - Batch tokenization does not support DP attention yet, and it will make everything goes to the first rank currently """ return batch_size > 0 and ( - self.server_args.enable_tokenizer_batch_encode + get_serving().enable_tokenizer_batch_encode or ( (not get_parallel().enable_dp_attention) and (not self._batch_has_text(batch_size, requests)) @@ -2464,7 +2472,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): state.time_stats.set_finished_time() meta_info["e2e_latency"] = state.time_stats.get_e2e_latency() - if self.server_args.speculative_algorithm: + if get_spec().speculative_algorithm: self._calculate_spec_decoding_metrics(meta_info, recv_obj, i) if self.enable_metrics: scheduler_time_stats = ( @@ -2787,7 +2795,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): ): # Total number of proposed draft tokens per request. num_proposed_drafts = recv_obj.spec_verify_ct[i] * ( - self.server_args.speculative_num_draft_tokens - 1 + get_spec().speculative_num_draft_tokens - 1 ) num_correct_drafts = recv_obj.spec_num_correct_drafts[i] @@ -3484,7 +3492,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # This flag will be used in _tokenize_one_request to determine processing path if should_dispatch: obj.need_wait_for_mm_inputs = True - if self.server_args.encoder_transfer_backend in [ + if get_disagg().encoder_transfer_backend in [ "zmq_to_scheduler", "mooncake", ]: @@ -3603,13 +3611,13 @@ async def print_exception_wrapper(func): def get_processor_wrapper(server_args): return get_processor( - server_args.tokenizer_path, - tokenizer_mode=server_args.tokenizer_mode, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.revision, + get_serving().tokenizer_path, + tokenizer_mode=get_serving().tokenizer_mode, + trust_remote_code=get_model().trust_remote_code, + revision=get_model().revision, image_processor_backend=resolve_image_processor_backend(server_args), - tokenizer_backend=server_args.tokenizer_backend, - model_name=server_args.model_path, + tokenizer_backend=get_serving().tokenizer_backend, + model_name=get_model().model_path, ) diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index ab14d8f0b..707109322 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -567,7 +567,7 @@ class CPUGraphRunner: self.return_hidden_states_mode = ( CaptureHiddenMode.NULL if model_runner.is_draft_worker - else get_server_return_hidden_states_mode(model_runner.server_args) + else get_server_return_hidden_states_mode() ) self.enable_return_hidden_states = self.return_hidden_states_mode.need_capture() # bs -> compiled fn (text-only / skip_cross_attention=True) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 9bf587240..0212e809a 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -32,7 +32,7 @@ import warnings from dataclasses import dataclass from enum import IntEnum, auto from functools import total_ordering -from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Set, Tuple, Union +from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Set, Tuple, Union import torch @@ -231,11 +231,12 @@ def register_attn_tp_sequence_sharded_predicate( _attn_tp_sequence_sharded_predicate = predicate -def get_server_return_hidden_states_mode(server_args: Any) -> CaptureHiddenMode: - mode = getattr(server_args, "return_hidden_states_mode", None) +def get_server_return_hidden_states_mode() -> CaptureHiddenMode: + features = get_exec().features + mode = features.return_hidden_states_mode if mode == "last": return CaptureHiddenMode.LAST - if mode == "full" or getattr(server_args, "enable_return_hidden_states", False): + if mode == "full" or features.enable_return_hidden_states: return CaptureHiddenMode.FULL return CaptureHiddenMode.NULL @@ -719,7 +720,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): if model_runner.is_draft_worker else max( batch.return_hidden_states_mode, - get_server_return_hidden_states_mode(model_runner.server_args), + get_server_return_hidden_states_mode(), ) ) capture_hidden_mode = get_required_capture_hidden_mode( diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index a7a69572f..017198c76 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -726,7 +726,6 @@ class ModelRunner: self._token_oracle_manager = None return self._token_oracle_manager = install_token_oracle_from_env( - server_args=self.server_args, vocab_size=self.model_config.vocab_size, ) diff --git a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py index a41a320f9..0fdd4e5ef 100644 --- a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py @@ -261,8 +261,7 @@ def capture_prefill_graph( if ( model_runner.spec_algorithm.is_eagle() and not model_runner.is_draft_worker - and get_server_return_hidden_states_mode(model_runner.server_args) - < CaptureHiddenMode.FULL + and get_server_return_hidden_states_mode() < CaptureHiddenMode.FULL and not check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE) ): logger.info( diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 17b277162..20426de03 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -219,7 +219,7 @@ class BaseRunner(ABC): self.return_hidden_states_mode = ( CaptureHiddenMode.NULL if model_runner.is_draft_worker - else get_server_return_hidden_states_mode(model_runner.server_args) + else get_server_return_hidden_states_mode() ) self.enable_return_hidden_states = self.return_hidden_states_mode.need_capture() self.attn_tp_size = get_parallel().attn_tp_size @@ -403,7 +403,7 @@ class BaseRunner(ABC): capture_hidden_mode = ( CaptureHiddenMode.NULL if mr.is_draft_worker - else get_server_return_hidden_states_mode(mr.server_args) + else get_server_return_hidden_states_mode() ) num_tokens_per_req = 1 # A PD prefill target worker's pool has no SpeculativeState, so a diff --git a/python/sglang/srt/models/sarvam_moe.py b/python/sglang/srt/models/sarvam_moe.py index e54dd6e79..63373a9a2 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -67,7 +67,6 @@ from sglang.srt.runtime_context import ( get_memory, get_model, get_parallel, - get_server_args, get_stream, ) from sglang.srt.utils import ( @@ -159,12 +158,13 @@ for backend in CONCAT_ROPE_BACKENDS: AttentionBackendRegistry.register(backend, _handle_concat_rope_backend) -def get_attn_forward_method(server_args, forward_batch) -> AttnForwardMethod: +def get_attn_forward_method(forward_batch) -> AttnForwardMethod: + prefill_backend, decode_backend = attention_backends() is_decode = forward_batch.forward_mode.is_decode_or_idle() if is_decode: - backend = server_args.decode_attention_backend or server_args.attention_backend + backend = decode_backend else: - backend = server_args.prefill_attention_backend or server_args.attention_backend + backend = prefill_backend if ( forward_batch.forward_mode.is_extend_without_speculative() and backend == "fa3" @@ -456,7 +456,6 @@ class SarvamMoEMLAAttention(nn.Module): self.max_position_embeddings = max_position_embeddings self.kv_cache_dtype = get_model().kv_cache_dtype - self._server_args = None self.current_attention_backend = None if self.q_lora_rank is None: @@ -761,11 +760,9 @@ class SarvamMoEMLAAttention(nn.Module): q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) k_pe = latent_cache[..., self.kv_lora_rank :].unsqueeze(1) - if self._server_args is None: - self._server_args = get_server_args() self._set_current_attention_backend(forward_batch) - forward_method = get_attn_forward_method(self._server_args, forward_batch) + forward_method = get_attn_forward_method(forward_batch) if forward_method == AttnForwardMethod.MHA_PREFILL: return self._run_mha_prefill( @@ -875,10 +872,8 @@ class SarvamMoEMLAAttention(nn.Module): q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) k_pe = latent_cache[..., self.kv_lora_rank :].unsqueeze(1) - if self._server_args is None: - self._server_args = get_server_args() self._set_current_attention_backend(forward_batch) - forward_method = get_attn_forward_method(self._server_args, forward_batch) + forward_method = get_attn_forward_method(forward_batch) if forward_method == AttnForwardMethod.MHA_PREFILL: output = self._run_mha_prefill( @@ -933,11 +928,9 @@ class SarvamMoEMLAAttention(nn.Module): q_nope_out, k_nope, q_pe, k_pe, forward_batch, zero_allocator = inner_state - if self._server_args is None: - self._server_args = get_server_args() self._set_current_attention_backend(forward_batch) - forward_method = get_attn_forward_method(self._server_args, forward_batch) + forward_method = get_attn_forward_method(forward_batch) if forward_method == AttnForwardMethod.MLA_SEPARATE_ROPE: attn_output = self.attn_mqa( diff --git a/python/sglang/srt/multimodal/cache/identity.py b/python/sglang/srt/multimodal/cache/identity.py index dfa04256c..f7a61d19b 100644 --- a/python/sglang/srt/multimodal/cache/identity.py +++ b/python/sglang/srt/multimodal/cache/identity.py @@ -9,7 +9,7 @@ import struct from dataclasses import dataclass from enum import Enum from pathlib import Path -from typing import TYPE_CHECKING, Any, Mapping, Optional, Protocol, runtime_checkable +from typing import Any, Mapping, Optional, Protocol, runtime_checkable from urllib.parse import unquote, urlparse import numpy as np @@ -17,8 +17,7 @@ import torch import transformers from PIL import Image -if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs +from sglang.srt.runtime_context import get_mm, get_model CONTENT_HASH_PREFIX = "sha256:" _SHA256_HEX_LENGTH = 64 @@ -379,11 +378,17 @@ def resolve_multimodal_item_hash( def build_processor_fingerprint( processor: Any, hf_config: Any, - server_args: ServerArgs, *, extra: Optional[Mapping[str, Any]] = None, ) -> str: - """Fingerprint preprocessing choices that can change processor output.""" + """Fingerprint preprocessing choices that can change processor output. + + Every config value comes from the published bags, which is the only source + that answers the *effective* preprocessing config. Taking any of them from + a handed ``ServerArgs`` would let two callers with the same effective + config disagree on the digest -- and an omitted one silently fingerprint + the empty config, which is how incompatible artifacts get reused. + """ processor_payload = ( processor.preprocess_fingerprint_payload() if isinstance(processor, PreprocessFingerprintProvider) @@ -395,10 +400,10 @@ def build_processor_fingerprint( "processor_class": f"{type(processor).__module__}.{type(processor).__qualname__}", "model_type": hf_payload.get("model_type"), "architectures": hf_payload.get("architectures"), - "model_revision": server_args.revision, - "processor_revision": server_args.revision, - "disable_fast_image_processor": server_args.disable_fast_image_processor, - "mm_process_config": server_args.mm_process_config or {}, + "model_revision": get_model().revision, + "processor_revision": get_model().revision, + "disable_fast_image_processor": get_mm().disable_fast_image_processor, + "mm_process_config": get_mm().mm_process_config or {}, "processor": processor_payload, "extra": extra or {}, } diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 123ea51fa..5b097f40d 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -39,6 +39,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import ( MmItemMemoryPool, get_mm_feature_pool_size_per_worker, ) +from sglang.srt.runtime_context import get_mm from sglang.srt.utils import ( CLIENT_MEDIA_EXCEPTIONS, configure_media_url_security, @@ -212,10 +213,10 @@ class BaseMultimodalProcessor(ABC): self.server_args = server_args self.transport_mode = transport_mode configure_media_url_security( - server_args.allowed_media_domains, + get_mm().allowed_media_domains, server_args.media_url_max_file_size_mb, ) - configured_mm_feature_transport = server_args.mm_feature_transport + configured_mm_feature_transport = get_mm().mm_feature_transport self.mm_feature_transport = ( configured_mm_feature_transport if configured_mm_feature_transport in ("cpu", "cuda_ipc", "cuda_vmm") @@ -231,7 +232,7 @@ class BaseMultimodalProcessor(ABC): self.disable_fast_image_processor = self.image_processor_backend == "pil" self.skip_tokenizer_init = server_args.skip_tokenizer_init - mm_process_config = self.server_args.mm_process_config + mm_process_config = get_mm().mm_process_config self.image_config = mm_process_config.get("image", {}) self.video_config = mm_process_config.get("video", {}) self.audio_config = mm_process_config.get("audio", {}) @@ -255,7 +256,7 @@ class BaseMultimodalProcessor(ABC): # The fingerprint is needed only to build artifact keys. Avoid inspecting # processor state when this processor will never retain artifacts. self.processor_fingerprint = ( - build_processor_fingerprint(self, hf_config, server_args) + build_processor_fingerprint(self, hf_config) if self.mm_preprocess_cache.enabled else None ) diff --git a/python/sglang/srt/multimodal/processors/mimo_v2.py b/python/sglang/srt/multimodal/processors/mimo_v2.py index cd9f5fc47..022bb41cc 100644 --- a/python/sglang/srt/multimodal/processors/mimo_v2.py +++ b/python/sglang/srt/multimodal/processors/mimo_v2.py @@ -39,6 +39,7 @@ from sglang.srt.multimodal.processors.mimo_audio import ( MiMoAudioPipeline, ) from sglang.srt.multimodal.processors.qwen_vl import smart_nframes +from sglang.srt.runtime_context import get_device from sglang.srt.utils import ImageData, VideoData from sglang.srt.utils.common import download_remote_media from sglang.utils import logger @@ -1588,7 +1589,7 @@ class MiMoV2Processor(BaseMultimodalProcessor): processor_config, "video_end_token_id" ) self.use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get() - device = server_args.device if self.use_image_processor_gpu else None + device = get_device().device if self.use_image_processor_gpu else None self.mimo_processor = MiMoProcessor( tokenizer=self._processor.tokenizer, diff --git a/python/sglang/srt/ray/data_parallel_controller.py b/python/sglang/srt/ray/data_parallel_controller.py index 138f0e6fd..1d3d87b27 100644 --- a/python/sglang/srt/ray/data_parallel_controller.py +++ b/python/sglang/srt/ray/data_parallel_controller.py @@ -31,7 +31,11 @@ from sglang.srt.ray.engine import ( _get_bundle_node_ip, _resolve_bundle_indices, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import ( + configured_attn_cp_size, + configured_pp_size, + 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 @@ -144,7 +148,10 @@ class RayDataParallelController(DataParallelController): for node_idx in range(nnodes): bundle_idx = self.bundle_for_node[node_idx] pp_range, tp_range, pp_per_node, tp_per_node = _calculate_rank_ranges( - nnodes, server_args.pp_size, server_args.tp_size, node_rank=node_idx + nnodes, + configured_pp_size(), + server_args.tp_size, + node_rank=node_idx, ) for pp_rank in pp_range: for tp_rank in tp_range: @@ -161,7 +168,7 @@ class RayDataParallelController(DataParallelController): tp_rank, server_args.tp_size, get_parallel().dp_size, - server_args.attn_cp_size, + configured_attn_cp_size(), ) rank_port_args = PortArgs.init_new( server_args, actual_dp_rank, worker_ports @@ -202,7 +209,7 @@ class RayDataParallelController(DataParallelController): world_size = _compute_world_size(server_args) bundle_indices = _resolve_bundle_indices(self.pg, world_size) - ranks_per_tp_group = server_args.tp_size * server_args.pp_size + ranks_per_tp_group = server_args.tp_size * configured_pp_size() if dp_rank is not None: start_rank = dp_rank * ranks_per_tp_group end_rank = start_rank + ranks_per_tp_group @@ -232,7 +239,7 @@ class RayDataParallelController(DataParallelController): tp_rank, server_args.tp_size, get_parallel().dp_size, - server_args.attn_cp_size, + configured_attn_cp_size(), ) rank_port_args = PortArgs.init_new( server_args, actual_dp_rank, worker_ports diff --git a/python/sglang/srt/ray/engine.py b/python/sglang/srt/ray/engine.py index f50523efa..be245fadd 100644 --- a/python/sglang/srt/ray/engine.py +++ b/python/sglang/srt/ray/engine.py @@ -32,7 +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.runtime_context import configured_pp_size, get_parallel from sglang.srt.server_args import PortArgs, ServerArgs logger = logging.getLogger(__name__) @@ -109,8 +109,8 @@ def _compute_world_size(server_args: ServerArgs) -> int: Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size. """ if get_parallel().enable_dp_attention: - return server_args.tp_size * server_args.pp_size - return get_parallel().dp_size * server_args.tp_size * server_args.pp_size + return server_args.tp_size * configured_pp_size() + return get_parallel().dp_size * server_args.tp_size * configured_pp_size() def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]: @@ -269,10 +269,10 @@ class RayEngine(Engine): ) if get_parallel().enable_dp_attention: - total_gpus = server_args.tp_size * server_args.pp_size + total_gpus = server_args.tp_size * configured_pp_size() else: total_gpus = ( - get_parallel().dp_size * server_args.tp_size * server_args.pp_size + get_parallel().dp_size * server_args.tp_size * configured_pp_size() ) nnodes = server_args.nnodes @@ -332,7 +332,7 @@ class RayEngine(Engine): pp_range, tp_range, pp_per_node, tp_per_node = ( _calculate_rank_ranges( nnodes, - server_args.pp_size, + configured_pp_size(), server_args.tp_size, node_rank=node_idx, ) @@ -449,16 +449,16 @@ class RayEngine(Engine): 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 + total_gpus = server_args.tp_size * configured_pp_size() else: total_gpus = ( - get_parallel().dp_size * server_args.tp_size * server_args.pp_size + get_parallel().dp_size * server_args.tp_size * configured_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={get_parallel().dp_size}, " - f"tp_size={server_args.tp_size}, pp_size={server_args.pp_size}, " + f"tp_size={server_args.tp_size}, pp_size={configured_pp_size()}, " f"enable_dp_attention={get_parallel().enable_dp_attention}" ) diff --git a/test/registered/mock_model/test_self_unit_install.py b/test/registered/mock_model/test_self_unit_install.py index 0bcc35feb..45f52056e 100644 --- a/test/registered/mock_model/test_self_unit_install.py +++ b/test/registered/mock_model/test_self_unit_install.py @@ -2,13 +2,13 @@ from __future__ import annotations import os import unittest -from types import SimpleNamespace os.environ["SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE"] = "1" from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env from sglang.srt.kv_canary.token_oracle.oracle import HashOracle from sglang.srt.layers.sampler import _CUSTOM_SAMPLER_FACTORIES +from sglang.srt.runtime_context import get_context from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.test_utils import CustomTestCase @@ -16,23 +16,26 @@ register_cuda_ci(est_time=60, stage="extra-a", runner_config="1-gpu-small") register_amd_ci(est_time=60, suite="extra-a-test-1-gpu-small-amd") -def _make_server_args(*, sampling_backend: str) -> SimpleNamespace: - return SimpleNamespace(sampling_backend=sampling_backend) +def _publish(case, *, sampling_backend: str) -> None: + """The gate reads the published config, so the test publishes one.""" + override = get_context().override_server_args(sampling_backend=sampling_backend) + override.install() + case.addCleanup(override.restore) class TestInstallTokenOracleFromEnv(CustomTestCase): def test_install_token_oracle_from_env_disabled_returns_none(self) -> None: """Verify server-arg-disabled token oracle installation (sampling_backend != 'token_oracle') returns no TokenOracleManager.""" - server_args = _make_server_args(sampling_backend="auto") - hook = install_token_oracle_from_env(server_args=server_args, vocab_size=1000) + _publish(self, sampling_backend="auto") + hook = install_token_oracle_from_env(vocab_size=1000) self.assertIsNone(hook) def test_install_token_oracle_from_env_enabled_registers_oracle_backend( self, ) -> None: """Verify token oracle installation via sampling_backend='token_oracle' registers the oracle backend.""" - server_args = _make_server_args(sampling_backend="token_oracle") - hook = install_token_oracle_from_env(server_args=server_args, vocab_size=512) + _publish(self, sampling_backend="token_oracle") + hook = install_token_oracle_from_env(vocab_size=512) self.assertIsNotNone(hook) self.assertIn("token_oracle", _CUSTOM_SAMPLER_FACTORIES) @@ -40,8 +43,8 @@ class TestInstallTokenOracleFromEnv(CustomTestCase): self, ) -> None: """Verify token oracle installation via sampling_backend='token_oracle' returns a TokenOracleManager wrapping a HashOracle.""" - server_args = _make_server_args(sampling_backend="token_oracle") - hook = install_token_oracle_from_env(server_args=server_args, vocab_size=256) + _publish(self, sampling_backend="token_oracle") + hook = install_token_oracle_from_env(vocab_size=256) self.assertIsNotNone(hook) self.assertIsInstance(hook.oracle, HashOracle) self.assertEqual(hook.oracle.vocab_size, 256) diff --git a/test/registered/unit/constrained/test_base_grammar_backend.py b/test/registered/unit/constrained/test_base_grammar_backend.py index 13568c2cb..52c4f8d29 100644 --- a/test/registered/unit/constrained/test_base_grammar_backend.py +++ b/test/registered/unit/constrained/test_base_grammar_backend.py @@ -28,6 +28,7 @@ from sglang.srt.constrained.base_grammar_backend import ( create_grammar_backend, register_grammar_backend, ) +from sglang.srt.runtime_context import get_context # noqa: E402 from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(2.0, "base-a-test-cpu") @@ -231,18 +232,39 @@ class TestCreateGrammarBackend(unittest.TestCase): GRAMMAR_BACKEND_REGISTRY.clear() GRAMMAR_BACKEND_REGISTRY.update(self._saved) + def _publish(self, **fields): + """Set the config the factory reads. + + The factory takes every config value off the published bags, so a test + that sets one on the handed object would be setting something the + factory does not read -- which is how a mismatch between the two used + to stay invisible here. + """ + override = get_context().override_server_args(**fields) + override.install() + self.addCleanup(override.restore) + def _make_server_args( - self, backend="none", reasoning_parser=None, enable_strict_thinking=False + self, + backend="none", + reasoning_parser=None, + enable_strict_thinking=False, + **fields ): + published = { + "grammar_backend": backend, + "reasoning_parser": reasoning_parser, + "enable_strict_thinking": enable_strict_thinking, + "constrained_json_whitespace_pattern": None, + "constrained_json_disable_any_whitespace": False, + } + published.update(fields) + self._publish(**published) + # Handed on to plugin-registered backends; not a config source. args = MagicMock() args.override = lambda source, **updates: [ setattr(args, key, value) for key, value in updates.items() ] - args.grammar_backend = backend - args.reasoning_parser = reasoning_parser - args.enable_strict_thinking = enable_strict_thinking - args.constrained_json_whitespace_pattern = None - args.constrained_json_disable_any_whitespace = False return args def test_none_backend_returns_none(self): @@ -293,8 +315,9 @@ class TestCreateGrammarBackend(unittest.TestCase): def test_outlines_backend(self, mock_outlines_cls): mock_backend = MagicMock(spec=BaseGrammarBackend) mock_outlines_cls.return_value = mock_backend - args = self._make_server_args("outlines") - args.constrained_json_whitespace_pattern = r"\s*" + args = self._make_server_args( + "outlines", constrained_json_whitespace_pattern=r"\s*" + ) result = create_grammar_backend(args, "tok", 32000) mock_outlines_cls.assert_called_once_with("tok", whitespace_pattern=r"\s*") @@ -304,8 +327,9 @@ class TestCreateGrammarBackend(unittest.TestCase): def test_xgrammar_backend(self, mock_xgrammar_cls): mock_backend = MagicMock(spec=BaseGrammarBackend) mock_xgrammar_cls.return_value = mock_backend - args = self._make_server_args("xgrammar") - args.constrained_json_disable_any_whitespace = True + args = self._make_server_args( + "xgrammar", constrained_json_disable_any_whitespace=True + ) result = create_grammar_backend(args, "tok", 32000, {1, 2}) mock_xgrammar_cls.assert_called_once_with( @@ -336,9 +360,11 @@ class TestCreateGrammarBackend(unittest.TestCase): def test_llguidance_backend(self, mock_guidance_cls): mock_backend = MagicMock(spec=BaseGrammarBackend) mock_guidance_cls.return_value = mock_backend - args = self._make_server_args("llguidance") - args.constrained_json_disable_any_whitespace = False - args.constrained_json_whitespace_pattern = r"\s+" + args = self._make_server_args( + "llguidance", + constrained_json_disable_any_whitespace=False, + constrained_json_whitespace_pattern=r"\s+", + ) result = create_grammar_backend(args, "tok", 32000, {1, 2}) mock_guidance_cls.assert_called_once_with( diff --git a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py index 1a467b079..022b92c4c 100644 --- a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py +++ b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py @@ -42,7 +42,7 @@ from sglang.srt.multimodal.kimi_k3_image_processing import ( materialize_kimi_k3_cpu_features, prepare_kimi_k3_encoder_inputs, ) -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 resolve_encoder_transfer_backend from sglang.srt.utils import ImageData from sglang.test.ci.ci_register import register_cpu_ci @@ -78,29 +78,32 @@ def test_kimi_k3_encoder_transfer_backend_auto_avoids_tp_fanout(): def test_epd_language_only_rejects_missing_dispatched_embedding(): - server_args = SimpleNamespace( + override = get_context().override_server_args( language_only=True, encoder_transfer_backend="zmq_to_tokenizer", ) - request = SimpleNamespace(need_wait_for_mm_inputs=True) + override.install() + try: + request = SimpleNamespace(need_wait_for_mm_inputs=True) - with pytest.raises(HTTPException) as exc_info: - _reject_missing_dispatched_encoder_embedding(server_args, request, None) + with pytest.raises(HTTPException) as exc_info: + _reject_missing_dispatched_encoder_embedding(request, None) - assert getattr(exc_info.value, "status_code", None) == 503 + assert getattr(exc_info.value, "status_code", None) == 503 + finally: + override.restore() def test_epd_rejection_reads_the_resolved_transfer_backend(): - """Tripwire for step 12: this guard fires on the *resolved* backend. + """This guard fires on the *resolved* backend. The record is produced by actual resolution -- a language-only Kimi-K3 launch at TP2, whose `encoder_transfer_backend` starts at the argument default `"auto"` (`ENCODER_TRANSFER_BACKEND_CHOICES[0]`) and is filled in - by `resolve_encoder_transfer_backend` to `"zmq_to_tokenizer"`. Today the - guard therefore rejects. When step 12 makes the instance raw, this same - launch hands the guard a record still at `"auto"`, the rejection silently - stops, and *this test fails* -- which is the signal to give this reader - the resolved value (per-engine overlay or bag) rather than the record. + by `resolve_encoder_transfer_backend` to `"zmq_to_tokenizer"`. The guard + reads that resolved value out of the published bags, so the rejection + survives the record going raw: what a reader must never do is go back to + the record for this field. Fixed doubles cannot trip on that change, so the record here must come from resolution, not a SimpleNamespace. """ @@ -188,20 +191,30 @@ def test_epd_rejection_reads_the_resolved_transfer_backend(): shutil.rmtree(config_dir, ignore_errors=True) assert resolved.encoder_transfer_backend == "zmq_to_tokenizer" - request = SimpleNamespace(need_wait_for_mm_inputs=True) - with pytest.raises(HTTPException) as exc_info: - _reject_missing_dispatched_encoder_embedding(resolved, request, None) - assert getattr(exc_info.value, "status_code", None) == 503 + # Publish that record: the guard reads the resolved value out of the bags, + # so a raw record does not silently disable the rejection. + publish(resolved, role="tokenizer") + try: + request = SimpleNamespace(need_wait_for_mm_inputs=True) + with pytest.raises(HTTPException) as exc_info: + _reject_missing_dispatched_encoder_embedding(request, None) + assert getattr(exc_info.value, "status_code", None) == 503 + finally: + reset_context() def test_epd_allows_local_processing_when_request_was_not_dispatched(): - server_args = SimpleNamespace( + override = get_context().override_server_args( language_only=True, encoder_transfer_backend="zmq_to_tokenizer", ) - request = SimpleNamespace(need_wait_for_mm_inputs=False) + override.install() + try: + request = SimpleNamespace(need_wait_for_mm_inputs=False) - _reject_missing_dispatched_encoder_embedding(server_args, request, None) + _reject_missing_dispatched_encoder_embedding(request, None) + finally: + override.restore() def _encoder(model_type="kimi_k3"): diff --git a/test/registered/unit/managers/test_batch_result_processor_hidden_states.py b/test/registered/unit/managers/test_batch_result_processor_hidden_states.py index ceb843287..c96081bbd 100644 --- a/test/registered/unit/managers/test_batch_result_processor_hidden_states.py +++ b/test/registered/unit/managers/test_batch_result_processor_hidden_states.py @@ -8,13 +8,21 @@ from sglang.srt.managers.scheduler_components.batch_result_processor import ( SchedulerBatchResultProcessor, ) from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode +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 register_cpu_ci(est_time=2, suite="base-a-test-cpu") -def _make_processor(server_mode: str = "full") -> SchedulerBatchResultProcessor: +def _make_processor(case, server_mode: str = "full") -> SchedulerBatchResultProcessor: + # The server-side hidden-state ceiling is a bag leaf. + override = get_context().override_server_args( + enable_return_hidden_states=True, + return_hidden_states_mode=server_mode, + ) + override.install() + case.addCleanup(override.restore) metrics_reporter = Mock() metrics_reporter.num_generated_tokens = 0 metrics_reporter.forward_ct_decode = 0 @@ -26,8 +34,6 @@ def _make_processor(server_mode: str = "full") -> SchedulerBatchResultProcessor: server_args=SimpleNamespace( enable_metrics=False, enable_hisparse=False, - enable_return_hidden_states=True, - return_hidden_states_mode=server_mode, ), model_config=SimpleNamespace(think_end_ids=None), token_to_kv_pool_allocator=Mock(), @@ -139,7 +145,7 @@ class TestPrefillHiddenStateOffsets(CustomTestCase): can_run_cuda_graph=False, skipped_output_comm=False, ) - processor = _make_processor(server_mode) + processor = _make_processor(self, server_mode) with ( patch( @@ -160,7 +166,7 @@ class TestPrefillHiddenStateOffsets(CustomTestCase): class TestDecodeHiddenStateRetention(CustomTestCase): def test_last_mode_multi_step_storage_stays_bounded(self): - processor = _make_processor() + processor = _make_processor(self) req = _DecodeReq() batch = SimpleNamespace( reqs=[req], diff --git a/test/registered/unit/managers/test_hidden_state_server_mode.py b/test/registered/unit/managers/test_hidden_state_server_mode.py index d5730d2b3..267478bf1 100644 --- a/test/registered/unit/managers/test_hidden_state_server_mode.py +++ b/test/registered/unit/managers/test_hidden_state_server_mode.py @@ -4,6 +4,7 @@ from unittest.mock import Mock from sglang.srt.managers.io_struct import GenerateReqInput from sglang.srt.managers.tokenizer_manager import TokenizerManager +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 @@ -11,19 +12,21 @@ register_cpu_ci(est_time=1, suite="base-a-test-cpu") class TestHiddenStateServerMode(CustomTestCase): - @staticmethod - def _make_tokenizer_manager(mode): + def _make_tokenizer_manager(self, mode): + # The server-side hidden-state mode is a bag leaf. + override = get_context().override_server_args( + enable_return_hidden_states=mode is not None, + return_hidden_states_mode=mode, + ) + override.install() + self.addCleanup(override.restore) manager = TokenizerManager.__new__(TokenizerManager) manager.context_len = 128 manager.num_reserved_tokens = 0 manager.allow_auto_truncate = False manager.validate_total_tokens = False manager.is_generation = True - manager.server_args = SimpleNamespace( - enable_return_hidden_states=mode is not None, - return_hidden_states_mode=mode, - enable_custom_logit_processor=False, - ) + manager.server_args = SimpleNamespace(enable_custom_logit_processor=False) manager._validate_token_ids_logprob = Mock() return manager diff --git a/test/registered/unit/managers/test_mm_process_config.py b/test/registered/unit/managers/test_mm_process_config.py index 3922a481f..04a9e82f1 100644 --- a/test/registered/unit/managers/test_mm_process_config.py +++ b/test/registered/unit/managers/test_mm_process_config.py @@ -7,6 +7,7 @@ from unittest.mock import MagicMock, patch import torch from sglang.srt.environ import envs +from sglang.srt.runtime_context import get_context from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -80,14 +81,20 @@ class TestBaseProcessorConfigExtraction(CustomTestCase): BaseMultimodalProcessor, ) + # The multimodal config comes from the bags. + override = get_context().override_server_args( + mm_process_config=mm_process_config, + allowed_media_domains=[], + ) + override.install() + self.addCleanup(override.restore) + server_args = MagicMock() - server_args.mm_process_config = mm_process_config server_args.mm_processor_worker_num = mm_processor_worker_num server_args.mm_io_worker_num = mm_io_worker_num server_args.mm_preprocess_cache_size_mb = None server_args.tokenizer_worker_num = 1 server_args.trust_mm_content_hashes = False - server_args.allowed_media_domains = [] server_args.media_url_max_file_size_mb = 64 hf_config = MagicMock() @@ -170,8 +177,14 @@ class TestBaseProcessorConfigExtraction(CustomTestCase): class TestMultimodalFeatureTransportRuntime(CustomTestCase): - @staticmethod - def _server_args(mm_feature_transport): + def _server_args(self, mm_feature_transport): + override = get_context().override_server_args( + mm_feature_transport=mm_feature_transport, + mm_process_config={}, + allowed_media_domains=[], + ) + override.install() + self.addCleanup(override.restore) return SimpleNamespace( mm_feature_transport=mm_feature_transport, image_processor_backend="auto", @@ -197,8 +210,7 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase): return processor def test_cuda_ipc_pool_uses_resolved_server_arg(self): - # The processor module can be imported before this instance is built; - # transport policy must still resolve from the instance's ServerArgs. + # Transport policy resolves from the mm bag, so the test publishes it. from sglang.srt.multimodal.processors import base_processor with ( @@ -792,16 +804,21 @@ class TestDoubleBosGuard(CustomTestCase): BaseMultimodalProcessor, ) + override = get_context().override_server_args( + mm_process_config={}, + mm_feature_transport="cpu", + allowed_media_domains=[], + ) + override.install() + self.addCleanup(override.restore) + server_args = MagicMock() - server_args.mm_process_config = {} server_args.mm_processor_worker_num = 0 server_args.mm_io_worker_num = 0 - server_args.mm_feature_transport = "cpu" server_args.disable_fast_image_processor = True server_args.mm_preprocess_cache_size_mb = None server_args.tokenizer_worker_num = 1 server_args.trust_mm_content_hashes = False - server_args.allowed_media_domains = [] server_args.media_url_max_file_size_mb = 64 mock_hf_processor = MagicMock() diff --git a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py index 15f89dbe7..b44f172e8 100644 --- a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py +++ b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py @@ -35,6 +35,7 @@ from sglang.srt.managers.tokenizer_manager import ( # noqa: E402 from sglang.srt.observability.req_time_stats import ( # noqa: E402 APIServerReqTimeStats, ) +from sglang.srt.runtime_context import get_context register_cpu_ci(est_time=15, suite="base-a-test-cpu") @@ -103,8 +104,15 @@ _PER_REQUEST_OPTIONAL_FIELDS = frozenset( ) -def _make_tokenizer_manager() -> TokenizerManager: - """Create a TokenizerManager with mocked dependencies, bypassing __init__.""" +def _make_tokenizer_manager(case) -> TokenizerManager: + """Create a TokenizerManager with mocked dependencies, bypassing __init__. + + The config it reads comes from the bags, so the stand-in needs a published + config rather than attributes on a mock. + """ + override = get_context().override_server_args(speculative_algorithm=None) + override.install() + case.addCleanup(override.restore) tm = TokenizerManager.__new__(TokenizerManager) tm.server_args = MagicMock() tm._config_updates = [] @@ -212,7 +220,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase): def test_abort_removes_rid_from_state(self): """After _handle_abort_req, rid should be removed from rid_to_state.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "abort_test_rid" state = _make_req_state(rid) tm.rid_to_state[rid] = state @@ -224,7 +232,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase): def test_abort_allows_resubmit_same_rid(self): """After abort, _init_req_state should accept the same rid again.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "resubmit_after_abort_rid" state = _make_req_state(rid) tm.rid_to_state[rid] = state @@ -245,7 +253,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase): def test_abort_sets_finished_and_notifies(self): """_handle_abort_req should mark state as finished and set the event.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "abort_notify_rid" state = _make_req_state(rid) tm.rid_to_state[rid] = state @@ -266,7 +274,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase): def test_batch_output_removes_rid_on_finish(self): """When a request finishes in _handle_batch_output, rid should be removed.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "batch_finish_rid" state = _make_req_state(rid) tm.rid_to_state[rid] = state @@ -278,7 +286,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase): def test_batch_output_allows_resubmit_after_finish(self): """After a request finishes, the same rid can be resubmitted.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "batch_resubmit_rid" state = _make_req_state(rid) tm.rid_to_state[rid] = state @@ -299,7 +307,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase): def test_batch_output_keeps_rid_when_not_finished(self): """When a request is not yet finished, rid should remain in rid_to_state.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "batch_ongoing_rid" state = _make_req_state(rid) tm.rid_to_state[rid] = state @@ -316,7 +324,7 @@ class TestInitReqStateDuplicateDetection(CustomTestCase): def test_duplicate_rid_raises_error(self): """_init_req_state should raise ValueError if rid already exists.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "duplicate_rid" state = _make_req_state(rid) tm.rid_to_state[rid] = state @@ -334,7 +342,7 @@ class TestInitReqStateDuplicateDetection(CustomTestCase): def test_unique_rid_succeeds(self): """_init_req_state should succeed with a unique rid.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "unique_rid" obj = Mock(spec=GenerateReqInput) @@ -353,7 +361,7 @@ class TestResubmitAfterCompletion(CustomTestCase): def test_complete_then_resubmit_same_rid(self): """A request that completes normally should allow resubmission with the same rid.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "complete_resubmit_rid" # Phase 1: simulate a request in rid_to_state, then complete it @@ -379,7 +387,7 @@ class TestResubmitAfterCompletion(CustomTestCase): def test_abort_then_resubmit_same_rid(self): """An aborted request should allow resubmission with the same rid.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "abort_resubmit_rid" # Phase 1: simulate a request, then abort it @@ -413,9 +421,9 @@ class _DummyAsyncCM: return False -def _make_tm_for_generate() -> TokenizerManager: +def _make_tm_for_generate(case) -> TokenizerManager: """Augment the mocked TokenizerManager with what generate_request needs.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(case) tm.server_args.language_only = False tm.server_args.tokenizer_worker_num = 1 tm.server_args.enable_strict_thinking = False @@ -450,7 +458,7 @@ class TestDiscardPendingReqStates(CustomTestCase): """Direct tests for _discard_pending_req_states.""" def test_discard_single(self): - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rid = "d_single" tm.rid_to_state[rid] = _make_req_state(rid) obj = Mock(spec=GenerateReqInput) @@ -460,7 +468,7 @@ class TestDiscardPendingReqStates(CustomTestCase): self.assertNotIn(rid, tm.rid_to_state) def test_discard_batch_removes_all(self): - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) rids = ["d0", "d1", "d2"] for r in rids: tm.rid_to_state[r] = _make_req_state(r) @@ -473,7 +481,7 @@ class TestDiscardPendingReqStates(CustomTestCase): def test_discard_ignores_already_removed(self): """Popping a rid that is no longer present must not raise.""" - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) tm.rid_to_state["p1"] = _make_req_state("p1") obj = Mock(spec=GenerateReqInput) obj.is_single = False @@ -484,7 +492,7 @@ class TestDiscardPendingReqStates(CustomTestCase): class TestParallelStreamTaskCleanup(CustomTestCase): def test_failing_choice_cancels_and_closes_sibling_waiters(self): - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) async def drive(): sibling_closed = asyncio.Event() @@ -512,7 +520,7 @@ class TestParallelStreamTaskCleanup(CustomTestCase): asyncio.run(drive()) def test_failing_non_stream_choice_cancels_and_closes_sibling_waiters(self): - tm = _make_tokenizer_manager() + tm = _make_tokenizer_manager(self) async def drive(): sibling_closed = asyncio.Event() @@ -546,7 +554,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase): """ def test_single_failure_before_dispatch_cleans_up(self): - tm = _make_tm_for_generate() + tm = _make_tm_for_generate(self) rid = "single_overlen" obj = _make_generate_obj(rid, is_single=True) # Simulate over-length rejection during tokenization/validation. @@ -566,7 +574,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase): self.assertNotIn(rid, tm.rid_to_state) def test_batch_failure_before_dispatch_cleans_up_all(self): - tm = _make_tm_for_generate() + tm = _make_tm_for_generate(self) rids = ["b0", "b1", "b2"] obj = _make_generate_obj(list(rids), is_single=False) @@ -588,7 +596,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase): self.assertNotIn(r, tm.rid_to_state) def test_thinking_budget_rejects_runtime_without_strict_thinking(self): - tm = _make_tm_for_generate() + tm = _make_tm_for_generate(self) obj = GenerateReqInput( text="hello", rid="thinking-budget", diff --git a/test/registered/unit/model_executor/runner/test_hidden_state_graph_recapture.py b/test/registered/unit/model_executor/runner/test_hidden_state_graph_recapture.py index 70cdcde0a..eb65396bc 100644 --- a/test/registered/unit/model_executor/runner/test_hidden_state_graph_recapture.py +++ b/test/registered/unit/model_executor/runner/test_hidden_state_graph_recapture.py @@ -15,6 +15,7 @@ from sglang.srt.model_executor.runner.decode_cuda_graph_runner import ( from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import ( PrefillCudaGraphRunner, ) +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 @@ -23,31 +24,29 @@ register_cpu_ci(est_time=1, suite="base-a-test-cpu") class TestHiddenStateGraphRecapture(CustomTestCase): def test_server_mode_sets_graph_capture_ceiling(self): - disabled = SimpleNamespace( - enable_return_hidden_states=False, - return_hidden_states_mode=None, - ) - last = SimpleNamespace( - enable_return_hidden_states=True, - return_hidden_states_mode="last", - ) - full = SimpleNamespace( - enable_return_hidden_states=True, - return_hidden_states_mode="full", - ) - - self.assertEqual( - get_server_return_hidden_states_mode(disabled), - CaptureHiddenMode.NULL, - ) - self.assertEqual( - get_server_return_hidden_states_mode(last), - CaptureHiddenMode.LAST, - ) - self.assertEqual( - get_server_return_hidden_states_mode(full), - CaptureHiddenMode.FULL, + cases = ( + (dict(enable_return_hidden_states=False), CaptureHiddenMode.NULL), + ( + dict( + enable_return_hidden_states=True, return_hidden_states_mode="last" + ), + CaptureHiddenMode.LAST, + ), + ( + dict( + enable_return_hidden_states=True, return_hidden_states_mode="full" + ), + CaptureHiddenMode.FULL, + ), ) + for fields, expected in cases: + with self.subTest(**fields): + override = get_context().override_server_args(**fields) + override.install() + try: + self.assertEqual(get_server_return_hidden_states_mode(), expected) + finally: + override.restore() @staticmethod def _make_runner(runner_cls, capture_hidden_mode): diff --git a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py index 81baa61b8..a00bb5c9b 100644 --- a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py +++ b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py @@ -17,6 +17,7 @@ from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import ( PrefillCudaGraphRunner, ) from sglang.srt.model_executor.runner.shape_key import ShapeKey +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 @@ -111,13 +112,17 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase): def test_eagle_target_tc_piecewise_skips_last_mode_capture(self): eager_runner = object() + # The server-side hidden-state ceiling is a bag leaf. + override = get_context().override_server_args( + enable_return_hidden_states=True, + return_hidden_states_mode="last", + ) + override.install() + self.addCleanup(override.restore) model_runner = SimpleNamespace( is_draft_worker=False, spec_algorithm=SimpleNamespace(is_eagle=lambda: True), - server_args=SimpleNamespace( - enable_return_hidden_states=True, - return_hidden_states_mode="last", - ), + server_args=SimpleNamespace(), ) with patch.object( diff --git a/test/registered/unit/models/test_kimi_k25.py b/test/registered/unit/models/test_kimi_k25.py index cf3f27ea6..88a7e19a6 100644 --- a/test/registered/unit/models/test_kimi_k25.py +++ b/test/registered/unit/models/test_kimi_k25.py @@ -66,7 +66,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import ( DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY, CudaIpcTensorTransportProxy, ) -from sglang.srt.runtime_context import get_context, get_parallel +from sglang.srt.runtime_context import get_context, get_parallel, publish, reset_context from sglang.srt.server_args import ServerArgs from sglang.srt.utils import ImageData from sglang.test.ci.ci_register import register_cpu_ci @@ -642,11 +642,9 @@ def _k3_preprocess_config( ) def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls): server_args = SimpleNamespace( - mm_feature_transport="cpu", image_processor_backend="auto", disable_fast_image_processor=False, skip_tokenizer_init=False, - mm_process_config={}, mm_io_worker_num=0, mm_processor_worker_num=0, tokenizer_worker_num=1, @@ -654,32 +652,35 @@ def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls trust_mm_content_hashes=False, base_gpu_id=0, rl_on_policy_target=None, - allowed_media_domains=[], media_url_max_file_size_mb=64, ) - processor = processor_cls( - hf_config=SimpleNamespace(media_placeholder_token_id=42), - server_args=server_args, - _processor=_HFProcessor(), - transport_mode=None, - ) - try: - worker_processor = asyncio.run( - processor.mm_processor_executor.run(lambda *, processor: processor) + # The multimodal config comes from the bags. + with get_context().override_server_args( + mm_feature_transport="cpu", mm_process_config={}, allowed_media_domains=[] + ): + processor = processor_cls( + hf_config=SimpleNamespace(media_placeholder_token_id=42), + server_args=server_args, + _processor=_HFProcessor(), + transport_mode=None, ) - assert isinstance(processor._processor, wrapper_cls) - assert isinstance(worker_processor, wrapper_cls) - assert worker_processor is not processor._processor - if processor_cls is KimiK3ImageProcessor: - fingerprint_config = processor.preprocess_fingerprint_payload()[ - "wrapped_processor" - ] - assert isinstance(fingerprint_config, KimiK3PreprocessConfig) - assert fingerprint_config.patch_size == 14 - finally: - processor.mm_processor_executor.shutdown() - processor.io_executor.shutdown() - processor.cpu_executor.shutdown() + try: + worker_processor = asyncio.run( + processor.mm_processor_executor.run(lambda *, processor: processor) + ) + assert isinstance(processor._processor, wrapper_cls) + assert isinstance(worker_processor, wrapper_cls) + assert worker_processor is not processor._processor + if processor_cls is KimiK3ImageProcessor: + fingerprint_config = processor.preprocess_fingerprint_payload()[ + "wrapped_processor" + ] + assert isinstance(fingerprint_config, KimiK3PreprocessConfig) + assert fingerprint_config.patch_size == 14 + finally: + processor.mm_processor_executor.shutdown() + processor.io_executor.shutdown() + processor.cpu_executor.shutdown() def test_kimi_k3_expands_image_placeholders_with_original_dimensions(): @@ -844,82 +845,88 @@ def test_kimi_k3_normal_cache_path_connects_real_producer_to_model_consumer(): tokenizer_worker_num=1, mm_preprocess_cache_size_mb=1, ) - processor = KimiK3ImageProcessor( - hf_config=hf_config, - server_args=server_args, - _processor=hf_processor, - transport_mode=None, - ) - image = Image.new("RGB", (28, 28), color=(1, 2, 3)) - encoded_image = io.BytesIO() - image.save(encoded_image, format="PNG") - image_data = ImageData( - url="data:image/png;base64," - + base64.b64encode(encoded_image.getvalue()).decode() - ) - request = SimpleNamespace(video_data=None, mm_content_hashes=None) - - class _Tower(nn.Module): - device = torch.device("cpu") - patch_size = 14 - - def __init__(self): - super().__init__() - self.patch_embed = SimpleNamespace( - proj=SimpleNamespace(weight=torch.empty(1, dtype=torch.float32)) - ) - - def forward(self, pixel_values, _grid_thws): - return pixel_values - - model = KimiK3ForConditionalGeneration.__new__(KimiK3ForConditionalGeneration) - nn.Module.__init__(model) - model.use_data_parallel = False - model.vision_tower = _Tower() - model.mm_projector = _Projector() - + publish(server_args, role="tokenizer") try: - with ( - patch( - "sglang.srt.multimodal.processors.kimi_k3.is_cuda", return_value=True - ), - patch.object( - processor, - "prepare_artifact_batch", - wraps=processor.prepare_artifact_batch, - ) as prepare_artifacts, - ): - cold = asyncio.run( - processor.process_mm_data_async([image_data], [1, 42, 2], request) - ) - hot = asyncio.run( - processor.process_mm_data_async([image_data], [3, 42, 4], request) - ) - cold_items = pickle.loads(pickle.dumps(cold.mm_items)) - hot_items = pickle.loads(pickle.dumps(hot.mm_items)) + processor = KimiK3ImageProcessor( + hf_config=hf_config, + server_args=server_args, + _processor=hf_processor, + transport_mode=None, + ) + image = Image.new("RGB", (28, 28), color=(1, 2, 3)) + encoded_image = io.BytesIO() + image.save(encoded_image, format="PNG") + image_data = ImageData( + url="data:image/png;base64," + + base64.b64encode(encoded_image.getvalue()).decode() + ) + request = SimpleNamespace(video_data=None, mm_content_hashes=None) - with ( - patch("sglang.srt.models.kimi_k3.configured_tp_size", return_value=1), - patch( - "sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images", - return_value=( - torch.ones((4, 3), dtype=torch.float32), - torch.tensor([[1, 2, 2]], dtype=torch.int64), + class _Tower(nn.Module): + device = torch.device("cpu") + patch_size = 14 + + def __init__(self): + super().__init__() + self.patch_embed = SimpleNamespace( + proj=SimpleNamespace(weight=torch.empty(1, dtype=torch.float32)) + ) + + def forward(self, pixel_values, _grid_thws): + return pixel_values + + model = KimiK3ForConditionalGeneration.__new__(KimiK3ForConditionalGeneration) + nn.Module.__init__(model) + model.use_data_parallel = False + model.vision_tower = _Tower() + model.mm_projector = _Projector() + + try: + with ( + patch( + "sglang.srt.multimodal.processors.kimi_k3.is_cuda", + return_value=True, ), - ), - ): - cold_features = model.get_image_feature(cold_items) - hot_features = model.get_image_feature(hot_items) - finally: - processor.shutdown() + patch.object( + processor, + "prepare_artifact_batch", + wraps=processor.prepare_artifact_batch, + ) as prepare_artifacts, + ): + cold = asyncio.run( + processor.process_mm_data_async([image_data], [1, 42, 2], request) + ) + hot = asyncio.run( + processor.process_mm_data_async([image_data], [3, 42, 4], request) + ) + cold_items = pickle.loads(pickle.dumps(cold.mm_items)) + hot_items = pickle.loads(pickle.dumps(hot.mm_items)) - assert prepare_artifacts.call_count == 1 - assert cold.mm_items[0].hash == hot.mm_items[0].hash - assert cold.mm_items[0].offsets == hot.mm_items[0].offsets == [(3, 3)] - assert ( - cold_items[0].model_specific_data[DEFERRED_PREPROCESSING_KEY].backend == "gpu" - ) - torch.testing.assert_close(cold_features, hot_features) + with ( + patch("sglang.srt.models.kimi_k3.configured_tp_size", return_value=1), + patch( + "sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images", + return_value=( + torch.ones((4, 3), dtype=torch.float32), + torch.tensor([[1, 2, 2]], dtype=torch.int64), + ), + ), + ): + cold_features = model.get_image_feature(cold_items) + hot_features = model.get_image_feature(hot_items) + finally: + processor.shutdown() + + assert prepare_artifacts.call_count == 1 + assert cold.mm_items[0].hash == hot.mm_items[0].hash + assert cold.mm_items[0].offsets == hot.mm_items[0].offsets == [(3, 3)] + assert ( + cold_items[0].model_specific_data[DEFERRED_PREPROCESSING_KEY].backend + == "gpu" + ) + torch.testing.assert_close(cold_features, hot_features) + finally: + reset_context() def test_kimi_k3_model_accepts_mixed_cached_eager_and_deferred_artifacts(): diff --git a/test/registered/unit/multimodal/rust/qwen/_fixtures.py b/test/registered/unit/multimodal/rust/qwen/_fixtures.py index 76e22fe7d..28534a362 100644 --- a/test/registered/unit/multimodal/rust/qwen/_fixtures.py +++ b/test/registered/unit/multimodal/rust/qwen/_fixtures.py @@ -21,11 +21,13 @@ maybe_stub_sgl_kernel() from sglang.srt.multimodal.processors.qwen_vl import ( # noqa: E402 QwenVLImageProcessor, ) +from sglang.srt.runtime_context import publish, reset_context # noqa: E402 +from sglang.srt.server_args import ServerArgs # noqa: E402 register_cpu_ci(est_time=0, suite="base-a-test-cpu", disabled="Qwen test fixtures") -def make_processor(config, image_processor_cls=None): +def make_processor(case, config, image_processor_cls=None): """A ``QwenVLImageProcessor`` over a tiny hand-built tokenizer. ``image_processor_cls`` picks the HF backend; they resample differently.""" image_processor_cls = image_processor_cls or HfQwenImageProcessor @@ -87,6 +89,18 @@ def make_processor(config, image_processor_cls=None): allowed_media_domains=[], media_url_max_file_size_mb=64, ) + # The processor reads its media policy, transport and per-modality limits + # from the mm bag, so the fixture publishes before building it. + publish( + ServerArgs( + model_path="dummy", + mm_feature_transport=server_args.mm_feature_transport, + mm_process_config=server_args.mm_process_config, + allowed_media_domains=server_args.allowed_media_domains, + ), + role="tokenizer", + ) + case.addCleanup(reset_context) return QwenVLImageProcessor( hf_config, server_args, processor, None, skip_mm_pool=True ) diff --git a/test/registered/unit/multimodal/rust/qwen/test_e2e_parity.py b/test/registered/unit/multimodal/rust/qwen/test_e2e_parity.py index 1f2f542a5..75cec9d64 100644 --- a/test/registered/unit/multimodal/rust/qwen/test_e2e_parity.py +++ b/test/registered/unit/multimodal/rust/qwen/test_e2e_parity.py @@ -50,7 +50,9 @@ class TestQwenE2eParity(CustomTestCase): import transformers.models.qwen2_vl as qwen2_vl self.processor = make_processor( - PROCESSOR_CONFIGS["qwen2_5_vl"], getattr(qwen2_vl, self.image_processor) + self, + PROCESSOR_CONFIGS["qwen2_5_vl"], + getattr(qwen2_vl, self.image_processor), ) def tearDown(self): diff --git a/test/registered/unit/multimodal/rust/qwen/test_native_mm_host.py b/test/registered/unit/multimodal/rust/qwen/test_native_mm_host.py index f53618b84..e7c078576 100644 --- a/test/registered/unit/multimodal/rust/qwen/test_native_mm_host.py +++ b/test/registered/unit/multimodal/rust/qwen/test_native_mm_host.py @@ -50,7 +50,7 @@ class TestQwenNativeMmHashes(CustomTestCase): from sglang.srt.managers.multimodal_processor import import_processors import_processors("sglang.srt.multimodal.processors") - self.processor = make_processor(PROCESSOR_CONFIGS["qwen2_5_vl"]) + self.processor = make_processor(self, PROCESSOR_CONFIGS["qwen2_5_vl"]) def tearDown(self): self.processor.io_executor.shutdown() diff --git a/test/registered/unit/multimodal/test_gpu_feature_transport.py b/test/registered/unit/multimodal/test_gpu_feature_transport.py index d104b27fe..2c5aab74a 100644 --- a/test/registered/unit/multimodal/test_gpu_feature_transport.py +++ b/test/registered/unit/multimodal/test_gpu_feature_transport.py @@ -42,6 +42,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase): def test_model_class_controls_cuda_vmm_opt_in(self): from sglang.srt.managers.tokenizer_manager import TokenizerManager + from sglang.srt.runtime_context import get_context class SupportedModel: supports_cuda_vmm_feature_transport = True @@ -49,8 +50,10 @@ class TestCudaVmmFeatureTransport(unittest.TestCase): class UnsupportedModel: pass + override = get_context().override_server_args(mm_feature_transport="cuda_vmm") + override.install() + self.addCleanup(override.restore) manager = object.__new__(TokenizerManager) - manager.server_args = SimpleNamespace(mm_feature_transport="cuda_vmm") manager.model_config = object() with patch( @@ -70,9 +73,12 @@ class TestCudaVmmFeatureTransport(unittest.TestCase): def test_cpu_transport_skips_model_opt_in_lookup(self): from sglang.srt.managers.tokenizer_manager import TokenizerManager + from sglang.srt.runtime_context import get_context + override = get_context().override_server_args(mm_feature_transport="cpu") + override.install() + self.addCleanup(override.restore) manager = object.__new__(TokenizerManager) - manager.server_args = SimpleNamespace(mm_feature_transport="cpu") manager.model_config = object() with patch( diff --git a/test/registered/unit/multimodal/test_preprocess_cache.py b/test/registered/unit/multimodal/test_preprocess_cache.py index 6db0e5481..e6ea8298f 100644 --- a/test/registered/unit/multimodal/test_preprocess_cache.py +++ b/test/registered/unit/multimodal/test_preprocess_cache.py @@ -21,6 +21,7 @@ from sglang.srt.multimodal.cache import ( resolve_multimodal_item_hash, snapshot_media, ) +from sglang.srt.runtime_context import publish, reset_context from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci @@ -216,26 +217,32 @@ class TestMediaIdentity(unittest.TestCase): return {"model_type": "vlm", "architectures": ["VLM"]} config = Config() - args = ServerArgs( - model_path="dummy", - revision="model-revision", - disable_fast_image_processor=False, - mm_process_config={"image": {"max_pixels": 1024}}, - ) - base = build_processor_fingerprint(Processor("gpu"), config, args) - changed_backend = build_processor_fingerprint(Processor("cpu"), config, args) - changed_args = ServerArgs( - model_path="dummy", - revision="model-revision", - disable_fast_image_processor=False, - mm_process_config={"image": {"max_pixels": 2048}}, - ) - changed_config = build_processor_fingerprint( - Processor("gpu"), config, changed_args - ) + def fingerprint(processor, mm_process_config): + # The digest reads the effective config, so the test publishes it + # rather than handing one in: that is the only source the function + # has, and two callers with the same effective config must agree. + publish( + ServerArgs( + model_path="dummy", + revision="model-revision", + disable_fast_image_processor=False, + mm_process_config=mm_process_config, + ), + role="test", + ) + return build_processor_fingerprint(processor, config) + + self.addCleanup(reset_context) + small = {"image": {"max_pixels": 1024}} + base = fingerprint(Processor("gpu"), small) + changed_backend = fingerprint(Processor("cpu"), small) + changed_config = fingerprint(Processor("gpu"), {"image": {"max_pixels": 2048}}) + same_again = fingerprint(Processor("gpu"), small) + self.assertNotEqual(base, changed_backend) self.assertNotEqual(base, changed_config) + self.assertEqual(base, same_again) def test_item_hash_namespace_covers_identity_and_processor_output(self): digest = snapshot_media(b"image").content_digest diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 5d462fe83..c3b2719b5 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -2161,6 +2161,7 @@ class TestGrpcServerArgs(CustomTestCase): tokenizer_manager=MagicMock(), template_manager=MagicMock(), scheduler_info={}, + grpc_port=server_args.grpc_port, ) self.assertEqual(handle, "handle") diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index 2b718ca7a..fdbeb8677 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -52,6 +52,23 @@ _SLOT_OWNERS = ("srt/runtime_context.py", "srt/server_args.py", "srt/arg_groups/ # live topology cannot answer there. The test below asserts this map is exactly # the set of call sites, so the reasons cannot drift away from the code. _CONFIGURED_SIZE_CALL_SITES = { + ("srt/entrypoints/engine.py", "configured_pp_size"): ( + "the launch path decides how many scheduler processes to spawn; it runs " + "before any of them exists, so there is no group to ask" + ), + ("srt/ray/engine.py", "configured_pp_size"): ( + "the Ray driver sizes the actor placement group; the actors it is about " + "to create are the ones that will hold the process groups" + ), + ("srt/ray/data_parallel_controller.py", "configured_pp_size"): ( + "same placement arithmetic on the DP path -- ranks per TP group, " + "computed in the driver before the actors start" + ), + ("srt/ray/data_parallel_controller.py", "configured_attn_cp_size"): ( + "the attention-CP factor of that same placement arithmetic, and the one " + "size whose live value cannot express the configured intent when " + "attn_cp_size > moe_dp_size aliases the groups" + ), ("srt/layers/attention/dsa/dsa_indexer.py", "configured_pp_size"): ( "gates `pp_size > 1 and not get_pp_group()...`; the short circuit is the " "point, since with PP off the group is never touched, which is what lets " @@ -91,6 +108,10 @@ _CONFIGURED_SIZE_CALL_SITES = { "the consumer count is configured fan-out arithmetic (tp_size // " "dp_size), which is what the record answered before" ), + ("srt/disaggregation/encode_server.py", "configured_tp_size"): ( + "the encode server's launch entry sizes its workers before it has " + "spawned any of them" + ), ("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_launch_path_reads_configured_sizes.py b/test/registered/unit/test_launch_path_reads_configured_sizes.py new file mode 100644 index 000000000..3d4fcac91 --- /dev/null +++ b/test/registered/unit/test_launch_path_reads_configured_sizes.py @@ -0,0 +1,280 @@ +"""Launch paths read the configured parallel sizes, not the live ones. + +`get_parallel().pp_size` and its four siblings are read-through properties over +the process groups, so they answer only after distributed init. The launcher +decides how many processes to spawn *before* that, and a live read there raises +`Distributed environment is not initialized` -- a startup crash no unit test +reaches, because nothing short of booting a server runs the launcher. +""" + +import ast +import pathlib +import unittest + +import sglang +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=9, suite="base-a-test-cpu") + +_PACKAGE_ROOT = pathlib.Path(sglang.__file__).resolve().parent + +# Live-shadowed sizes a launch path is known to have read. ParallelContext +# shadows more properties than these (every `_v(name, ...)` one raises the same +# "Distributed environment is not initialized"); this dict carries the ones a +# `configured_*` accessor answers, so it is a remedy map, not a census. +_LIVE_SHADOWED = { + "tp_size": "configured_tp_size()", + "pp_size": "configured_pp_size()", + "moe_dp_size": "configured_moe_dp_size()", + "attn_cp_size": "configured_attn_cp_size()", + "dcp_size": "a configured accessor (none exists yet; add one beside configured_pp_size)", +} + +# Launch paths that decide how many children to spawn are derived below +# from the spawn itself. These launch without a size-driven spawn, so no +# derivation reaches them and they are carried by hand. +_HAND_CARRIED = ( + "srt/entrypoints/http_server.py", + "srt/entrypoints/sidecar.py", + "srt/ray/data_parallel_controller.py", + "srt/ray/engine.py", + "srt/ray/http_server.py", +) + + +def _multiprocessing_names(tree): + """Names bound to multiprocessing, to one of its start contexts, or to the + process constructors themselves.""" + modules, constructors = set(), set() + for node in ast.walk(tree): + if isinstance(node, ast.Import): + for a in node.names: + if a.name == "multiprocessing" or a.name.startswith("multiprocessing."): + modules.add(a.asname or a.name.split(".")[0]) + elif a.name == "torch.multiprocessing": + modules.add(a.asname or "torch") + elif isinstance(node, ast.ImportFrom): + if node.module in ( + "multiprocessing", + "multiprocessing.context", + "torch.multiprocessing", + ): + constructors |= { + a.asname or a.name for a in node.names if a.name == "Process" + } + elif node.module == "concurrent.futures": + constructors |= { + a.asname or a.name + for a in node.names + if a.name == "ProcessPoolExecutor" + } + for node in ast.walk(tree): + if isinstance(node, ast.Assign) and isinstance(node.value, ast.Call): + func = node.value.func + if ( + isinstance(func, ast.Attribute) + and func.attr == "get_context" + and isinstance(func.value, ast.Name) + and func.value.id in modules + ): + modules |= {t.id for t in node.targets if isinstance(t, ast.Name)} + return modules, constructors + + +def _configured_accessors() -> frozenset: + """The `configured_*_size()` names `runtime_context` exports. + + Derived from that module, so a new accessor keeps its launcher watched + without a second list here. + """ + tree = ast.parse( + (_PACKAGE_ROOT / "srt/runtime_context.py").read_text(encoding="utf-8-sig") + ) + names = frozenset( + node.name + for node in tree.body + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + and node.name.startswith("configured_") + and node.name.endswith("_size") + ) + assert names, ( + "no configured_*_size accessors found in runtime_context; the " + "derivation is broken, not the tree" + ) + return names + + +def _spawns_from_a_size(tree) -> bool: + """Does any function here construct a child process *and* read one of the + five sizes -- live off the parallel bag, or through its `configured_*_size()` + answer? That is a spawn count decided from the topology. + + Counting the configured read too is what keeps a launcher watched after it + is converted. Deriving on the live read alone means the file drops out of + the scan the moment it stops offending, so the guard would only ever watch + the launchers that already fail it. + """ + configured = _configured_accessors() + modules, constructors = _multiprocessing_names(tree) + names, qualified = _parallel_bag_names(tree) + aliases = _bag_aliases(tree, names, qualified) + for fn in ast.walk(tree): + if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)): + continue + spawns = reads = False + for node in ast.walk(fn): + if isinstance(node, ast.Call): + func = node.func + if isinstance(func, ast.Attribute) and func.attr in ( + "Process", + "ProcessPoolExecutor", + "Popen", + "spawn", + ): + # `mp.Process(`, `mp.get_context("spawn").Process(` and + # `subprocess.Popen(` all reach a child process; the + # receiver of a chained call is itself a call, so this + # cannot require a bare Name. + spawns = True + elif isinstance(func, ast.Name) and func.id in constructors: + spawns = True + if (isinstance(func, ast.Name) and func.id in configured) or ( + isinstance(func, ast.Attribute) and func.attr in configured + ): + reads = True + elif ( + isinstance(node, ast.Attribute) + and node.attr in _LIVE_SHADOWED + and ( + _is_parallel_bag_call(node.value, names, qualified) + or (isinstance(node.value, ast.Name) and node.value.id in aliases) + ) + ): + # A record read (`server_args.tp_size`) sizes a spawn too, but + # it cannot raise pre-dist; only the bag read is this guard's + # subject, so only it forces a module into _PRE_DIST. + reads = True + if spawns and reads: + return True + return False + + +def _parallel_bag_names(tree): + """What this module calls `get_parallel`, plus any runtime_context alias. + + A literal-name match reads only one spelling; an aliased import or a + module-qualified call is the same read with a different surface. + """ + names, modules = set(), set() + for node in ast.walk(tree): + if ( + isinstance(node, ast.ImportFrom) + and node.module + and node.module.endswith("runtime_context") + ): + names |= { + a.asname or a.name for a in node.names if a.name == "get_parallel" + } + elif isinstance(node, ast.Import): + for a in node.names: + if a.name.endswith("runtime_context"): + modules.add(a.asname or a.name.split(".")[0]) + return names, modules + + +def _is_parallel_bag_call(node, names, modules) -> bool: + if not isinstance(node, ast.Call): + return False + if isinstance(node.func, ast.Name): + return node.func.id in names + return ( + isinstance(node.func, ast.Attribute) + and node.func.attr == "get_parallel" + and isinstance(node.func.value, ast.Name) + and node.func.value.id in modules + ) + + +def _bag_aliases(tree, names, qualified): + """Locals bound to the parallel bag: `p = get_parallel()` then `p.pp_size` + is the same read one line later.""" + return { + target.id + for node in ast.walk(tree) + if isinstance(node, ast.Assign) + and _is_parallel_bag_call(node.value, names, qualified) + for target in node.targets + if isinstance(target, ast.Name) + } + + +def _launch_paths(): + """(relative path, tree) per module that runs before its process groups. + + A module that sizes a spawn loop from a parallel-bag size is derived from + the spawn itself; `_HAND_CARRIED` holds the launch entries that spawn + nothing, which no derivation can reach. + """ + seen = {} + for path in sorted(_PACKAGE_ROOT.rglob("*.py")): + source = path.read_text() + # Every spawn shape below names Process, ProcessPoolExecutor or Popen. + if not any(name in source for name in ("Process", "Popen", "spawn")): + continue + try: + tree = ast.parse(source) + except SyntaxError: + continue + if _spawns_from_a_size(tree): + seen[str(path.relative_to(_PACKAGE_ROOT))] = tree + for rel in _HAND_CARRIED: + seen.setdefault(rel, ast.parse((_PACKAGE_ROOT / rel).read_text())) + return sorted(seen.items()) + + +class TestLaunchPathsReadConfiguredSizes(CustomTestCase): + def test_no_live_topology_read_before_distributed_init(self): + offenders = [] + for rel, tree in _launch_paths(): + names, modules = _parallel_bag_names(tree) + aliases = _bag_aliases(tree, names, modules) + for node in ast.walk(tree): + if isinstance(node, ast.Attribute) and node.attr in _LIVE_SHADOWED: + base = node.value + if _is_parallel_bag_call(base, names, modules) or ( + isinstance(base, ast.Name) and base.id in aliases + ): + offenders.append( + f"{rel}:{node.lineno} reads the live {node.attr}; " + f"use {_LIVE_SHADOWED[node.attr]}" + ) + elif ( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Name) + and node.func.id == "getattr" + and len(node.args) >= 2 + and isinstance(node.args[1], ast.Constant) + and node.args[1].value in _LIVE_SHADOWED + and ( + _is_parallel_bag_call(node.args[0], names, modules) + or ( + isinstance(node.args[0], ast.Name) + and node.args[0].id in aliases + ) + ) + ): + offenders.append( + f"{rel}:{node.lineno} reads the live " + f"{node.args[1].value} through getattr; " + f"use {_LIVE_SHADOWED[node.args[1].value]}" + ) + self.assertEqual( + offenders, + [], + "launch paths run before distributed init:\n " + "\n ".join(offenders), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/test_publish_precedes_bag_reads.py b/test/registered/unit/test_publish_precedes_bag_reads.py index 1dee9c003..a9383f248 100644 --- a/test/registered/unit/test_publish_precedes_bag_reads.py +++ b/test/registered/unit/test_publish_precedes_bag_reads.py @@ -104,9 +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/managers/detokenizer_manager.py", "run_detokenizer_process"): ( - "DetokenizerManager reads the handed instance at this revision" - ), } # `publish` itself and its named wrappers live here; a call inside them is the diff --git a/test/registered/unit/test_split_attention_backend_decisions.py b/test/registered/unit/test_split_attention_backend_decisions.py index 3812d658e..e01129f16 100644 --- a/test/registered/unit/test_split_attention_backend_decisions.py +++ b/test/registered/unit/test_split_attention_backend_decisions.py @@ -49,6 +49,7 @@ _PAIR_READERS = { "batch_overlap/two_batch_overlap.py": "prefill (extend positions)", "managers/scheduler.py": "prefill (truncation align knobs)", "entrypoints/engine.py": "either half (flashinfer version floor)", + "models/sarvam_moe.py": "the half serving the forward (attn dispatch)", } diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index 8426d3940..0a3b39de7 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -129,6 +129,7 @@ _ENV_MATRIX = (({}, {"SGLANG_IS_IN_CI": "true"}),) _PASSED = frozenset({"model_path", "device", "random_seed"}) _EXPOSED = { + ("entrypoints/sidecar.py", "grpc_port"), ("configs/embedding_model_spec.py", "chunked_prefill_size"), ("configs/embedding_model_spec.py", "cuda_graph_config"), ("configs/embedding_model_spec.py", "disable_radix_cache"), @@ -143,24 +144,10 @@ _EXPOSED = { ("configs/model_config.py", "quantization"), ("configs/model_config.py", "speculative_algorithm"), ("configs/model_config.py", "speculative_draft_model_quantization"), - ("constrained/base_grammar_backend.py", "grammar_backend"), - ("constrained/base_grammar_backend.py", "reasoning_parser"), ("disaggregation/common/conn.py", "disaggregation_bootstrap_port"), ("disaggregation/common/conn.py", "pp_size"), ("disaggregation/decode_kvcache_offload_manager.py", "hicache_io_backend"), ("disaggregation/decode_kvcache_offload_manager.py", "served_model_name"), - ("disaggregation/encode_receiver.py", "disaggregation_ib_device"), - ("disaggregation/encode_receiver.py", "encoder_transfer_backend"), - ("disaggregation/encode_receiver.py", "mooncake_ib_device"), - ("disaggregation/encode_receiver.py", "tokenizer_path"), - ("disaggregation/encode_server.py", "allowed_media_domains"), - ("disaggregation/encode_server.py", "device"), - ("disaggregation/encode_server.py", "encoder_transfer_backend"), - ("disaggregation/encode_server.py", "load_format"), - ("disaggregation/encode_server.py", "mm_process_config"), - ("disaggregation/encode_server.py", "model_path"), - ("disaggregation/encode_server.py", "served_model_name"), - ("disaggregation/encode_server.py", "tokenizer_path"), ("disaggregation/utils.py", "disaggregation_transfer_backend"), ("distributed/bootstrap.py", "disable_custom_all_reduce"), ("distributed/bootstrap.py", "enable_symm_mem"), @@ -201,28 +188,14 @@ _EXPOSED = { ("elastic_ep/expert_backup_manager.py", "load_format"), ("elastic_ep/expert_backup_manager.py", "mooncake_ib_device"), ("entrypoints/engine.py", "attn_cp_size"), - ("entrypoints/engine.py", "detokenizer_worker_num"), - ("entrypoints/engine.py", "dtype"), ("entrypoints/engine.py", "enable_symm_mem"), - ("entrypoints/engine.py", "ep_join_mode"), - ("entrypoints/engine.py", "load_format"), - ("entrypoints/engine.py", "model_path"), ("entrypoints/engine.py", "moe_dp_size"), - ("entrypoints/engine.py", "pp_size"), - ("entrypoints/engine.py", "quantization"), ("entrypoints/engine.py", "reasoning_parser"), ( "entrypoints/engine.py", "remote_instance_weight_loader_start_seed_via_transfer_engine", ), ("entrypoints/engine.py", "tool_call_parser"), - ("entrypoints/http_server.py", "disaggregation_mode"), - ("entrypoints/http_server.py", "ep_join_mode"), - ("entrypoints/http_server.py", "grpc_port"), - ("entrypoints/http_server.py", "model_path"), - ("entrypoints/http_server.py", "served_model_name"), - ("entrypoints/http_server.py", "skip_server_warmup"), - ("entrypoints/sidecar.py", "grpc_port"), ("eplb/eplb_manager.py", "ep_dispatch_algorithm"), ("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"), ("eplb/expert_distribution.py", "deepep_mode"), @@ -236,7 +209,6 @@ _EXPOSED = { ("kv_canary/capacities.py", "chunked_prefill_size"), ("kv_canary/capacities.py", "cuda_graph_config"), ("kv_canary/capacities.py", "speculative_num_draft_tokens"), - ("kv_canary/token_oracle/install.py", "sampling_backend"), ("layers/cp/base.py", "attn_cp_size"), ("layers/cp/base.py", "cp_strategy"), ("layers/cp/base.py", "enable_prefill_cp"), @@ -259,9 +231,6 @@ _EXPOSED = { ("managers/data_parallel_controller.py", "moe_dp_size"), ("managers/data_parallel_controller.py", "pp_size"), ("managers/data_parallel_controller.py", "soft_watchdog_timeout"), - ("managers/detokenizer_manager.py", "soft_watchdog_timeout"), - ("managers/detokenizer_manager.py", "tokenizer_path"), - ("managers/detokenizer_manager.py", "tool_call_parser"), ("managers/disagg_service.py", "disaggregation_bootstrap_port"), ("managers/disagg_service.py", "disaggregation_mode"), ("managers/disagg_service.py", "disaggregation_transfer_backend"), @@ -279,27 +248,7 @@ _EXPOSED = { ("managers/scheduler.py", "pp_size"), ("managers/scheduler.py", "soft_watchdog_timeout"), ("managers/scheduler.py", "speculative_algorithm"), - ( - "managers/scheduler_components/new_token_ratio_tracker.py", - "schedule_conservativeness", - ), - ("managers/tokenizer_manager.py", "disable_radix_cache"), - ("managers/tokenizer_manager.py", "disaggregation_mode"), - ("managers/tokenizer_manager.py", "disaggregation_transfer_backend"), - ("managers/tokenizer_manager.py", "enable_lora"), - ("managers/tokenizer_manager.py", "enable_tokenizer_batch_encode"), - ("managers/tokenizer_manager.py", "encoder_transfer_backend"), - ("managers/tokenizer_manager.py", "limit_mm_data_per_request"), - ("managers/tokenizer_manager.py", "lora_paths"), - ("managers/tokenizer_manager.py", "mm_feature_transport"), - ("managers/tokenizer_manager.py", "model_path"), - ("managers/tokenizer_manager.py", "preferred_sampling_params"), - ("managers/tokenizer_manager.py", "return_hidden_states_mode"), ("managers/tokenizer_manager.py", "served_model_name"), - ("managers/tokenizer_manager.py", "soft_watchdog_timeout"), - ("managers/tokenizer_manager.py", "speculative_algorithm"), - ("managers/tokenizer_manager.py", "speculative_num_draft_tokens"), - ("managers/tokenizer_manager.py", "tokenizer_path"), ("managers/tp_worker.py", "disable_overlap_schedule"), ("managers/tp_worker.py", "model_path"), ("managers/tp_worker.py", "random_seed"), @@ -313,8 +262,6 @@ _EXPOSED = { ("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "served_model_name"), ("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"), - ("model_executor/forward_batch_info.py", "return_hidden_states_mode"), ("model_executor/model_runner.py", "device"), ("model_executor/model_runner.py", "speculative_algorithm"), ("model_executor/model_runner.py", "speculative_draft_attention_backend"), @@ -353,22 +300,11 @@ _EXPOSED = { "model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py", "cuda_graph_config", ), - ("models/sarvam_moe.py", "attention_backend"), - ("models/sarvam_moe.py", "decode_attention_backend"), - ("models/sarvam_moe.py", "prefill_attention_backend"), - ("multimodal/cache/identity.py", "mm_process_config"), - ("multimodal/processors/base_processor.py", "allowed_media_domains"), ("multimodal/processors/base_processor.py", "image_processor_backend"), - ("multimodal/processors/base_processor.py", "mm_feature_transport"), - ("multimodal/processors/base_processor.py", "mm_process_config"), - ("multimodal/processors/mimo_v2.py", "device"), ("observability/metrics_collector.py", "disaggregation_mode"), ("observability/metrics_collector.py", "prefill_delayer_max_delay_passes"), ("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", "pp_size"), - ("ray/engine.py", "pp_size"), ("speculative/adaptive_spec_params.py", "speculative_algorithm"), ("speculative/adaptive_spec_params.py", "speculative_eagle_topk"), ("speculative/dflash_worker_v2.py", "speculative_draft_window_size"), @@ -435,29 +371,20 @@ _EXPOSED_CUDA_ONLY: frozenset = frozenset() _OVERRIDDEN_AND_READ = { ("configs/model_config.py", "dtype"), ("configs/model_config.py", "model_path"), - ("constrained/base_grammar_backend.py", "grammar_backend"), ("disaggregation/decode_kvcache_offload_manager.py", "hicache_storage_backend"), ( "disaggregation/decode_kvcache_offload_manager.py", "hicache_storage_backend_extra_config", ), - ("disaggregation/encode_server.py", "load_format"), - ("disaggregation/encode_server.py", "model_path"), ( "distributed/device_communicators/mooncake_transfer_engine.py", "hicache_storage_backend", ), ("dllm/config.py", "model_path"), ("elastic_ep/expert_backup_manager.py", "load_format"), - ("entrypoints/engine.py", "dtype"), - ("entrypoints/engine.py", "load_format"), - ("entrypoints/engine.py", "model_path"), - ("entrypoints/http_server.py", "model_path"), ("kv_canary/api.py", "speculative_num_steps"), ("kv_canary/capacities.py", "speculative_num_draft_tokens"), ("managers/scheduler.py", "hicache_storage_backend"), - ("managers/tokenizer_manager.py", "model_path"), - ("managers/tokenizer_manager.py", "speculative_num_draft_tokens"), ("managers/tp_worker.py", "model_path"), ("mem_cache/hiradix_cache.py", "hicache_storage_backend"), ("mem_cache/hiradix_cache.py", "hicache_storage_backend_extra_config"),