From 58d4baba5c4f59eb20eba581069e6a391f207e03 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 09:45:21 +0000 Subject: [PATCH 01/26] config: whole-object readbacks report the projection, not the fields `/server_info`, its gRPC and in-process twins, and the scheduler's internal-state dump all handed out the record itself: three via `dataclasses.asdict(server_args)` and one via `dict(vars(server_args))`. Both read the fields, which carry resolution's result only for as long as declarations materialize onto the record -- and the point of declaring is that they will stop. Left alone, these endpoints would quietly start reporting what the operator typed instead of what resolution decided. `ServerArgs.resolved_dict()` is the whole-object shape of `resolution_result`: every field, read through the declarations, nested dataclasses expanded the way `asdict` expands them. The four exits report it. Values are unchanged today -- 15 launch shapes agree field for field across all 476 -- so this is the placement, not a new answer. The `vars()` base was also leaking: the resolution bookkeeping (`_raw_input`, the declaration stash, the materialization marker) and the `ModelConfig` memo crossed IPC into `/server_info`'s `internal_states` block. The projection is fields only, and a test pins that the dump is exactly the fields. --- python/sglang/srt/arg_groups/overrides.py | 36 +++++++++++++++++++ python/sglang/srt/entrypoints/engine.py | 2 +- python/sglang/srt/entrypoints/grpc_bridge.py | 3 +- python/sglang/srt/entrypoints/http_server.py | 3 +- .../sglang/srt/managers/tokenizer_manager.py | 2 +- python/sglang/srt/runtime_context.py | 13 +++---- python/sglang/srt/server_args.py | 15 ++++++++ .../unit/entrypoints/test_server_info.py | 8 ++--- .../test_resolution_declarations.py | 26 ++++++++++++++ 9 files changed, 92 insertions(+), 16 deletions(-) diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 2f16545287ca..b0d7838812c2 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -29,6 +29,7 @@ from __future__ import annotations +import copy import dataclasses import inspect import json @@ -380,6 +381,41 @@ def resolution_result(server_args: Any, field: str, default: Any = None) -> Any: return getattr(server_args, field, default) +def resolution_projection(server_args: Any) -> Dict[str, Any]: + """Every field's resolved value, nested dataclasses expanded. + + The whole-object shape of ``resolution_result``, for the exits that hand out + the entire configuration (``/server_info``, the gRPC and engine readbacks). + They used ``dataclasses.asdict``, which reads the fields -- correct only for + as long as declarations materialize onto the record, and the point of + declaring is that they will not. Field values only: the private resolution + bookkeeping and the ``model_config`` memo that a ``vars()`` dump carried into + the readback are not configuration. + """ + return { + field.name: _plain(resolution_result(server_args, field.name)) + for field in dataclasses.fields(server_args) + } + + +def _plain(value: Any) -> Any: + """``dataclasses.asdict``'s conversion, applied to one value: dataclasses + become dicts, containers recurse, everything else is deep-copied (a caller + mutating the dump must not reach the live configuration).""" + if dataclasses.is_dataclass(value) and not isinstance(value, type): + return { + field.name: _plain(getattr(value, field.name)) + for field in dataclasses.fields(value) + } + if isinstance(value, tuple) and hasattr(value, "_fields"): # namedtuple + return type(value)(*(_plain(item) for item in value)) + if isinstance(value, (list, tuple)): + return type(value)(_plain(item) for item in value) + if isinstance(value, dict): + return type(value)((_plain(k), _plain(v)) for k, v in value.items()) + return copy.deepcopy(value) + + def resolved_view(server_args: Any) -> ResolvedView: """Read-only view of the resolving configuration for mid-resolution code that is not a pass (``__post_init__`` handlers and hooks). Internal to diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 118b05f852f8..e4e625c6062d 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -1342,7 +1342,7 @@ def get_server_info(self): ) return msgspec_to_builtins( { - **dataclasses.asdict(self.tokenizer_manager.server_args), + **self.tokenizer_manager.server_args.resolved_dict(), **self._scheduler_init_result.scheduler_infos[0], "startup_time": self.tokenizer_manager.startup_time, "internal_states": internal_states, diff --git a/python/sglang/srt/entrypoints/grpc_bridge.py b/python/sglang/srt/entrypoints/grpc_bridge.py index 56dfc1acfb1f..9349d3f92358 100644 --- a/python/sglang/srt/entrypoints/grpc_bridge.py +++ b/python/sglang/srt/entrypoints/grpc_bridge.py @@ -7,7 +7,6 @@ """ import asyncio -import dataclasses import json import logging from types import SimpleNamespace @@ -423,7 +422,7 @@ def get_model_info(self) -> str: return json.dumps(result, default=str) def get_server_info(self) -> str: - result: Dict[str, Any] = dataclasses.asdict(self.tokenizer_manager.server_args) + result: Dict[str, Any] = self.tokenizer_manager.server_args.resolved_dict() result.update(self.scheduler_info) result["kv_events"] = ( self.tokenizer_manager.server_args.describe_kv_events_publisher() diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index df6d09fc310d..f6006872f7f2 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -811,10 +811,9 @@ async def server_info(): server_args = _global_state.tokenizer_manager.server_args - # server_args.model_config is not serializable but should be excluded by asdict. return msgspec_to_builtins( { - **dataclasses.asdict(server_args), + **server_args.resolved_dict(), **_global_state.scheduler_info, "startup_time": _global_state.tokenizer_manager.startup_time, "internal_states": internal_states, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index c21fe1754c92..1c71ad387a6a 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -2040,7 +2040,7 @@ def _dump_config_snapshot(self) -> Optional[Dict[str, Any]]: part that cannot be reconstructed afterwards. """ try: - return self.resolved_config_dict(dataclasses.asdict(self.server_args)) + return self.resolved_config_dict(self.server_args.resolved_dict()) except Exception as e: logger.error(f"Failed to snapshot the resolved config for the dump: {e!r}") return None diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 20fd371209fa..fa95274503df 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -969,11 +969,12 @@ def resolved_server_args_dict(self, base: dict | None = None) -> dict: in a readback: HiCache attach/detach, the generated forward-pass-metrics endpoint, tunables set via ``/set_internal_state``. - ``base`` defaults to ``dict(vars(server_args))`` (matching the legacy - ``vars`` dump); pass ``dataclasses.asdict(server_args)`` when nested - dataclass fields must be expanded first. Override leaves are flat - ``ServerArgs`` field names, so overlaying them onto the top level of - either base is exact. + ``base`` defaults to ``server_args.resolved_dict()`` -- the record's + fields as resolution decided them, nested dataclasses expanded. (It used + to be ``dict(vars(server_args))``, which carried the private resolution + bookkeeping and the ``model_config`` memo into the readback.) Override + leaves are flat ``ServerArgs`` field names, so overlaying them onto the + top level of the base is exact. The log is per process: it carries what *this* process overrode. A weight reload records ``model_path`` and ``load_format`` from the @@ -984,7 +985,7 @@ def resolved_server_args_dict(self, base: dict | None = None) -> dict: The top-level ``/server_info`` fields are the startup record, not this dump. """ - d = dict(vars(self.server_args)) if base is None else dict(base) + d = self.server_args.resolved_dict() if base is None else dict(base) for _source, fields in self._overrides_log: d.update(fields) return d diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index fb5c69cdff8a..d8e483f7b7e7 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -3729,6 +3729,21 @@ def resolve_once(self) -> None: # handlers ran, not how far they got. self._declarations_materialized = True + def resolved_dict(self) -> Dict[str, Any]: + """This configuration as a plain dict of resolved field values. + + What the whole-object readbacks report (`/server_info` and its gRPC and + in-process twins). `dataclasses.asdict(self)` reads the fields, which + carry resolution's result only while declarations materialize onto the + record; this reads the declarations, so it keeps answering with what + resolution decided once they stop. Nested dataclass fields are expanded + the way `asdict` expands them; the private resolution bookkeeping and the + `model_config` memo are not fields and do not appear. + """ + from sglang.srt.arg_groups.overrides import resolution_projection + + return resolution_projection(self) + def replace_resolved(self, source: str, **changes: Any) -> ServerArgs: """A copy of this record that stays resolved, and says what it changed. diff --git a/test/registered/unit/entrypoints/test_server_info.py b/test/registered/unit/entrypoints/test_server_info.py index 56d72cc1b742..8aa9b327605e 100644 --- a/test/registered/unit/entrypoints/test_server_info.py +++ b/test/registered/unit/entrypoints/test_server_info.py @@ -292,7 +292,7 @@ def test_the_recorded_update_reaches_the_readback_overlay(self): tokenizer_manager.record_config_updates("test", weight_version="v2") self.assertEqual(tokenizer_manager.config_value("weight_version"), "v2") overlaid = tokenizer_manager.resolved_config_dict( - dataclasses.asdict(server_args) + server_args.resolved_dict() ) self.assertEqual(overlaid["weight_version"], "v2") finally: @@ -313,9 +313,9 @@ class TestServerInfoExistingFieldsPreserved(CustomTestCase): """ def test_every_server_args_field_appears_in_response(self): - # `dataclasses.asdict(server_args)` is spread into the response; - # asserting every dataclass field surfaces is the strongest - # backward-compat guarantee that's still implementation-agnostic. + # `server_args.resolved_dict()` is spread into the response; asserting + # every dataclass field surfaces is the strongest backward-compat + # guarantee that's still implementation-agnostic. args = ServerArgs(model_path="dummy") info = _call_server_info_with(args) diff --git a/test/registered/unit/server_args/test_resolution_declarations.py b/test/registered/unit/server_args/test_resolution_declarations.py index 2925811613f4..331e17779fc8 100644 --- a/test/registered/unit/server_args/test_resolution_declarations.py +++ b/test/registered/unit/server_args/test_resolution_declarations.py @@ -376,6 +376,32 @@ def test_the_projection_input_is_the_resolved_configuration(self): + "\n ".join(differences), ) + def test_the_whole_object_readback_carries_only_fields(self): + """`/server_info` and its gRPC and in-process twins report + `ServerArgs.resolved_dict()`. + + The dump is exactly the field names, carrying the resolution result + for each. It holds none of the resolution bookkeeping (`_raw_input`, the + declaration stash, the finished flag) and no `ModelConfig` memo: none of + that is configuration, and all of it would cross IPC with the + readback. + """ + server_args = self._resolve({"tp_size": 2}) + dump = server_args.resolved_dict() + self.assertEqual( + sorted(dump), + sorted(field.name for field in dataclasses.fields(server_args)), + "the readback dump is no longer exactly the fields", + ) + leaked = sorted( + name + for name in vars(server_args) + if name not in dump and not name.startswith("__") + ) + self.assertNotEqual( + leaked, [], "nothing to leak any more -- this check is now vacuous" + ) + def test_every_published_leaf_is_what_resolution_decided(self): """One hop further than the check above: the leaf a reader reads. From 37084f438512d8815e0263d37e2c4156f8b94764 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 10:08:55 +0000 Subject: [PATCH 02/26] config: every process entry publishes its own config Three constructors published defensively because each could be built with nothing published before it: `ModelRunner`, `TokenizerManager`, `MMEncoder`. That made "are the config bags available here?" a question about which constructor happened to have run, which is the wrong place for it: code between the process entry and that constructor cannot read a bag, and a `publish` that lands late re-projects the bags and silently drops any `override()` taken in between. The six entries that relied on it now publish for themselves: `init_multi_tokenizer` (the multi-tokenizer worker reads its record from shared memory), the two `benchmark/one_batch` work functions (run inline for tp_size == 1 and spawned per rank otherwise), the encoder's gRPC entry, and its spawned TP and DP workers. `publish` resolves through the idempotent gate, so a spawned child that receives a resolved record is unaffected. No behavior change: the constructors' `ensure_published` becomes a no-op at each of these, and an override taken between an entry and its constructor now survives instead of being discarded. `test_publish_precedes_bag_reads` pins all six as publishing entries, so the walk checks each one's publish against the bag reads it reaches. --- python/sglang/benchmark/one_batch.py | 7 ++++++- python/sglang/srt/disaggregation/encoder/grpc_server.py | 3 ++- python/sglang/srt/disaggregation/encoder/runtime.py | 2 ++ python/sglang/srt/disaggregation/encoder/server.py | 2 ++ python/sglang/srt/entrypoints/http_server.py | 3 +++ test/registered/unit/test_publish_precedes_bag_reads.py | 9 +++++++++ 6 files changed, 24 insertions(+), 2 deletions(-) diff --git a/python/sglang/benchmark/one_batch.py b/python/sglang/benchmark/one_batch.py index 38addce4b388..ea1863500c3f 100644 --- a/python/sglang/benchmark/one_batch.py +++ b/python/sglang/benchmark/one_batch.py @@ -82,7 +82,7 @@ from sglang.srt.model_executor.cuda_graph_config import Phase, cuda_graph_fully_disabled from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_parallel, get_schedule +from sglang.srt.runtime_context import get_parallel, get_schedule, publish from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.speculative.spec_info import SpeculativeAlgorithm @@ -681,6 +681,8 @@ def correctness_test( gpu_id, tp_rank, ): + publish(server_args, role="scheduler") + # Configure the logger configure_logger(server_args, prefix=f" TP{tp_rank}") rank_print = print if tp_rank == 0 else lambda *args, **kwargs: None @@ -881,6 +883,9 @@ def latency_test( gpu_id, tp_rank, ): + # `main` runs this inline for tp_size == 1 and spawns it per rank otherwise; + # a spawned child arrives with nothing published. + publish(server_args, role="scheduler") initialize_moe_config(server_args) initialize_fp8_gemm_config(server_args) initialize_fp4_gemm_config(server_args) diff --git a/python/sglang/srt/disaggregation/encoder/grpc_server.py b/python/sglang/srt/disaggregation/encoder/grpc_server.py index 163dc276d0b8..91024e8594ef 100644 --- a/python/sglang/srt/disaggregation/encoder/grpc_server.py +++ b/python/sglang/srt/disaggregation/encoder/grpc_server.py @@ -24,7 +24,7 @@ from sglang.srt.disaggregation.encoder.server import MMEncoder, launch_encoder from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle from sglang.srt.managers.schedule_batch import Modality -from sglang.srt.runtime_context import get_disagg +from sglang.srt.runtime_context import get_disagg, publish from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import random_uuid from sglang.srt.utils.network import NetworkAddress, get_zmq_socket @@ -201,6 +201,7 @@ async def SchedulerReceiveUrl( async def serve_grpc_encoder(server_args: ServerArgs): + publish(server_args, role="encoder") ctx = mp.get_context("spawn") zmq_ctx = zmq.asyncio.Context(10) ipc_path_prefix = random_uuid() diff --git a/python/sglang/srt/disaggregation/encoder/runtime.py b/python/sglang/srt/disaggregation/encoder/runtime.py index 330a8b9330b0..433f0e9fc73d 100644 --- a/python/sglang/srt/disaggregation/encoder/runtime.py +++ b/python/sglang/srt/disaggregation/encoder/runtime.py @@ -53,6 +53,7 @@ get_observability, get_parallel, get_serving, + publish, ) from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import configure_logger, random_uuid, set_prometheus_multiproc_dir @@ -1473,6 +1474,7 @@ def launch_dp_worker( dispatch_path: str, result_path: str, ): + publish(server_args, role="encoder") try: configure_logger(server_args, prefix=f" encode_dp_worker[{dp_rank}]") asyncio.run( diff --git a/python/sglang/srt/disaggregation/encoder/server.py b/python/sglang/srt/disaggregation/encoder/server.py index 3d9596d81d58..ad716af57c98 100644 --- a/python/sglang/srt/disaggregation/encoder/server.py +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -63,6 +63,7 @@ get_exec, get_mm, get_model, + publish, ) from sglang.srt.server_args import ServerArgs from sglang.srt.utils import configure_media_url_security @@ -2036,6 +2037,7 @@ async def _handle_encoder_worker_request(encoder: MMEncoder, request): def launch_encoder(server_args, schedule_path, dist_init_method, rank): + publish(server_args, role="encoder") try: asyncio.run(run_encoder(server_args, schedule_path, dist_init_method, rank)) except KeyboardInterrupt: diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index f6006872f7f2..69575931f735 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -232,6 +232,8 @@ async def init_multi_tokenizer() -> ServerArgs: server_args.api_key is None ), "API key is not supported in multi-tokenizer mode" + publish(server_args, role="tokenizer") + # Create a new ipc name for the current process port_args.tokenizer_ipc_name = ( f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}" @@ -488,6 +490,7 @@ async def custom_handler(request: Request): get_model, get_parallel, get_serving, + publish, ) elastic_ep_router.route_class = ORJSONRoute diff --git a/test/registered/unit/test_publish_precedes_bag_reads.py b/test/registered/unit/test_publish_precedes_bag_reads.py index 8429b0915da5..fc5c6b34e535 100644 --- a/test/registered/unit/test_publish_precedes_bag_reads.py +++ b/test/registered/unit/test_publish_precedes_bag_reads.py @@ -87,6 +87,15 @@ "run_expert_backup_manager_process", ), ("srt/weight_cache/daemon.py", "load"), + # The multi-tokenizer worker, the benchmark work functions (run + # inline or spawned per rank), and the encoder's gRPC / spawned-TP / + # spawned-DP entries. + ("srt/entrypoints/http_server.py", "init_multi_tokenizer"), + ("benchmark/one_batch.py", "latency_test"), + ("benchmark/one_batch.py", "correctness_test"), + ("srt/disaggregation/encoder/grpc_server.py", "serve_grpc_encoder"), + ("srt/disaggregation/encoder/server.py", "launch_encoder"), + ("srt/disaggregation/encoder/runtime.py", "launch_dp_worker"), } ) From de39aea7e55cbf914cd53ce2a66f1c44e51c637d Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 10:17:33 +0000 Subject: [PATCH 03/26] config: a constructor checks that the config is published, it does not publish `ModelRunner`, `TokenizerManager` and `MMEncoder` published defensively through `ensure_published`. With every process entry publishing for itself, a constructor arriving unpublished no longer means "this is a standalone build" -- it means an entry was missed, and publishing here would work by accident while re-projecting the bags over whatever the process had. `assert_published` replaces it: same identity-and-role check, and it raises with what is published instead, naming the entry as the place to publish. The draft runner is unchanged (it deliberately does not publish, so it does not check either). The four manual tests that built one of these objects directly now publish in their setup, which is what they are: the process entry for that test. The two constructor entries leave `_KNOWN_ENTRIES`, and the constructor-publisher census is down to the two that really are entries -- `Engine` and `SchedulerActor`. --- .../srt/disaggregation/encoder/server.py | 4 +- .../sglang/srt/managers/tokenizer_manager.py | 4 +- .../sglang/srt/model_executor/model_runner.py | 11 +-- python/sglang/srt/runtime_context.py | 49 ++++++---- python/sglang/test/config_publishers.py | 4 +- test/manual/test_forward_split_prefill.py | 3 + test/manual/test_tokenizer_batch_encode.py | 2 + test/manual/test_tokenizer_manager.py | 5 + test/manual/test_vlm_accuracy.py | 11 ++- .../unit/test_publish_precedes_bag_reads.py | 2 - test/registered/unit/test_runtime_context.py | 92 ++++++++----------- 11 files changed, 95 insertions(+), 92 deletions(-) diff --git a/python/sglang/srt/disaggregation/encoder/server.py b/python/sglang/srt/disaggregation/encoder/server.py index ad716af57c98..cef8e884ba9d 100644 --- a/python/sglang/srt/disaggregation/encoder/server.py +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -57,7 +57,7 @@ ) from sglang.srt.observability.metrics_collector import EncoderMetricsCollector from sglang.srt.runtime_context import ( - ensure_published, + assert_published, get_device, get_disagg, get_exec, @@ -449,7 +449,7 @@ def __init__( ``base_gpu_id + rank`` — the DP launcher's per-worker placement. It is this instance's value, not a config change, so it travels as an argument.""" - ensure_published(server_args, role="encoder") + assert_published(server_args, role="encoder") logger.info(f"init MMEncoder {rank}/{server_args.tp_size}") self.server_args = server_args configure_media_url_security( diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 1c71ad387a6a..221877577e82 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -123,7 +123,7 @@ ) from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers from sglang.srt.runtime_context import ( - ensure_published, + assert_published, get_context, get_device, get_disagg, @@ -409,7 +409,7 @@ def __init__( ): # Parse args self.server_args = server_args - ensure_published(server_args, role="tokenizer") + assert_published(server_args, role="tokenizer") self.startup_time: Optional[Dict[str, Any]] = None self.elastic_worker_count = get_parallel().config.dp_size self.elastic_pending_ep_size = None diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index b990b4914c49..cd3eb88ed685 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -168,7 +168,7 @@ ) from sglang.srt.platforms import current_platform from sglang.srt.runtime_context import ( - ensure_published, + assert_published, get_context, get_device, get_exec, @@ -332,13 +332,10 @@ def __init__( self.dist_port = nccl_port self.server_args = server_args self.is_draft_worker = is_draft_worker - # Set the global server_args in the scheduler process (target worker - # only, so a draft init cannot clobber target-derived global state). - # Before the constructor's bag reads (page_size below): a standalone - # construction (benchmark/one_batch, the manual runner tests) has no - # earlier publish. + # The process entry published; a draft runner is not one (it must not + # clobber the target's config), so only the target checks. if not is_draft_worker: - ensure_published(server_args, role="scheduler") + assert_published(server_args, role="scheduler") # Set by maybe_init_lora_manager; stays None when LoRA is off and on # draft runners, which serve adapters' target model unadapted. self.lora_manager: Optional[LoRAManager] = None diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index fa95274503df..c94b32c6ff5a 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -1375,29 +1375,38 @@ def publish(server_args, *, role: str, hf_config: Any = None) -> RuntimeContext: return _CONTEXT -def ensure_published(server_args, *, role: str) -> RuntimeContext: - """Publish unless this exact record is already published under this role. - - Three constructors publish defensively, because each can be built with - nothing published before it -- `ModelRunner` (a benchmark harness, the - manual runner tests), `TokenizerManager`, and `MMEncoder` (spawned encoder - workers). Inside a process that already published the same record, - publishing again re-projects the bags: every `override()` taken between the - two calls is discarded, and the provenance log with it. - - No override sits in one of those windows today, so this removes a hazard - rather than a live bug. It is worth removing anyway: the drop is silent, it - depends on where a constructor happens to sit relative to the overrides - around it, and `publish` now says what a re-projection discarded so the - next one is loud. - - So these callers ask for the end state -- this record, this role, published - -- and get a no-op when that already holds. An engine rebuild still calls - `publish` directly, because there the reset is the point. +def assert_published(server_args, *, role: str) -> RuntimeContext: + """This record, under this role, is already published -- or fail loud. + + Publishing is the process entry's job: `run_scheduler_process`, + `init_multi_tokenizer`, a spawned encoder worker, the benchmark work + functions. A constructor arriving here unpublished means one of those + entries is missing. + + A `publish` at this point re-projects the bags over a live process, + discarding every `override()` taken since and the provenance log with it, + so this raises. """ if _CONTEXT._server_args is server_args and _CONTEXT._publish_role == role: return _CONTEXT - return publish(server_args, role=role) + if _CONTEXT._server_args is None: + detail = "nothing is published in this process" + elif _CONTEXT._server_args is not server_args: + detail = ( + "a different record is published " + f"(role={_CONTEXT._publish_role!r}); this constructor was handed " + "one the process never published" + ) + else: + detail = ( + f"this record is published under role " + f"{_CONTEXT._publish_role!r}, not {role!r}" + ) + raise RuntimeError( + f"config not published for role {role!r}: {detail}. The process entry " + "publishes -- add publish(server_args, role=...) there rather than " + "publishing from a constructor." + ) def publish_role() -> str | None: diff --git a/python/sglang/test/config_publishers.py b/python/sglang/test/config_publishers.py index f473403a796b..65544200bb8b 100644 --- a/python/sglang/test/config_publishers.py +++ b/python/sglang/test/config_publishers.py @@ -1,8 +1,8 @@ """Who installs the startup record into the runtime context, derived from code. Two guards need this answer and neither should keep its own list: matching the -spellings by hand is how `ensure_published` once read as *not* publishing, -which turned a correct module into a reported violation. A publisher is +spellings by hand is how the constructors' old defensive publish once read as +*not* publishing, which turned a correct module into a reported violation. A publisher is defined by what it does -- it reaches ``RuntimeContext.set_server_args`` -- and a *constructor* publisher is an ``__init__`` that calls one. """ diff --git a/test/manual/test_forward_split_prefill.py b/test/manual/test_forward_split_prefill.py index c54c6456cd97..c020793c1a25 100644 --- a/test/manual/test_forward_split_prefill.py +++ b/test/manual/test_forward_split_prefill.py @@ -19,6 +19,7 @@ from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.model_runner import ModelRunner +from sglang.srt.runtime_context import publish from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.speculative.spec_info import SpeculativeAlgorithm @@ -52,6 +53,8 @@ def setUpClass(cls): cls.port_args = PortArgs.init_new(cls.server_args) + publish(cls.server_args, role="scheduler") + # Load model and tokenizer cls.model_config = ModelConfig.from_server_args(cls.server_args) cls.model_runner = ModelRunner( diff --git a/test/manual/test_tokenizer_batch_encode.py b/test/manual/test_tokenizer_batch_encode.py index 31b9dd9c40d8..31d87730478c 100644 --- a/test/manual/test_tokenizer_batch_encode.py +++ b/test/manual/test_tokenizer_batch_encode.py @@ -15,6 +15,7 @@ from sglang.srt.managers.io_struct import GenerateReqInput from sglang.srt.managers.tokenizer_manager import TokenizerManager +from sglang.srt.runtime_context import publish from sglang.srt.server_args import PortArgs, ServerArgs from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST @@ -39,6 +40,7 @@ def setUp(self): ): mock_tokenizer.return_value = Mock(vocab_size=32000) + publish(self.server_args, role="tokenizer") self.tokenizer_manager = TokenizerManager(self.server_args, self.port_args) def test_batch_encode_enabled(self): diff --git a/test/manual/test_tokenizer_manager.py b/test/manual/test_tokenizer_manager.py index fb1df645449e..3b9f63ba2678 100644 --- a/test/manual/test_tokenizer_manager.py +++ b/test/manual/test_tokenizer_manager.py @@ -24,6 +24,7 @@ TokenizerManager, ) from sglang.srt.observability.req_time_stats import APIServerReqTimeStats +from sglang.srt.runtime_context import publish from sglang.srt.server_args import PortArgs, ServerArgs from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST @@ -45,6 +46,7 @@ def setUp(self): ) as mock_tokenizer, ): mock_tokenizer.return_value = Mock(vocab_size=32000) + publish(self.server_args, role="tokenizer") self.tokenizer_manager = TokenizerManager(self.server_args, self.port_args) def test_detect_single_string(self): @@ -143,6 +145,7 @@ def setUp(self): ) as mock_tokenizer, ): mock_tokenizer.return_value = Mock(vocab_size=32000) + publish(self.server_args, role="tokenizer") self.tokenizer_manager = TokenizerManager(self.server_args, self.port_args) def test_prepare_single_string_input(self): @@ -203,6 +206,7 @@ def setUp(self): ) as mock_tokenizer, ): mock_tokenizer.return_value = Mock(vocab_size=32000) + publish(self.server_args, role="tokenizer") self.tokenizer_manager = TokenizerManager(self.server_args, self.port_args) def test_extract_single_string_results(self): @@ -327,6 +331,7 @@ def setUp(self): ) as mock_tokenizer, ): mock_tokenizer.return_value = Mock(vocab_size=32000) + publish(self.server_args, role="tokenizer") self.tokenizer_manager = TokenizerManager(self.server_args, self.port_args) def test_full_workflow_single_string(self): diff --git a/test/manual/test_vlm_accuracy.py b/test/manual/test_vlm_accuracy.py index 0387f227a4bb..88fdfec01c2a 100644 --- a/test/manual/test_vlm_accuracy.py +++ b/test/manual/test_vlm_accuracy.py @@ -20,6 +20,7 @@ from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor from sglang.srt.parser.conversation import generate_chat_conv +from sglang.srt.runtime_context import publish from sglang.srt.server_args import ServerArgs from sglang.test.test_utils import download_image_with_retry @@ -141,16 +142,18 @@ def get_processor_output(self, req: Optional[ChatCompletionRequest] = None): return inputs def get_sglang_model(self): + server_args = ServerArgs( + model_path=self.model_path, + disable_cuda_graph=True, + ) + publish(server_args, role="scheduler") self.model_runner = ModelRunner( model_config=ModelConfig(self.model_path, model_override_args="{}"), mem_fraction_static=0.8, gpu_id=0, ps=ParallelState.trivial(), nccl_port=12435, - server_args=ServerArgs( - model_path=self.model_path, - disable_cuda_graph=True, - ), + server_args=server_args, ) return self.model_runner.model diff --git a/test/registered/unit/test_publish_precedes_bag_reads.py b/test/registered/unit/test_publish_precedes_bag_reads.py index fc5c6b34e535..78ced6117755 100644 --- a/test/registered/unit/test_publish_precedes_bag_reads.py +++ b/test/registered/unit/test_publish_precedes_bag_reads.py @@ -78,9 +78,7 @@ "run_data_parallel_controller_process", ), ("srt/ray/scheduler_actor.py", "__init__"), - ("srt/disaggregation/encoder/server.py", "__init__"), ("srt/disaggregation/encoder/http_server.py", "launch_server"), - ("srt/managers/tokenizer_manager.py", "__init__"), ("srt/entrypoints/engine.py", "_launch_subprocesses"), ( "srt/elastic_ep/expert_backup_manager.py", diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index b655102d14a1..11d4fd836a92 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -19,7 +19,7 @@ ParallelContext, RuntimeContext, _FlagGroupBase, - ensure_published, + assert_published, get_context, get_exec, get_flags, @@ -255,35 +255,30 @@ def test_reset_context_clears_owned_store(self): get_server_args() -class TestEnsurePublished(_IsolatedServerArgs): - """A defensive publish must not re-project over a live process. +class TestAssertPublished(_IsolatedServerArgs): + """Publishing is the process entry's job; the constructors only check. - Three constructors publish because each can be built with nothing published - first: `ModelRunner`, `TokenizerManager`, `MMEncoder`. Inside a process that - already published the same record, publishing again re-projects the bags -- - discarding every `override()` taken since, and the provenance log with it. - - No current override sits in one of those windows, so what these assertions - protect is the mechanism, not a reproduction: the drop is silent and depends - on where a constructor happens to sit relative to the overrides around it. + `ModelRunner`, `TokenizerManager` and `MMEncoder` assert. A publish inside + a process that has already published re-projects the bags, discarding every + `override()` taken since and the provenance log with it, so a constructor + that finds nothing published fails loud. """ def _record(self, **fields): return ServerArgs(model_path="dummy", **fields) - def test_a_second_publish_of_the_same_record_keeps_the_overrides(self): + def test_the_check_leaves_a_live_process_alone(self): record = self._record(grammar_backend="xgrammar") publish(record, role="scheduler") get_context().override("grammar.import_fallback", grammar_backend="none") - ensure_published(record, role="scheduler") + assert_published(record, role="scheduler") self.assertEqual( get_exec().kernel.grammar_backend, "none", - "the constructor's publish re-projected the bags, so the import " - "fallback was discarded and the process reports a backend it is " - "not using", + "the check re-projected the bags, so the import fallback was " + "discarded and the process reports a backend it is not using", ) self.assertEqual( len(get_context().overrides_log()), @@ -291,45 +286,49 @@ def test_a_second_publish_of_the_same_record_keeps_the_overrides(self): "the provenance of the override went with it", ) - def test_a_different_record_is_published(self): + def test_a_different_record_fails(self): first = self._record(grammar_backend="xgrammar") publish(first, role="scheduler") second = self._record(grammar_backend="llguidance") - ensure_published(second, role="scheduler") + with self.assertRaisesRegex(RuntimeError, "a different record is published"): + assert_published(second, role="scheduler") - self.assertIs(get_server_args(), second) - self.assertEqual(get_exec().kernel.grammar_backend, "llguidance") + self.assertIs( + get_server_args(), + first, + "the failing check published anyway", + ) - def test_an_empty_slot_is_published(self): - """The standalone case the defensive publish exists for.""" + def test_an_empty_slot_fails(self): + """An empty slot fails.""" reset_context() record = self._record(grammar_backend="xgrammar") - ensure_published(record, role="scheduler") - - self.assertIs(get_server_args(), record) - self.assertEqual(publish_role(), "scheduler") + with self.assertRaisesRegex( + RuntimeError, "nothing is published in this process" + ): + assert_published(record, role="scheduler") - def test_the_same_record_under_a_different_role_is_republished(self): + def test_the_same_record_under_a_different_role_fails(self): """The role decides which namespaces this process may read.""" record = self._record() publish(record, role="tokenizer") - ensure_published(record, role="scheduler") + with self.assertRaisesRegex(RuntimeError, "published under role 'tokenizer'"): + assert_published(record, role="scheduler") - self.assertEqual(publish_role(), "scheduler") + self.assertEqual(publish_role(), "tokenizer") - def test_every_constructor_that_publishes_is_classified(self): - """A new constructor publish has to say which of the two it is. + def test_no_constructor_publishes_outside_the_two_entries(self): + """Publishing from an `__init__` is an entry's job or a bug. - Publishing in a constructor is right when the constructor *is* the - entry -- a spawned worker, the Ray actor that stands in for - `run_scheduler_process`, an `Engine` being (re)built, where resetting - the bags is the point -- and wrong when the process is already live - with the same record, where it silently drops overrides. The - difference is not visible in the syntax, so the census is pinned: - adding one fails here until it is classified. + It is right when the constructor *is* the entry -- an `Engine` being + (re)built, the Ray actor that stands in for `run_scheduler_process`, + where resetting the bags is the point. It is wrong anywhere else, + because the process is already live with a record and re-projecting + drops its overrides. The census is pinned, so a new constructor publish + fails here until it is one of the two. Both the publisher set and "which `__init__` reaches one" come from `sglang.test.config_publishers`, which derives them from the code -- @@ -347,24 +346,11 @@ def test_every_constructor_that_publishes_is_classified(self): self.assertEqual( constructor_publishers(srt), { - # Entries: nothing published yet, or a rebuild that must not - # inherit the previous engine's runtime overrides. ("entrypoints/engine.py", "Engine", "publish"), ("ray/scheduler_actor.py", "SchedulerActor", "publish"), - # Defensive: the process is usually already live with this - # record, and `launch_server` publishes before building the - # in-process encoder. - ("disaggregation/encoder/server.py", "MMEncoder", "ensure_published"), - ( - "managers/tokenizer_manager.py", - "TokenizerManager", - "ensure_published", - ), - ("model_executor/model_runner.py", "ModelRunner", "ensure_published"), }, - "a constructor publishes and this census does not know which kind " - "it is; an entry uses publish(), one that may run inside a live " - "process with the same record uses ensure_published()", + "a constructor publishes and it is not one of the two entries; " + "publish at the process entry and let the constructor assert", ) From 3ee95a3b59315bd8fe99e4cc9ea9bdcde370838f Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 10:54:12 +0000 Subject: [PATCH 04/26] config: stop threading server_args through functions that never read it Forty functions took a `server_args` and never mentioned it. Each one is a reason for its callers to hold a record, and for *their* callers to pass one, which is how a config object ends up threaded through call graphs that have no config decision in them -- `require_mlp_sync`, `require_mlp_tp_gather`, `get_cuda_graph_max_batch_size` and friends read the bags and the live topology already. Removing the parameter cascades: five rounds, each one making another caller's parameter dead, until nothing was left. 104 arguments dropped at 41 call sites. `create_offloader_from_server_args` is now `create_offloader`, since it takes none. Eleven functions keep theirs, and all eleven are contracts rather than dead weight: the platform and spec-algo hooks whose base signature the pipeline calls (`apply_server_args_defaults`, `handle_server_args`), constructors whose siblings read the record (`BaseKVManager`, `RadixCacheCpp`, `TokenizerMetricsCollector`, the expert-distribution gatherer), overridden methods (`StackStrategy.build`, `launch_tensor_parallel_group`, `load_kv_cache_scales`), `release_req`, and `StartupWeightLoadOptions.from_server_args`, which needs a rename with it. --- python/sglang/benchmark/one_batch.py | 4 +-- .../mooncake_transfer_engine.py | 6 ++-- python/sglang/srt/elastic_ep/elastic_ep.py | 4 +-- python/sglang/srt/entrypoints/http_server.py | 2 -- python/sglang/srt/eplb/eplb_manager.py | 3 -- python/sglang/srt/eplb/expert_distribution.py | 6 ++-- python/sglang/srt/eplb/expert_location.py | 28 ++++------------ .../layers/attention/flashinfer_backend.py | 6 ++-- .../layers/attention/linear/gdn_backend.py | 1 - .../layers/attention/linear/kda_backend.py | 1 - .../srt/layers/attention/linear/utils.py | 8 ++--- .../layers/deep_gemm_wrapper/compile_utils.py | 2 +- python/sglang/srt/lora/lora_manager.py | 1 - python/sglang/srt/managers/disagg_service.py | 13 ++------ .../srt/managers/multi_tokenizer_mixin.py | 2 +- python/sglang/srt/managers/overlap_utils.py | 4 +-- python/sglang/srt/managers/scheduler.py | 13 +++----- .../batch_result_processor.py | 8 +---- .../managers/scheduler_components/dp_attn.py | 4 +-- .../srt/managers/tokenizer_control_mixin.py | 4 +-- .../sglang/srt/managers/tokenizer_manager.py | 6 ++-- .../hybrid_cache/hybrid_pool_assembler.py | 2 -- .../srt/mem_cache/kv_cache_configurator.py | 8 +---- .../srt/model_executor/cpu_graph_runner.py | 8 ++--- .../srt/model_executor/forward_batch_info.py | 4 +-- .../sglang/srt/model_executor/model_runner.py | 33 +++++-------------- .../load_model_utils.py | 10 ++---- .../ngram_embedding_manager.py | 2 -- .../spec_aux_hidden_state.py | 2 -- .../srt/model_executor/pool_configurator.py | 2 -- .../runner/base_cuda_graph_runner.py | 4 +-- .../srt/model_executor/runner/base_runner.py | 8 ++--- .../runner/decode_cuda_graph_runner.py | 10 +++--- .../srt/model_executor/runner/eager_runner.py | 4 +-- .../runner/prefill_cuda_graph_runner.py | 8 ++--- python/sglang/srt/models/kimi_k3.py | 5 ++- .../dspark_components/dspark_config.py | 2 +- .../dspark_components/dspark_planner.py | 2 +- .../dspark_components/dspark_worker_v2.py | 2 +- .../eagle_draft_cuda_graph_runner.py | 8 ++--- .../eagle_draft_extend_cuda_graph_runner.py | 8 ++--- .../frozen_kv_mtp_cuda_graph_runner.py | 8 ++--- ...er_eagle_draft_extend_cuda_graph_runner.py | 8 ++--- .../multi_layer_eagle_worker_v2.py | 2 +- python/sglang/srt/utils/common.py | 28 ++++++++-------- .../srt/utils/cuda_vmm_transport_utils.py | 4 +-- python/sglang/srt/utils/offloader.py | 3 +- .../entrypoints/test_http_server_warmup.py | 1 - .../mlx/test_attention_patching.py | 4 --- ...st_batch_result_processor_hidden_states.py | 4 --- ...t_batch_result_processor_mamba_boundary.py | 1 - ...est_batch_result_processor_spec_grammar.py | 1 - .../unit/spec/test_draft_per_runner_config.py | 3 +- 53 files changed, 112 insertions(+), 213 deletions(-) diff --git a/python/sglang/benchmark/one_batch.py b/python/sglang/benchmark/one_batch.py index ea1863500c3f..d6d12a499f1e 100644 --- a/python/sglang/benchmark/one_batch.py +++ b/python/sglang/benchmark/one_batch.py @@ -538,7 +538,7 @@ def decode(input_token_ids, batch, model_runner): def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner): - if require_mlp_sync(model_runner.server_args): + if require_mlp_sync(): prepare_mlp_sync_batch_raw( batch, model_runner=model_runner, @@ -548,7 +548,7 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner): tp_group=model_runner.tp_group, get_idle_batch=None, disable_cuda_graph=cuda_graph_fully_disabled(), - require_mlp_tp_gather=require_mlp_tp_gather(model_runner.server_args), + require_mlp_tp_gather=require_mlp_tp_gather(), disable_overlap_schedule=get_schedule().disable_overlap_schedule, offload_tags=set(), ) diff --git a/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py b/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py index 0eaad79a93d5..8567fe662857 100644 --- a/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py +++ b/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py @@ -14,7 +14,7 @@ from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs + pass logger = logging.getLogger(__name__) @@ -307,9 +307,7 @@ def get_mooncake_transfer_engine() -> Optional[MooncakeTransferEngine]: return _mooncake_transfer_engine -def maybe_init_shared_mooncake_transfer_engine( - *, server_args: ServerArgs, gpu_id: int -) -> None: +def maybe_init_shared_mooncake_transfer_engine(*, gpu_id: int) -> None: """ Need MooncakeTransferEngine when: 1) PD disaggregation uses mooncake for KV transfer (prefill/decode) diff --git a/python/sglang/srt/elastic_ep/elastic_ep.py b/python/sglang/srt/elastic_ep/elastic_ep.py index 1dc43078aa14..73a3489a2973 100644 --- a/python/sglang/srt/elastic_ep/elastic_ep.py +++ b/python/sglang/srt/elastic_ep/elastic_ep.py @@ -111,14 +111,14 @@ def init(cls, server_args: ServerArgs): inst.ep_join_rank_offset = get_parallel().config.ep_join_rank_offset if server_args.is_ep_joiner: - cls._init_joiner_state(inst, server_args) + cls._init_joiner_state(inst) cls._instance = inst return cls._instance @classmethod - def _init_joiner_state(cls, inst: ElasticEPState, server_args: ServerArgs) -> None: + def _init_joiner_state(cls, inst: ElasticEPState) -> None: global_rank = torch.distributed.get_rank() inst.active_ranks.zero_() inst.active_ranks[global_rank] = 1 diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 69575931f735..13817d468720 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -2142,7 +2142,6 @@ def _get_vlm_warmup_image_base64(model_info: dict) -> str: async def _send_disaggregation_warmup_requests( - server_args: ServerArgs, url: str, headers: Dict[str, str], ssl_verify: Union[bool, str], @@ -2321,7 +2320,6 @@ def _execute_server_warmup(server_args: ServerArgs): logger.info(f"Start of pd disaggregation warmup ...") status_codes = asyncio.run( _send_disaggregation_warmup_requests( - server_args=server_args, url=url, headers=headers, ssl_verify=ssl_verify, diff --git a/python/sglang/srt/eplb/eplb_manager.py b/python/sglang/srt/eplb/eplb_manager.py index 00738b40f5c5..d87ebcf1d012 100644 --- a/python/sglang/srt/eplb/eplb_manager.py +++ b/python/sglang/srt/eplb/eplb_manager.py @@ -193,7 +193,6 @@ def _compute_expert_location_metadata( ) -> ExpertLocationMetadata: if not broadcast_over_world: return ExpertLocationMetadata.init_by_eplb( - self._server_args, self._model_config, logical_count, ) @@ -204,7 +203,6 @@ def _compute_expert_location_metadata( # the mapping chosen for the expanded world. if dist.get_rank() == 0: computed_metadata = ExpertLocationMetadata.init_by_eplb( - self._server_args, self._model_config, logical_count, # Arbitrary append topologies may not preserve node divisibility. @@ -220,7 +218,6 @@ def _compute_expert_location_metadata( dist.broadcast(physical_to_logical_map, src=0) return ExpertLocationMetadata.init_by_mapping( - self._server_args, self._model_config, physical_to_logical_map, moe_ep_rank=self._elastic_global_rank(), diff --git a/python/sglang/srt/eplb/expert_distribution.py b/python/sglang/srt/eplb/expert_distribution.py index 693d69264141..ec6133956a7f 100644 --- a/python/sglang/srt/eplb/expert_distribution.py +++ b/python/sglang/srt/eplb/expert_distribution.py @@ -671,12 +671,10 @@ def init_new( expert_location_metadata: ExpertLocationMetadata, rank: int, ) -> _Accumulator: - return _Accumulator.get_class(server_args)( - server_args, expert_location_metadata, rank - ) + return _Accumulator.get_class()(server_args, expert_location_metadata, rank) @staticmethod - def get_class(server_args: ServerArgs) -> Type[_Accumulator]: + def get_class() -> Type[_Accumulator]: return { "stat": _StatAccumulator, "stat_approx": _StatAccumulator, diff --git a/python/sglang/srt/eplb/expert_location.py b/python/sglang/srt/eplb/expert_location.py index 19624be81ed6..9a10e2693659 100644 --- a/python/sglang/srt/eplb/expert_location.py +++ b/python/sglang/srt/eplb/expert_location.py @@ -32,7 +32,6 @@ if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig - from sglang.srt.server_args import ServerArgs logger = logging.getLogger(__name__) @@ -105,11 +104,9 @@ def __post_init__(self): # -------------------------------- construction ------------------------------------ @staticmethod - def init_trivial( - server_args: ServerArgs, model_config: ModelConfig, moe_ep_rank: int - ): + def init_trivial(model_config: ModelConfig, moe_ep_rank: int): """Trivial location - logical expert i corresponds to physical expert i""" - common = ExpertLocationMetadata._init_common(server_args, model_config) + common = ExpertLocationMetadata._init_common(model_config) if common is None: return None @@ -131,7 +128,6 @@ def init_trivial( ) return ExpertLocationMetadata.init_by_mapping( - server_args, model_config, physical_to_logical_map=physical_to_logical_map, moe_ep_rank=moe_ep_rank, @@ -139,7 +135,6 @@ def init_trivial( @staticmethod def init_by_mapping( - server_args: ServerArgs, model_config: ModelConfig, physical_to_logical_map, moe_ep_rank: int = None, @@ -148,7 +143,7 @@ def init_by_mapping( physical_to_logical_map = torch.tensor(physical_to_logical_map) physical_to_logical_map = physical_to_logical_map.to(get_device().device) - common = ExpertLocationMetadata._init_common(server_args, model_config) + common = ExpertLocationMetadata._init_common(model_config) if common is None: return None @@ -179,7 +174,6 @@ def init_by_mapping( @staticmethod def init_by_eplb( - server_args: ServerArgs, model_config: ModelConfig, logical_count: torch.Tensor, *, @@ -193,7 +187,7 @@ def init_by_eplb( from sglang.srt.runtime_context import get_parallel - common = ExpertLocationMetadata._init_common(server_args, model_config) + common = ExpertLocationMetadata._init_common(model_config) if common is None: return None @@ -229,7 +223,7 @@ def init_by_eplb( ) @staticmethod - def _init_common(server_args: ServerArgs, model_config: ModelConfig): + def _init_common(model_config: ModelConfig): from sglang.srt.runtime_context import get_exec, get_parallel model_config_for_expert_location = ( @@ -526,9 +520,6 @@ def broadcast_global_expert_location_metadata( src_rank: int = 0, group: Optional[torch.distributed.ProcessGroup] = None, ) -> ExpertLocationMetadata: - from sglang.srt.runtime_context import get_server_args - - server_args = get_server_args() metadata = get_global_expert_location_metadata() assert metadata is not None @@ -537,7 +528,6 @@ def broadcast_global_expert_location_metadata( metadata.physical_to_logical_map, src=src_rank, group=group ) metadata = ExpertLocationMetadata.init_by_mapping( - server_args, model_config, metadata.physical_to_logical_map, moe_ep_rank=moe_ep_rank, @@ -785,15 +775,12 @@ def from_model_config(model_config: ModelConfig): def compute_initial_expert_location_metadata( - server_args: ServerArgs, model_config: ModelConfig, moe_ep_rank: int, ) -> Optional[ExpertLocationMetadata]: data = get_exec().moe.init_expert_location if data == "trivial": - return ExpertLocationMetadata.init_trivial( - server_args, model_config, moe_ep_rank - ) + return ExpertLocationMetadata.init_trivial(model_config, moe_ep_rank) # TODO unify with the utils function if data.endswith(".pt"): @@ -808,7 +795,6 @@ def compute_initial_expert_location_metadata( "init_expert_location from init_by_mapping using ServerArgs.init_expert_location" ) return ExpertLocationMetadata.init_by_mapping( - server_args, model_config, **data_dict, moe_ep_rank=moe_ep_rank, @@ -818,7 +804,7 @@ def compute_initial_expert_location_metadata( "init_expert_location from init_by_eplb using ServerArgs.init_expert_location" ) return ExpertLocationMetadata.init_by_eplb( - server_args, model_config, logical_count=data_dict["logical_count"] + model_config, logical_count=data_dict["logical_count"] ) else: raise NotImplementedError( diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 0de53ce6475f..6d37cbb7c234 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -442,9 +442,7 @@ def __init__( ) else: self.workspace_buffer = global_workspace_buffer - max_bs = get_cuda_graph_max_batch_size( - model_runner.server_args, model_runner.req_to_token_pool.size - ) + max_bs = get_cuda_graph_max_batch_size(model_runner.req_to_token_pool.size) if kv_indptr_buf is None: self.kv_indptr = [ torch.zeros( @@ -2254,7 +2252,7 @@ def __init__( self.page_size = model_runner.page_size max_bs = get_cuda_graph_max_batch_size( - model_runner.server_args, model_runner.req_to_token_pool.size * self.topk + model_runner.req_to_token_pool.size * self.topk ) self.kv_indptr = torch.zeros( ( diff --git a/python/sglang/srt/layers/attention/linear/gdn_backend.py b/python/sglang/srt/layers/attention/linear/gdn_backend.py index e8033a62d61d..40669014fdc8 100644 --- a/python/sglang/srt/layers/attention/linear/gdn_backend.py +++ b/python/sglang/srt/layers/attention/linear/gdn_backend.py @@ -400,7 +400,6 @@ def __init__(self, model_runner: ModelRunner): self.verify_intermediate_state_indices = ( build_verify_intermediate_state_indices( self.req_to_token_pool.size, - model_runner.server_args, model_runner.device, ) ) diff --git a/python/sglang/srt/layers/attention/linear/kda_backend.py b/python/sglang/srt/layers/attention/linear/kda_backend.py index a0dbba998cb3..406d1c591f3e 100644 --- a/python/sglang/srt/layers/attention/linear/kda_backend.py +++ b/python/sglang/srt/layers/attention/linear/kda_backend.py @@ -408,7 +408,6 @@ def __init__(self, model_runner: ModelRunner): self.verify_intermediate_state_indices = ( build_verify_intermediate_state_indices( self.req_to_token_pool.size, - model_runner.server_args, model_runner.device, ) ) diff --git a/python/sglang/srt/layers/attention/linear/utils.py b/python/sglang/srt/layers/attention/linear/utils.py index 4e63068ec635..c8e469e54205 100644 --- a/python/sglang/srt/layers/attention/linear/utils.py +++ b/python/sglang/srt/layers/attention/linear/utils.py @@ -9,7 +9,7 @@ from sglang.srt.utils.common import rank0_log if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs + pass class LinearAttnKernelBackend(Enum): @@ -103,9 +103,7 @@ def resolve_linear_attn_backends( return backends -def build_verify_intermediate_state_indices( - pool_size: int, server_args: ServerArgs, device -): +def build_verify_intermediate_state_indices(pool_size: int, device): """Per-request row index into the speculative intermediate scratch (`intermediate_ssm` / `intermediate_conv_window`) for the MTP / target_verify path: request slot i owns scratch row i. @@ -123,7 +121,7 @@ def build_verify_intermediate_state_indices( from sglang.srt.utils.common import get_eager_max_batch_size - padded_bs = max(get_eager_max_batch_size(server_args, pool_size), pool_size) + padded_bs = max(get_eager_max_batch_size(pool_size), pool_size) indices = torch.arange(pool_size, dtype=torch.int32, device=device) if padded_bs > pool_size: indices = torch.cat( diff --git a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py index 1e2db45ff462..fe2b659f4557 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py @@ -477,7 +477,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None: cp = max(get_cp_padding_align_size(), 1) attn_tp_size = get_parallel().attn_tp_size - mlp_sync = require_mlp_sync(model_runner.server_args) + mlp_sync = require_mlp_sync() def _align(bs: int) -> int: # Align to lcm(cp, attn_tp_size) so the CP multiple isn't undone by a diff --git a/python/sglang/srt/lora/lora_manager.py b/python/sglang/srt/lora/lora_manager.py index ba07718a05e3..58f1ddee29d4 100644 --- a/python/sglang/srt/lora/lora_manager.py +++ b/python/sglang/srt/lora/lora_manager.py @@ -1030,7 +1030,6 @@ def init_lora_modules(self): def init_lora_cuda_graph_moe_buffers( *, - server_args: ServerArgs, model: torch.nn.Module, lora_manager: LoRAManager, dtype: torch.dtype, diff --git a/python/sglang/srt/managers/disagg_service.py b/python/sglang/srt/managers/disagg_service.py index be82b892b844..d710c601a47c 100644 --- a/python/sglang/srt/managers/disagg_service.py +++ b/python/sglang/srt/managers/disagg_service.py @@ -13,12 +13,9 @@ get_parallel, get_serving, ) -from sglang.srt.server_args import ServerArgs -def start_disagg_service( - server_args: ServerArgs, -): +def start_disagg_service(): # Start kv bootstrap server on prefill disagg_mode = DisaggregationMode(get_disagg().disaggregation_mode) transfer_backend = TransferBackend(get_disagg().disaggregation_transfer_backend) @@ -32,16 +29,12 @@ def start_disagg_service( host=get_serving().host, port=get_disagg().disaggregation_bootstrap_port, ) - maybe_create_ascend_config_store( - server_args=server_args, transfer_backend=transfer_backend - ) + maybe_create_ascend_config_store(transfer_backend=transfer_backend) return bootstrap_server -def maybe_create_ascend_config_store( - server_args: ServerArgs, transfer_backend: TransferBackend -) -> None: +def maybe_create_ascend_config_store(transfer_backend: TransferBackend) -> None: """Also called directly by the rust-server scheduler: there the KV bootstrap registry is served by the embedded rust server's api listener (one rust implementation covers every transfer backend — their diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 3ce1494d7b9d..03772d725621 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -478,7 +478,7 @@ def __init__( ) self._loop.call_soon_threadsafe(self._register_load_snapshot_reader) - self.disaggregation_bootstrap_server = start_disagg_service(self.server_args) + self.disaggregation_bootstrap_server = start_disagg_service() # Worker IPC names for pause/continue broadcasting self.all_worker_ipcs: set[str] = set() diff --git a/python/sglang/srt/managers/overlap_utils.py b/python/sglang/srt/managers/overlap_utils.py index ce78fa8da116..21db5f3960bc 100644 --- a/python/sglang/srt/managers/overlap_utils.py +++ b/python/sglang/srt/managers/overlap_utils.py @@ -18,14 +18,12 @@ from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.mem_cache.memory_pool import ReqToTokenPool - from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.ngram_info import NgramVerifyInput from sglang.srt.speculative.spec_info import SpeculativeAlgorithm def decide_needs_cpu_seq_lens( - server_args: ServerArgs, attn_backends: Sequence[AttentionBackend], ) -> bool: """Whether FutureMap must publish seq_lens_cpu / sum. @@ -53,7 +51,7 @@ def decide_needs_cpu_seq_lens( ) -def decide_needs_confidence_relay(server_args: ServerArgs) -> bool: +def decide_needs_confidence_relay() -> bool: from sglang.srt.speculative.ragged_verify import ( RaggedVerifyMode, read_ragged_verify_mode, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 56b2200cabad..24b2689c9ca5 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -416,7 +416,7 @@ def __init__( # init_soft_watchdog starts a daemon thread that reads these on its first tick. self.forward_ct: int = 0 self.cur_batch_for_debug: Optional[ScheduleBatch] = None - self.init_soft_watchdog(server_args) + self.init_soft_watchdog() # Parse args self.server_args = server_args @@ -911,7 +911,7 @@ def init_moe_gemm_config(self): initialize_bf16_gemm_config(self.server_args) # This must be called after initialize_moe_config - self.require_mlp_sync = require_mlp_sync(self.server_args) + self.require_mlp_sync = require_mlp_sync() def init_tp_model_worker(self): worker_kwargs = dict( @@ -1289,7 +1289,7 @@ def init_schedule_policy(self): self.new_token_ratio_tracker = NewTokenRatioTracker.from_config() - def init_soft_watchdog(self, server_args: ServerArgs): + def init_soft_watchdog(self): if (x := get_device().soft_watchdog_timeout) is not None: self.soft_watchdog = create_scheduler_watchdog( self, watchdog_timeout=x, soft=True @@ -1341,7 +1341,6 @@ def init_disaggregation(self): and self._hosts_rust_server() ): maybe_create_ascend_config_store( - server_args=self.server_args, transfer_backend=self.transfer_backend, ) @@ -1493,8 +1492,8 @@ def init_overlap(self): ) else: attn_backends = (self.tp_worker.model_runner.attn_backend,) - needs_cpu_seq_lens = decide_needs_cpu_seq_lens(self.server_args, attn_backends) - needs_confidence_relay = decide_needs_confidence_relay(self.server_args) + needs_cpu_seq_lens = decide_needs_cpu_seq_lens(attn_backends) + needs_confidence_relay = decide_needs_confidence_relay() self.future_map = self.spec_algorithm.create_future_map( self.device, self.req_to_token_pool, @@ -2089,7 +2088,6 @@ def init_dp_attn_adapter(self) -> None: tree_cache=self.tree_cache, offload_tags=self.weight_updater.offload_tags, ps=self.ps, - server_args=self.server_args, model_config=self.model_config, enable_overlap=self.enable_overlap, spec_algorithm=self.spec_algorithm, @@ -2213,7 +2211,6 @@ def init_batch_result_processor(self) -> None: disaggregation_mode=self.disaggregation_mode, enable_overlap=self.enable_overlap, enable_overlap_mlx=self.enable_overlap_mlx, - server_args=self.server_args, model_config=self.model_config, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, tree_cache=self.tree_cache, 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 42245d088fb5..2deb7bb92056 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -70,7 +70,6 @@ from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput - from sglang.srt.server_args import ServerArgs logger = logging.getLogger(__name__) @@ -81,7 +80,6 @@ class SchedulerBatchResultProcessor: disaggregation_mode: DisaggregationMode enable_overlap: bool enable_overlap_mlx: bool - server_args: ServerArgs model_config: ModelConfig token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator tree_cache: BasePrefixCache @@ -280,7 +278,6 @@ def process_batch_result_prefill( hidden_state_offset = 0 prefill_hidden_capture_mode = self._get_prefill_hidden_capture_mode( batch, - self.server_args, ) # Check finish conditions @@ -617,10 +614,7 @@ def _append_decode_hidden_states( ) @staticmethod - def _get_prefill_hidden_capture_mode( - batch: ScheduleBatch, - server_args: ServerArgs, - ) -> CaptureHiddenMode: + def _get_prefill_hidden_capture_mode(batch: ScheduleBatch) -> CaptureHiddenMode: return get_required_capture_hidden_mode( max( batch.return_hidden_states_mode, diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index b443da6271fc..595c551072ab 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -27,7 +27,6 @@ from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.observability.metrics_collector import DPCooperationInfo from sglang.srt.runtime_context import get_parallel, get_schedule -from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils.common import require_mlp_tp_gather @@ -401,7 +400,6 @@ class SchedulerDPAttnAdapter: tree_cache: BasePrefixCache offload_tags: set[str] ps: ParallelState - server_args: ServerArgs model_config: ModelConfig enable_overlap: bool spec_algorithm: SpeculativeAlgorithm @@ -417,7 +415,7 @@ def prepare_mlp_sync_batch(self, local_batch: ScheduleBatch): tp_group=self.tp_group, get_idle_batch=self.get_idle_batch, disable_cuda_graph=cuda_graph_fully_disabled(), - require_mlp_tp_gather=require_mlp_tp_gather(self.server_args), + require_mlp_tp_gather=require_mlp_tp_gather(), disable_overlap_schedule=get_schedule().disable_overlap_schedule, offload_tags=self.offload_tags, dwdp=get_parallel().config.dwdp_size > 1, diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index c7713a0e48c6..0860a5e9fb83 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -81,7 +81,7 @@ get_parallel, get_spec, ) -from sglang.srt.server_args import LoRARef, ServerArgs +from sglang.srt.server_args import LoRARef from sglang.srt.utils import ( get_bool_env_var, normalize_serialized_named_tensor_payloads, @@ -158,7 +158,7 @@ class TokenizerControlMixin: FanOutCommunicator, as opposed to data-plane inference requests multiplexed by rid. """ - def init_communicators(self: TokenizerManager, server_args: ServerArgs): + def init_communicators(self: TokenizerManager): dispatch_pairs = [] for spec in _COMMUNICATOR_SPECS: name, resp_type = spec[0], spec[1] diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 221877577e82..ee97103c872d 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -663,9 +663,7 @@ def init_disaggregation(self, *, start_pd_bootstrap_service: bool = True): 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) - if start_pd_bootstrap_service - else None + start_disagg_service() if start_pd_bootstrap_service else None ) # Single-source counter for auto-assigning fake bootstrap_room. self.fake_bootstrap_room_counter = 0 @@ -761,7 +759,7 @@ def init_request_dispatcher(self): (ElasticScaleUpdateReq, self.forward_elastic_scale_update), ] ) - self.init_communicators(self.server_args) + self.init_communicators() self.sampling_params_class = SamplingParams self.signal_handler_class = SignalHandler diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index b1900ea65d21..e29de3a0152a 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -410,7 +410,6 @@ def build_hybrid_swa_stack( def _deepseek_v4_num_host_pages( *, params: CacheInitParams, - server_args: ServerArgs, kvcache: Any, page_size: int, swa_page_size: int, @@ -511,7 +510,6 @@ def build_deepseek_v4_hicache_stack( } num_host_pages, swa_num_host_pages = _deepseek_v4_num_host_pages( params=params, - server_args=server_args, kvcache=kvcache, page_size=page_size, swa_page_size=kvcache.swa_page_size, diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 068fda733023..ed0215bfffdf 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -1152,7 +1152,6 @@ def _build_oot_dsa_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: kv_cache_dim=calculate_mla_kv_cache_dim( model_config=self.model_config, kv_cache_dtype=self.kv_cache_dtype, - server_args=self.server_args, ), enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, @@ -1356,7 +1355,6 @@ def _build_dsa_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: kv_cache_dim=calculate_mla_kv_cache_dim( model_config=self.model_config, kv_cache_dtype=self.kv_cache_dtype, - server_args=self.server_args, ), enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, @@ -1396,7 +1394,6 @@ def _build_hybrid_mla_swa_kv_pool( kv_cache_dim=calculate_mla_kv_cache_dim( model_config=self.model_config, kv_cache_dtype=self.kv_cache_dtype, - server_args=self.server_args, ), ) @@ -2199,10 +2196,7 @@ def _handle_max_mamba_cache(self, total_rest_memory): def calculate_mla_kv_cache_dim( - *, - model_config: ModelConfig, - kv_cache_dtype: torch.dtype, - server_args: ServerArgs, + *, model_config: ModelConfig, kv_cache_dtype: torch.dtype ) -> int: is_dsa_model = is_deepseek_dsa(model_config.hf_config) kv_cache_dtype = kv_cache_dtype diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index be0a80da3e3c..7a1b716bf4ce 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -598,10 +598,10 @@ def __init__(self, model_runner: ModelRunner): self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder - self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) - self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) - self.require_mlp_sync = require_mlp_sync(model_runner.server_args) - self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) + self.require_gathered_buffer = require_gathered_buffer() + self.require_mlp_tp_gather = require_mlp_tp_gather() + self.require_mlp_sync = require_mlp_sync() + self.require_attn_tp_gather = require_attn_tp_gather() self.enable_two_batch_overlap = ( model_runner.server_args.enable_two_batch_overlap ) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index b92d5ce07e4a..85e8157e0e92 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -986,14 +986,14 @@ def _maybe_init_non_generation_fields(self, batch: ScheduleBatch): pin_memory=is_pin_memory_available(batch.device), ).to(batch.device, non_blocking=True) - def adjust_num_token_non_padded_for_attn_tp(self, server_args) -> None: + def adjust_num_token_non_padded_for_attn_tp(self) -> None: """Make num_token_non_padded local to this attention-TP rank.""" from sglang.srt.utils.common import require_mlp_tp_gather dp_rank = get_parallel().attn_dp_rank assert self.global_num_tokens_cpu is not None - if require_mlp_tp_gather(server_args): + if require_mlp_tp_gather(): num_tokens_per_dp = self.global_num_tokens_cpu[dp_rank] else: num_tokens_per_dp = self.global_num_tokens_cpu[0] diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index cd3eb88ed685..03542f0b68fc 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -223,7 +223,7 @@ from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks from sglang.srt.utils.nvtx_utils import profile_range from sglang.srt.utils.offloader import ( - create_offloader_from_server_args, + create_offloader, get_offloader, set_offloader, ) @@ -270,7 +270,6 @@ class ModelRunnerOutput: def resolve_draft_attention_backend( *, draft_attention_backend: Optional[str], - server_args: ServerArgs, is_draft_worker: bool, ) -> Optional[str]: """The attention backend a runner uses because it is a draft runner. @@ -342,7 +341,6 @@ def __init__( self.device = get_device().device self.draft_attention_backend = resolve_draft_attention_backend( draft_attention_backend=draft_attention_backend, - server_args=server_args, is_draft_worker=is_draft_worker, ) # This runner's own load format, resolved before anything keys off it: @@ -430,9 +428,7 @@ def __init__( self.shared_read_done_event: Optional[torch.cuda.Event] = None # CPU offload - set_offloader( - create_offloader_from_server_args(server_args, dp_rank=self.ps.dp_rank) - ) + set_offloader(create_offloader(dp_rank=self.ps.dp_rank)) self._weight_checker = WeightChecker(get_model=lambda: self.model, ps=self.ps) @@ -591,7 +587,6 @@ def init_ngram_embedding_manager(self): model=self.model, model_config=self.model_config, req_to_token_pool=self.req_to_token_pool, - server_args=self.server_args, max_running_requests=self.max_running_requests, device=self.device, ) @@ -702,7 +697,6 @@ def maybe_init_expert_location_metadata(self): ) set_global_expert_location_metadata( compute_initial_expert_location_metadata( - server_args=self.server_args, model_config=self.model_config, moe_ep_rank=expert_rank, ) @@ -1076,9 +1070,7 @@ def init_torch_distributed(self): self.pre_model_load_memory = result.pre_model_load_memory def init_shared_mooncake_transfer_engine(self): - maybe_init_shared_mooncake_transfer_engine( - server_args=self.server_args, gpu_id=self.gpu_id - ) + maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id) def load_model(self): tic_total = time.perf_counter() @@ -1091,9 +1083,7 @@ def load_model(self): if self.device != "cpu": torch.set_num_threads(1) if self.device == "cuda": - maybe_downgrade_dtype_for_legacy_gpu( - server_args=self.server_args, model_config=self.model_config - ) + maybe_downgrade_dtype_for_legacy_gpu(model_config=self.model_config) set_cuda_arch() @@ -1113,7 +1103,6 @@ def load_model(self): # and derive the per-rank daemon socket. Idempotent across reloads. maybe_enable_ipc_weight_cache( load_config=self.load_config, - server_args=self.server_args, tp_size=self.ps.tp_size, pp_rank=self.ps.pp_rank, tp_rank=self.ps.tp_rank, @@ -1124,7 +1113,6 @@ def load_model(self): ) maybe_trigger_remote_instance_nccl_send_group( - server_args=self.server_args, tp_rank=self.ps.tp_rank, load_format=draft_load_format, ) @@ -1195,11 +1183,12 @@ def load_model(self): f"mem usage={self.weight_load_mem_usage:.2f} GB." ) - report_online_quantization(model=self.model, server_args=self.server_args) + report_online_quantization( + model=self.model, + ) maybe_register_debug_tensor_dump_hook( model=self.model, - server_args=self.server_args, spec_algorithm=self.spec_algorithm, is_draft_worker=self.is_draft_worker, tp_size=self.ps.tp_size, @@ -1281,7 +1270,6 @@ def init_lora_manager(self): ) if not cuda_graph_fully_disabled(): init_lora_cuda_graph_moe_buffers( - server_args=self.server_args, model=self.model, lora_manager=self.lora_manager, dtype=self.dtype, @@ -1471,13 +1459,11 @@ def _prepare_eager_forward_batch(self, forward_batch: ForwardBatch) -> None: if ( forward_batch.num_token_non_padded is not None and forward_batch.global_num_tokens_gpu is not None - and require_gathered_buffer(self.server_args) + and require_gathered_buffer() and not is_dsa_enable_prefill_cp() and not is_mla_prefill_cp_enabled() ): - forward_batch.adjust_num_token_non_padded_for_attn_tp( - server_args=self.server_args, - ) + forward_batch.adjust_num_token_non_padded_for_attn_tp() # Hisparse coordinator — backends now read it from self.model_runner. if self.hisparse_coordinator is not None: @@ -1926,7 +1912,6 @@ def _expand_eplb_metadata_for_scale( start=old_num_physical - num_local * initial_ep_size, ) new_metadata = ExpertLocationMetadata.init_by_mapping( - self.server_args, self.model_config, physical_to_logical_map=expanded_p2l, moe_ep_rank=self._elastic_global_rank(), diff --git a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py index 4051d5a25a41..8a3b92ae211d 100644 --- a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py +++ b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py @@ -67,9 +67,7 @@ class LoadedModel(msgspec.Struct, frozen=True, kw_only=True): startup_weight_load: Optional[Any] = None -def maybe_downgrade_dtype_for_legacy_gpu( - *, server_args: ServerArgs, model_config: ModelConfig -) -> None: +def maybe_downgrade_dtype_for_legacy_gpu(*, model_config: ModelConfig) -> None: if torch.cuda.get_device_capability()[0] < 8: logger.info( "Compute capability below sm80. Use float16 due to lack of bfloat16 support." @@ -85,7 +83,7 @@ def maybe_downgrade_dtype_for_legacy_gpu( def maybe_trigger_remote_instance_nccl_send_group( - *, server_args: ServerArgs, tp_rank: int, load_format: Optional[str] = None + *, tp_rank: int, load_format: Optional[str] = None ) -> None: """``load_format`` is this runner's effective format: a draft loading under ``--speculative-draft-draft-load-format`` needs its own send group, and the @@ -151,7 +149,7 @@ def resolve_sliding_window_size(model, model_config: ModelConfig) -> Optional[in return sliding_window_size -def report_online_quantization(*, model, server_args: ServerArgs) -> None: +def report_online_quantization(*, model) -> None: # TODO: Make sure all models have `quant_config` attribute, and all online quantization methods register which layers they actually quantize. quantized_layers = getattr( getattr(model, "quant_config", None), "quantized_layers", None @@ -170,7 +168,6 @@ def report_online_quantization(*, model, server_args: ServerArgs) -> None: def maybe_register_debug_tensor_dump_hook( *, model, - server_args: ServerArgs, spec_algorithm: SpeculativeAlgorithm, is_draft_worker: bool, tp_size: int, @@ -237,7 +234,6 @@ def build_load_config( def maybe_enable_ipc_weight_cache( *, load_config: LoadConfig, - server_args: ServerArgs, tp_size: int, pp_rank: int, tp_rank: int, diff --git a/python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py b/python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py index 4d034f781bf5..d4054cce1b2d 100644 --- a/python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py +++ b/python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py @@ -12,7 +12,6 @@ from sglang.srt.managers.schedule_batch import ForwardMode from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.runtime_context import get_schedule -from sglang.srt.server_args import ServerArgs if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req, ScheduleBatch @@ -33,7 +32,6 @@ def from_model( model: torch.nn.Module, model_config: ModelConfig, req_to_token_pool: ReqToTokenPool, - server_args: ServerArgs, max_running_requests: int, device: str, ): diff --git a/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py b/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py index 3e26fd8adb26..61a2a573b417 100644 --- a/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py +++ b/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py @@ -203,7 +203,6 @@ def _resolve_dflash_aux_hidden_state( config.dflash_draft_num_layers = int(draft_num_layers) config.dflash_target_layer_ids = target_layer_ids config.dflash_draft_cell_size_per_token = _resolve_dflash_draft_cell_size( - server_args=server_args, draft_model_config=draft_model_config, draft_num_layers=int(draft_num_layers), ) @@ -211,7 +210,6 @@ def _resolve_dflash_aux_hidden_state( def _resolve_dflash_draft_cell_size( *, - server_args: ServerArgs, draft_model_config: ModelConfig, draft_num_layers: int, ) -> int | None: diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index e38553218209..3f40b902b81a 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -259,7 +259,6 @@ def _compute_cell_size(self, kvc: KVCacheConfigurator, num_layers: int) -> int: calculate_mla_kv_cache_dim( model_config=model_config, kv_cache_dtype=kv_cache_dtype, - server_args=kvc.server_args, ) * effective_num_layers * kv_size @@ -465,7 +464,6 @@ def __init__(self, kvc: KVCacheConfigurator): calculate_mla_kv_cache_dim( model_config=model_config, kv_cache_dtype=kv_cache_dtype, - server_args=kvc.server_args, ) * kv_size ) diff --git a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py index f58717e19c19..58a076ea6012 100644 --- a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py @@ -74,7 +74,7 @@ def get_batch_sizes_to_capture( capture_bs = list(get_exec().graph.cuda_graph_config.decode.bs) num_max_requests = model_runner.req_to_token_pool.size - mul_base = get_cuda_graph_batch_size_alignment(server_args) + mul_base = get_cuda_graph_batch_size_alignment() # TBO splits each request's rows across two micro-batches, so the # alignment constraint applies per request rather than per token row. alignment_width = captured_req_width @@ -82,7 +82,7 @@ def get_batch_sizes_to_capture( alignment_width = 1 # pad `num_max_requests` to avoid being filtered out - num_max_requests = get_cuda_graph_max_batch_size(server_args, num_max_requests) + num_max_requests = get_cuda_graph_max_batch_size(num_max_requests) if max(capture_bs) > num_max_requests: # In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests # is very small. We add more values here to make sure we capture the maximum bs. diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 20a024aada67..1ebc26310e6d 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -351,7 +351,7 @@ def _alloc_dummy_decode_buffers( dp_size=get_parallel().config.dp_size, pp_size=get_parallel().config.pp_size, is_encoder_decoder=mr.model_config.is_encoder_decoder, - require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args), + require_mlp_tp_gather=require_mlp_tp_gather(), seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(), encoder_len_fill_value=( getattr(mr.model_config.hf_config, "max_source_positions", 0) @@ -535,9 +535,9 @@ def _dummy_run( ) # TP-gather requirements for global token metadata. - require_mlp_tp_gather_ = require_mlp_tp_gather(mr.server_args) - require_attn_tp_gather_ = require_attn_tp_gather(mr.server_args) - if require_gathered_buffer(mr.server_args): + require_mlp_tp_gather_ = require_mlp_tp_gather() + require_attn_tp_gather_ = require_attn_tp_gather() + if require_gathered_buffer(): assert require_mlp_tp_gather_ or require_attn_tp_gather_ if require_mlp_tp_gather_: diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 65323e63beb5..a954cd1816d1 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -227,10 +227,10 @@ def __init__( self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder - self.require_mlp_tp_gather = require_mlp_tp_gather( - model_runner.server_args - ) and not self._forward_is_dp_local(model_runner) - self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) + self.require_mlp_tp_gather = ( + require_mlp_tp_gather() and not self._forward_is_dp_local(model_runner) + ) + self.require_attn_tp_gather = require_attn_tp_gather() # Composite predicates derive from the instance values so the dp-local # draft exemption above stays consistent (require_gathered_buffer == # mlp_tp_gather or attn_tp_gather; require_mlp_sync adds dp attention). @@ -597,7 +597,7 @@ def _forward_is_dp_local(model_runner) -> bool: draft_is_deepseek_v4, ) - return not draft_is_deepseek_v4(server_args=model_runner.server_args) + return not draft_is_deepseek_v4() def _ragged_capture_slots(self, num_tokens: int) -> int: if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get(): diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 6476c3a87b48..83caa44d6de6 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -113,10 +113,10 @@ def __init__(self, model_runner: ModelRunner) -> None: # (expand_for_topk_draft) before the eager fallback. max_bs *= get_spec().speculative_eagle_topk # Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies. - max_bs = get_eager_max_batch_size(sa, max_bs) + max_bs = get_eager_max_batch_size(max_bs) prefill_ceiling = max(mr.max_total_num_tokens, max_prefill_buffer_tokens()) max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req) - if require_mlp_sync(sa): + if require_mlp_sync(): from sglang.srt.layers.cp.padding import get_cp_padding_align_size max_num_token = ceil_align(max_num_token, self.attn_tp_size) diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index 96d57bfca263..b10ca8f093c3 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -332,7 +332,7 @@ def __init__(self, model_runner: ModelRunner): embed_dtype=self.model_runner.dtype, enable_mamba_track=self.mamba_track_enabled, enable_num_token_non_padded=enable_num_token_non_padded(), - require_gathered_buffer=require_gathered_buffer(model_runner.server_args), + require_gathered_buffer=require_gathered_buffer(), enable_prefill_cp=( is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled() ), @@ -349,8 +349,8 @@ def __init__(self, model_runner: ModelRunner): self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None) self.dp_size = get_parallel().config.dp_size - self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) - self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) + self.require_mlp_tp_gather = require_mlp_tp_gather() + self.require_attn_tp_gather = require_attn_tp_gather() # --- backend --------------------------------------------------- # TcPiecewise resolves by running a compile pass that calls back into @@ -602,7 +602,7 @@ def _capture_num_token_non_padded(self, num_tokens: int) -> Optional[torch.Tenso buf = self.buffer_registry.get_slot("num_token_non_padded").buffer buf.fill_(num_tokens) - if require_gathered_buffer(self.model_runner.server_args): + if require_gathered_buffer(): local = compute_local_num_token_non_padded( global_num_token_non_padded=buf, num_tokens_per_dp=num_tokens, diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index a7a59d87a3e0..7c15a3ff1c8a 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -121,7 +121,6 @@ from sglang.srt.runtime_context import ( get_exec, get_parallel, - get_server_args, ) from sglang.srt.utils import is_blackwell_supported, is_hip, is_npu, make_layers from sglang.srt.utils.common import ( @@ -2146,7 +2145,7 @@ def __init__( self._dp_attention = is_dp_attention_enabled() # mlp-sync (DP attention OR MoE a2a/EP) pads extend batches to # attn_tp multiples; attention must then run on the real rows only. - self._trim_padded_attn = require_mlp_sync(get_server_args()) + self._trim_padded_attn = require_mlp_sync() # A layer runs MoE (vs a plain dense MLP) iff it is past the dense # prefix and on the MoE cadence — same predicate the mlp construction # below uses. @@ -2614,7 +2613,7 @@ def __init__( self.pp_group = get_pp_group() self.dspark_layers_to_capture: Optional[list[int]] = None self._dp_attention = is_dp_attention_enabled() - self._trim_padded_attn = require_mlp_sync(get_server_args()) + self._trim_padded_attn = require_mlp_sync() if self.pp_group.is_first_rank: embedding_quant_config = ( diff --git a/python/sglang/srt/speculative/dspark_components/dspark_config.py b/python/sglang/srt/speculative/dspark_components/dspark_config.py index 068e3ff7c20c..88bc14405b3b 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_config.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_config.py @@ -30,7 +30,7 @@ def get_dspark_sample_from_anchor(draft_hf_config: Any) -> bool: return bool(_cfg_get(draft_hf_config, "sample_from_anchor", True)) -def draft_is_deepseek_v4(*, server_args: ServerArgs) -> bool: +def draft_is_deepseek_v4() -> bool: from sglang.srt.configs.model_config import is_deepseek_v4 from sglang.srt.utils.hf_transformers_utils import get_config diff --git a/python/sglang/srt/speculative/dspark_components/dspark_planner.py b/python/sglang/srt/speculative/dspark_components/dspark_planner.py index 1bdfa7c4a203..edde1bbf1b1b 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_planner.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_planner.py @@ -172,7 +172,7 @@ def __init__( and is_dp_attention_enabled() and get_parallel().attn_tp_size == 1 and get_parallel().attn_cp_size == 1 - and require_mlp_tp_gather(self.server_args) + and require_mlp_tp_gather() and not get_schedule().disable_overlap_schedule and not get_spec().speculative_skip_dp_mlp_sync and get_disagg().disaggregation_mode == "null" diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index 53fe887c2c84..08c8d78967d4 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -103,7 +103,7 @@ def __init__( self.page_size = get_schedule().page_size self.device = target_worker.device - self._draft_is_moe = draft_is_deepseek_v4(server_args=server_args) + self._draft_is_moe = draft_is_deepseek_v4() self._draft_dp_context_enabled = ( get_parallel().config.enable_dp_attention and not self._draft_is_moe ) diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index dc21efc40524..ab2f5ce2758c 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -116,10 +116,10 @@ def __init__( self.pp_size = get_parallel().config.pp_size self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding - self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) - self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) - self.require_mlp_sync = require_mlp_sync(model_runner.server_args) - self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) + self.require_gathered_buffer = require_gathered_buffer() + self.require_mlp_tp_gather = require_mlp_tp_gather() + self.require_mlp_sync = require_mlp_sync() + self.require_attn_tp_gather = require_attn_tp_gather() self.enable_profile_cuda_graph = ( model_runner.server_args.enable_profile_cuda_graph ) diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 3c4c63d7facb..2ee2f7c2805d 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -102,10 +102,10 @@ def __init__( self.pp_size = get_parallel().config.pp_size self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding - self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) - self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) - self.require_mlp_sync = require_mlp_sync(model_runner.server_args) - self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) + self.require_gathered_buffer = require_gathered_buffer() + self.require_mlp_tp_gather = require_mlp_tp_gather() + self.require_mlp_sync = require_mlp_sync() + self.require_attn_tp_gather = require_attn_tp_gather() self.enable_profile_cuda_graph = ( model_runner.server_args.enable_profile_cuda_graph ) diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py index 3c7352035875..43914af2a92f 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py @@ -93,10 +93,10 @@ def __init__(self, frozen_kv_mtp_worker: FrozenKVMTPDraftWorker): self.device_module = torch.get_device_module(self.device) self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding - self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) - self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) - self.require_mlp_sync = require_mlp_sync(model_runner.server_args) - self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) + self.require_gathered_buffer = require_gathered_buffer() + self.require_mlp_tp_gather = require_mlp_tp_gather() + self.require_mlp_sync = require_mlp_sync() + self.require_attn_tp_gather = require_attn_tp_gather() self.tp_size = self.model_runner.ps.tp_size self.attn_dp_size = self.model_runner.ps.attn_dp_size self.pp_size = get_parallel().config.pp_size diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index aa15c97da51a..3dc1c6739bc6 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -158,10 +158,10 @@ def __init__(self, eagle_worker: MultiLayerEagleDraftWorker, step: int): self.pp_size = get_parallel().config.pp_size self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding - self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) - self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) - self.require_mlp_sync = require_mlp_sync(model_runner.server_args) - self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) + self.require_gathered_buffer = require_gathered_buffer() + self.require_mlp_tp_gather = require_mlp_tp_gather() + self.require_mlp_sync = require_mlp_sync() + self.require_attn_tp_gather = require_attn_tp_gather() self.enable_pdmux = model_runner.server_args.enable_pdmux self.speculative_num_steps = get_spec().speculative_num_steps self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index ea75a3fd2c16..7807e8960dab 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -846,7 +846,7 @@ def _draft_extend_for_decode( self.cuda_graph_runner_for_draft_extend.prune_draft_extend_logits ) else: - prune_logits = not require_gathered_buffer(self.server_args) + prune_logits = not require_gathered_buffer() if prune_logits: forward_batch.spec_info.select_index = select_index # Left unmarked on every platform: each de-tied runner has its own diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index d1fd11a023fd..c7972f9253a7 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -105,7 +105,7 @@ from sglang.srt.utils.video_decoder import _BACKEND, VideoDecoderWrapper if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs + pass logger = logging.getLogger(__name__) torch_release = pkg_version.parse(torch.__version__).release @@ -3719,7 +3719,7 @@ def with_value(self, new_value: T): self._value = None -def require_mlp_tp_gather(server_args: ServerArgs): +def require_mlp_tp_gather(): """ Check if the input of MLP is obtained by all-gather rather than all-reduce. This only happens when each MLP TP group contains multiple attention DP groups. """ @@ -3763,7 +3763,7 @@ def require_mlp_tp_gather(server_args: ServerArgs): return False -def require_attn_tp_gather(server_args: ServerArgs): +def require_attn_tp_gather(): """ Check if the input of attention is scattered. """ @@ -3790,35 +3790,33 @@ def require_attn_tp_gather(server_args: ServerArgs): return False -def require_gathered_buffer(server_args: ServerArgs): - return require_mlp_tp_gather(server_args) or require_attn_tp_gather(server_args) +def require_gathered_buffer(): + return require_mlp_tp_gather() or require_attn_tp_gather() -def require_mlp_sync(server_args: ServerArgs): +def require_mlp_sync(): from sglang.srt.runtime_context import get_parallel - return get_parallel().config.enable_dp_attention or require_gathered_buffer( - server_args - ) + return get_parallel().config.enable_dp_attention or require_gathered_buffer() -def get_cuda_graph_batch_size_alignment(server_args: ServerArgs) -> int: +def get_cuda_graph_batch_size_alignment() -> int: alignment = 1 if get_exec().overlap.enable_two_batch_overlap: alignment *= 2 - if require_gathered_buffer(server_args): + if require_gathered_buffer(): alignment *= get_parallel().attn_tp_size if alignment % get_parallel().attn_cp_size != 0: alignment *= get_parallel().attn_cp_size return alignment -def get_cuda_graph_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int: - return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment(server_args)) +def get_cuda_graph_max_batch_size(max_batch_size: int) -> int: + return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment()) -def get_eager_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int: - if not require_mlp_sync(server_args): +def get_eager_max_batch_size(max_batch_size: int) -> int: + if not require_mlp_sync(): return max_batch_size from sglang.srt.layers.cp.padding import get_cp_padding_align_size diff --git a/python/sglang/srt/utils/cuda_vmm_transport_utils.py b/python/sglang/srt/utils/cuda_vmm_transport_utils.py index 8525e5812d44..d115701bf448 100644 --- a/python/sglang/srt/utils/cuda_vmm_transport_utils.py +++ b/python/sglang/srt/utils/cuda_vmm_transport_utils.py @@ -161,7 +161,7 @@ def _contains_tensor_container(value) -> bool: ) -def get_vmm_feature_consumer_count(server_args) -> int: +def get_vmm_feature_consumer_count() -> int: if get_parallel().config.enable_dp_attention: return get_parallel().config.tp_size // get_parallel().config.dp_size return get_parallel().config.tp_size @@ -947,7 +947,7 @@ def __init__(self, server_args, mm_processor) -> None: memory_size=per_worker_pool_size, recycle_interval=MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL, base_gpu_id=server_args.base_gpu_id, - consumer_count=get_vmm_feature_consumer_count(server_args), + consumer_count=get_vmm_feature_consumer_count(), allow_posix_fallback=server_args.nnodes == 1, ) diff --git a/python/sglang/srt/utils/offloader.py b/python/sglang/srt/utils/offloader.py index 63c4a8954c27..e38b0f2a1928 100644 --- a/python/sglang/srt/utils/offloader.py +++ b/python/sglang/srt/utils/offloader.py @@ -17,7 +17,6 @@ get_parallel, get_stream, ) -from sglang.srt.server_args import ServerArgs from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available from sglang.srt.utils.host_shared_memory import ( HostSharedMemoryManager, @@ -66,7 +65,7 @@ def set_offloader(instance: BaseOffloader): _instance = instance -def create_offloader_from_server_args(server_args: ServerArgs, dp_rank: int): +def create_offloader(dp_rank: int): if get_exec().offload.cpu_offload_gb > 0: return OffloaderV1( cpu_offload_max_bytes=int(get_exec().offload.cpu_offload_gb * 1024**3) diff --git a/test/registered/unit/entrypoints/test_http_server_warmup.py b/test/registered/unit/entrypoints/test_http_server_warmup.py index e4391e456ca2..a74d02b85c06 100644 --- a/test/registered/unit/entrypoints/test_http_server_warmup.py +++ b/test/registered/unit/entrypoints/test_http_server_warmup.py @@ -57,7 +57,6 @@ def post(self, *args, **kwargs): with patch("sglang.srt.entrypoints.http_server.aiohttp.ClientSession", Session): status_codes = await _send_disaggregation_warmup_requests( - server_args=server_args, url="http://localhost:30000", headers={"Authorization": "Bearer token"}, ssl_verify=False, diff --git a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py index 14dc72d4eeee..662da5b09687 100644 --- a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py +++ b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py @@ -1200,10 +1200,6 @@ def test_finished_request_snapshots_before_release(self): disaggregation_mode=None, enable_overlap=False, enable_overlap_mlx=False, - server_args=SimpleNamespace( - disaggregation_decode_enable_offload_kvcache=False, - enable_hisparse=False, - ), model_config=None, token_to_kv_pool_allocator=None, tree_cache=tree_cache, 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 82d1beee5823..f606842e05ae 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 @@ -31,10 +31,6 @@ def _make_processor(case, server_mode: str = "full") -> SchedulerBatchResultProc disaggregation_mode=None, enable_overlap=False, enable_overlap_mlx=False, - server_args=SimpleNamespace( - enable_metrics=False, - enable_hisparse=False, - ), model_config=SimpleNamespace(think_end_ids=None), token_to_kv_pool_allocator=Mock(), tree_cache=None, diff --git a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py index 2cc5f4d46def..0c618ef04782 100644 --- a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py +++ b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py @@ -60,7 +60,6 @@ def _make_processor() -> SchedulerBatchResultProcessor: disaggregation_mode=None, enable_overlap=True, enable_overlap_mlx=False, - server_args=SimpleNamespace(), model_config=SimpleNamespace(think_end_ids=None), token_to_kv_pool_allocator=MagicMock(), tree_cache=SimpleNamespace(page_size=TRACK_INTERVAL), diff --git a/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py b/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py index 144d2b13a961..baee3db67c50 100644 --- a/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py +++ b/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py @@ -65,7 +65,6 @@ def _make_processor() -> SchedulerBatchResultProcessor: disaggregation_mode=None, enable_overlap=False, enable_overlap_mlx=False, - server_args=SimpleNamespace(enable_metrics=False), model_config=SimpleNamespace(think_end_ids=None), token_to_kv_pool_allocator=None, tree_cache=None, diff --git a/test/registered/unit/spec/test_draft_per_runner_config.py b/test/registered/unit/spec/test_draft_per_runner_config.py index 274a71a512fb..29c7ce928127 100644 --- a/test/registered/unit/spec/test_draft_per_runner_config.py +++ b/test/registered/unit/spec/test_draft_per_runner_config.py @@ -134,14 +134,13 @@ def test_the_draft_backend_applies_to_the_draft_runner_only(self): def test_an_unresolved_draft_falls_back_to_the_config_field(self): """The v2 workers pass no backend: --speculative-draft-attention-backend.""" - server_args = self._seed( + self._seed( attention_backend="fa3", speculative_draft_attention_backend="triton" ) def effective(*, is_draft_worker, passed=None): return resolve_draft_attention_backend( draft_attention_backend=passed, - server_args=server_args, is_draft_worker=is_draft_worker, ) From 9b414c7b3c04efd68f38baa9391da1e2bf4688a5 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 10:56:16 +0000 Subject: [PATCH 05/26] config: name the weight-load factories after what they read `StartupWeightLoadOptions.from_server_args` took a record and read sixteen config leaves off the bags instead -- the name said the opposite of what the body did, and the parameter kept every caller holding a record for it. It is `from_published_config(is_draft_worker=...)` now, and the manager factory above it `create_from_published_config`. `is_draft_worker` stays an argument: that is this runner's role, not the process's configuration. --- .../load_model_utils.py | 3 +-- .../startup_weight_load.py | 21 ++++++++++--------- .../test_startup_weight_load.py | 3 +-- 3 files changed, 13 insertions(+), 14 deletions(-) diff --git a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py index 8a3b92ae211d..9218179037a2 100644 --- a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py +++ b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py @@ -307,12 +307,11 @@ def load_model_with_memory_saver( StartupWeightLoadManager, ) - startup_weight_load = StartupWeightLoadManager.create_from_server_args( + startup_weight_load = StartupWeightLoadManager.create_from_published_config( loader=loader, model_config=model_config, load_config=load_config, device_config=device_config, - server_args=server_args, is_draft_worker=is_draft_worker, ) model = startup_weight_load.prepare() diff --git a/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py b/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py index 202e602c551b..3633856a5935 100644 --- a/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py +++ b/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py @@ -31,7 +31,6 @@ if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig - from sglang.srt.server_args import ServerArgs logger = logging.getLogger(__name__) @@ -98,12 +97,16 @@ class StartupWeightLoadOptions: prefetch_num_threads: int @classmethod - def from_server_args( + def from_published_config( cls, *, - server_args: ServerArgs, is_draft_worker: bool, ) -> StartupWeightLoadOptions: + """Everything this needs is a published leaf; nothing comes off a record. + + `is_draft_worker` is the exception and travels as an argument: it is + this runner's role, not the process's configuration. + """ cuda_graph_config = get_exec().graph.cuda_graph_config cuda_graph_enabled = any( getattr(cuda_graph_config, phase).backend != Backend.DISABLED @@ -267,29 +270,27 @@ def __init__( self._prefetch_failure_reported = False @classmethod - def create_from_server_args( + def create_from_published_config( cls, *, loader, model_config: ModelConfig, load_config: LoadConfig, device_config: DeviceConfig, - server_args: ServerArgs, is_draft_worker: bool, ) -> StartupWeightLoadManager: - """Build a manager straight from ``ServerArgs``. + """Build a manager from the published configuration. Callers on the model-loading path only decide *whether* to overlap; the - knowledge of which server arguments matter, and every support rule, - stays in this module. + knowledge of which config leaves matter, and every support rule, stays + in this module. """ return cls.create( loader=loader, model_config=model_config, load_config=load_config, device_config=device_config, - options=StartupWeightLoadOptions.from_server_args( - server_args=server_args, + options=StartupWeightLoadOptions.from_published_config( is_draft_worker=is_draft_worker, ), ) diff --git a/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py b/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py index ba5a924dc01d..dd183c3793d0 100644 --- a/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py +++ b/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py @@ -209,8 +209,7 @@ def test_options_accept_current_server_args_schema(self): # The parallel sizes come from the bags, so the config has to be published. publish(server_args, role="test") self.addCleanup(reset_context) - options = StartupWeightLoadOptions.from_server_args( - server_args=server_args, + options = StartupWeightLoadOptions.from_published_config( is_draft_worker=False, ) From 981bebcd2cf70199d7fa0871fd8347133bccd4ec Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 11:04:38 +0000 Subject: [PATCH 06/26] config: three managers read the bags instead of a handed record `EPLBManager` took a `ServerArgs` and read nine config leaves off it, all of them `exec.moe` or `parallel`; it reads the bags now and no longer takes a record at all. `LoRAManager` and the multimodal base processor lose the same kind of read -- `get_lora()` / `get_mm()` / `get_serving()` -- which is also what makes them follow a post-publish override instead of reporting startup. What stays on a record is what has to: the base processor's `base_gpu_id`, `tp_size` and `rl_on_policy_target` are read from *its own* instance because engines sharing a process each have their own, and the LoRA backend below still takes one. Four (file, field) pairs leave the supplied-instance exposure pin. --- python/sglang/srt/eplb/eplb_manager.py | 23 ++++++++----------- python/sglang/srt/lora/lora_manager.py | 12 ++++------ .../sglang/srt/model_executor/model_runner.py | 1 - .../runner/base_cuda_graph_runner.py | 1 - .../multimodal/processors/base_processor.py | 22 +++++++++--------- .../unit/managers/test_mm_process_config.py | 14 +++++------ test/registered/unit/models/test_kimi_k25.py | 15 +++++++----- .../unit/multimodal/rust/qwen/_fixtures.py | 8 +++++-- ...test_supplied_instance_exposure_ratchet.py | 5 ---- 9 files changed, 48 insertions(+), 53 deletions(-) diff --git a/python/sglang/srt/eplb/eplb_manager.py b/python/sglang/srt/eplb/eplb_manager.py index d87ebcf1d012..b591a9cd05f0 100644 --- a/python/sglang/srt/eplb/eplb_manager.py +++ b/python/sglang/srt/eplb/eplb_manager.py @@ -18,11 +18,10 @@ get_global_expert_location_metadata, ) from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater -from sglang.srt.runtime_context import get_model +from sglang.srt.runtime_context import get_exec, get_model, get_parallel if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig - from sglang.srt.server_args import ServerArgs logger = logging.getLogger(__name__) @@ -31,7 +30,6 @@ class EPLBManager: def __init__( self, *, - server_args: ServerArgs, model_config: ModelConfig, ps: Any, get_model: Callable[[], nn.Module], @@ -43,7 +41,6 @@ def __init__( # These collaborators are set on ModelRunner AFTER EPLBManager is # constructed (model load, expert_backup_client, weight_updater), so # they are read through getters at rebalance time, not captured here. - self._server_args = server_args self._model_config = model_config self._ps = ps self._get_model = get_model @@ -51,16 +48,16 @@ def __init__( self._get_expert_backup_client = get_expert_backup_client self._get_weight_updater = get_weight_updater self._rebalance_layers_per_chunk = ( - self._server_args.eplb_rebalance_layers_per_chunk + get_exec().moe.eplb_rebalance_layers_per_chunk ) - self._rebalance_num_iterations = self._server_args.eplb_rebalance_num_iterations + self._rebalance_num_iterations = get_exec().moe.eplb_rebalance_num_iterations self._rebalance_disabled_reason = None self._rebalance_disabled_logged = False # Otherwise, the circular buffer will contain stale data. If the case is needed, it can be implemented. assert ( - self._server_args.eplb_rebalance_num_iterations - >= self._server_args.expert_distribution_recorder_buffer_size + get_exec().moe.eplb_rebalance_num_iterations + >= get_exec().moe.expert_distribution_recorder_buffer_size ), "eplb_rebalance_num_iterations must be greater than expert_distribution_recorder_buffer_size" if not get_global_expert_distribution_recorder().recording: @@ -160,7 +157,7 @@ def rebalance(self): model=self._get_model(), new_expert_location_metadata=expert_location_metadata, update_layer_ids=chunk_layer_ids, - nnodes=self._server_args.nnodes, + nnodes=get_parallel().config.nnodes, tp_rank=( self._elastic_global_rank() if is_post_scale_rebalance @@ -169,7 +166,7 @@ def rebalance(self): use_flat_topology=is_post_scale_rebalance, expert_backup_client=self._get_expert_backup_client(), update_weights_from_disk_callable=self._get_weight_updater().update_weights_from_disk, - ep_dispatch_algorithm=self._server_args.ep_dispatch_algorithm, + ep_dispatch_algorithm=get_exec().moe.ep_dispatch_algorithm, init_lplb_solvers_callable=lambda: init_lplb_solvers( model_config=self._model_config ), @@ -224,7 +221,7 @@ def _compute_expert_location_metadata( ) def _elastic_global_rank(self) -> int: - return self._ps.tp_rank + self._server_args.ep_join_rank_offset + return self._ps.tp_rank + get_parallel().config.ep_join_rank_offset def _check_rebalance_needed(self, average_utilization_rate_over_window): if average_utilization_rate_over_window is None: @@ -232,10 +229,10 @@ def _check_rebalance_needed(self, average_utilization_rate_over_window): if ( average_utilization_rate_over_window - > self._server_args.eplb_min_rebalancing_utilization_threshold + > get_exec().moe.eplb_min_rebalancing_utilization_threshold ): logger.info( - f"[EPLBManager] Skipped ep rebalancing: current GPU utilization {average_utilization_rate_over_window:.2f} > minimum rebalance threshold {self._server_args.eplb_min_rebalancing_utilization_threshold:.2f}" + f"[EPLBManager] Skipped ep rebalancing: current GPU utilization {average_utilization_rate_over_window:.2f} > minimum rebalance threshold {get_exec().moe.eplb_min_rebalancing_utilization_threshold:.2f}" ) return False diff --git a/python/sglang/srt/lora/lora_manager.py b/python/sglang/srt/lora/lora_manager.py index 58f1ddee29d4..f83cededf288 100644 --- a/python/sglang/srt/lora/lora_manager.py +++ b/python/sglang/srt/lora/lora_manager.py @@ -94,19 +94,17 @@ def __init__( self.attn_tp_size: int = get_parallel().attn_tp_size self.lora_added_tokens_size: Optional[int] = None self.enable_lora_overlap_loading: Optional[bool] = ( - server_args.enable_lora_overlap_loading + get_lora().enable_lora_overlap_loading ) self.pending_lora_load_events = {} - self.eviction_policy = server_args.lora_eviction_policy + self.eviction_policy = get_lora().lora_eviction_policy self.enable_dp_attention: bool = get_parallel().config.enable_dp_attention self._experts_shared_outer_override: Optional[bool] = ( - server_args.experts_shared_outer_loras - ) - self.lora_use_virtual_experts: bool = server_args.lora_use_virtual_experts - self.lora_strict_loading: bool = getattr( - server_args, "lora_strict_loading", False + get_lora().experts_shared_outer_loras ) + self.lora_use_virtual_experts: bool = get_lora().lora_use_virtual_experts + self.lora_strict_loading: bool = get_lora().lora_strict_loading self.speculative_algorithm: Optional[str] = get_spec().speculative_algorithm # LoRA backend for running sgemm kernels diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 03542f0b68fc..f7947fb61a83 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -721,7 +721,6 @@ def maybe_init_lplb_solvers(self): def maybe_init_eplb_manager(self): self.eplb_manager = ( EPLBManager( - server_args=self.server_args, model_config=self.model_config, ps=self.ps, get_model=lambda: self.model, diff --git a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py index 58a076ea6012..880ea3136dd7 100644 --- a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py @@ -70,7 +70,6 @@ def get_batch_sizes_to_capture( constraints and clamps to req_to_token_pool.size. """ - server_args = model_runner.server_args capture_bs = list(get_exec().graph.cuda_graph_config.decode.bs) num_max_requests = model_runner.req_to_token_pool.size diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 2ba4e508c507..ff3f60a8c7a9 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -40,7 +40,7 @@ MmItemMemoryPool, get_mm_feature_pool_size_per_worker, ) -from sglang.srt.runtime_context import get_mm +from sglang.srt.runtime_context import get_mm, get_serving from sglang.srt.utils import ( CLIENT_MEDIA_EXCEPTIONS, configure_media_url_security, @@ -215,7 +215,7 @@ def __init__( self.transport_mode = transport_mode configure_media_url_security( get_mm().allowed_media_domains, - server_args.media_url_max_file_size_mb, + get_mm().media_url_max_file_size_mb, ) configured_mm_feature_transport = get_mm().mm_feature_transport self.mm_feature_transport = ( @@ -227,11 +227,11 @@ def __init__( self.use_ipc_pool_handle_cache = ( self.use_cuda_ipc and envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get() ) - self.image_processor_backend = server_args.image_processor_backend - if server_args.disable_fast_image_processor: + self.image_processor_backend = get_mm().image_processor_backend + if get_mm().disable_fast_image_processor: self.image_processor_backend = "pil" self.disable_fast_image_processor = self.image_processor_backend == "pil" - self.skip_tokenizer_init = server_args.skip_tokenizer_init + self.skip_tokenizer_init = get_serving().skip_tokenizer_init mm_process_config = get_mm().mm_process_config self.image_config = mm_process_config.get("image", {}) @@ -241,19 +241,19 @@ def __init__( # Each tokenizer worker is a separate process with its own CPU cache. # Split the requested service-wide budget so increasing worker count # does not silently multiply host-memory usage. - requested_cache_mb = self.server_args.mm_preprocess_cache_size_mb + requested_cache_mb = get_mm().mm_preprocess_cache_size_mb total_cache_mb = ( self.auto_mm_preprocess_cache_size_mb if requested_cache_mb is None else requested_cache_mb ) - tokenizer_worker_num = max(int(self.server_args.tokenizer_worker_num), 1) + tokenizer_worker_num = max(int(get_serving().tokenizer_worker_num), 1) worker_cache_bytes = total_cache_mb * 1024 * 1024 // tokenizer_worker_num self.mm_preprocess_cache = MultimodalPreprocessCache( max_size_bytes=worker_cache_bytes, max_entries=8192, ) - self.trust_mm_content_hashes = bool(self.server_args.trust_mm_content_hashes) + self.trust_mm_content_hashes = bool(get_mm().trust_mm_content_hashes) # The fingerprint is needed only to build artifact keys. Avoid inspecting # processor state when this processor will never retain artifacts. self.processor_fingerprint = ( @@ -288,7 +288,7 @@ def __init__( # FIXME: not accurate, model and image specific self.NUM_TOKEN_PER_FRAME = 330 - requested_mm_io_worker_num = self.server_args.mm_io_worker_num + requested_mm_io_worker_num = get_mm().mm_io_worker_num env_mm_io_worker_num = os.environ.get("SGLANG_IO_WORKERS") if requested_mm_io_worker_num: self.mm_io_worker_num = requested_mm_io_worker_num @@ -310,7 +310,7 @@ def __init__( io_worker_source, ) skip_mm_pool = kwargs.get("skip_mm_pool", False) - requested_mm_processor_worker_num = self.server_args.mm_processor_worker_num + requested_mm_processor_worker_num = get_mm().mm_processor_worker_num self.mm_processor_worker_num = ( 1 if skip_mm_pool @@ -402,7 +402,7 @@ def __init__( # SGLANG_MM_FEATURE_CACHE_MB is the total pool budget across all # tokenizer workers. Each worker gets an equal share so that adding # workers doesn't multiply the GPU-side footprint. - worker_num = self.server_args.tokenizer_worker_num + worker_num = get_serving().tokenizer_worker_num per_worker_pool_size = get_mm_feature_pool_size_per_worker( MM_FEATURE_CACHE_SIZE, worker_num ) diff --git a/test/registered/unit/managers/test_mm_process_config.py b/test/registered/unit/managers/test_mm_process_config.py index 04a9e82f187f..a5e8adc30970 100644 --- a/test/registered/unit/managers/test_mm_process_config.py +++ b/test/registered/unit/managers/test_mm_process_config.py @@ -81,21 +81,21 @@ def _make_processor( BaseMultimodalProcessor, ) - # The multimodal config comes from the bags. + # The worker counts and the cache budget are bag leaves. override = get_context().override_server_args( mm_process_config=mm_process_config, allowed_media_domains=[], + mm_processor_worker_num=mm_processor_worker_num, + mm_io_worker_num=mm_io_worker_num, + mm_preprocess_cache_size_mb=None, + tokenizer_worker_num=1, + trust_mm_content_hashes=False, + media_url_max_file_size_mb=64, ) override.install() self.addCleanup(override.restore) server_args = MagicMock() - 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.media_url_max_file_size_mb = 64 hf_config = MagicMock() mock_hf_processor = MagicMock() diff --git a/test/registered/unit/models/test_kimi_k25.py b/test/registered/unit/models/test_kimi_k25.py index dbb2046b2bff..9e4591fe647d 100644 --- a/test/registered/unit/models/test_kimi_k25.py +++ b/test/registered/unit/models/test_kimi_k25.py @@ -647,7 +647,16 @@ def _k3_preprocess_config( ], ) def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls): + # The per-worker instance reads: device identity and the RL target. server_args = SimpleNamespace( + base_gpu_id=0, + rl_on_policy_target=None, + tp_size=1, + ) + with get_context().override_server_args( + mm_feature_transport="cpu", + mm_process_config={}, + allowed_media_domains=[], image_processor_backend="auto", disable_fast_image_processor=False, skip_tokenizer_init=False, @@ -656,13 +665,7 @@ def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls tokenizer_worker_num=1, mm_preprocess_cache_size_mb=0, trust_mm_content_hashes=False, - base_gpu_id=0, - rl_on_policy_target=None, media_url_max_file_size_mb=64, - ) - # 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), diff --git a/test/registered/unit/multimodal/rust/qwen/_fixtures.py b/test/registered/unit/multimodal/rust/qwen/_fixtures.py index 28534a362ea3..3ef0a8448d73 100644 --- a/test/registered/unit/multimodal/rust/qwen/_fixtures.py +++ b/test/registered/unit/multimodal/rust/qwen/_fixtures.py @@ -89,14 +89,18 @@ def make_processor(case, 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. + # The processor reads its media policy, transport, per-modality limits and + # image-processor backend from the mm bag, so the fixture publishes before + # building it. The backend matters here: left at the default the fast image + # processor runs and sends the tensor to `cuda:`, which a + # CPU-only host cannot do. 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, + disable_fast_image_processor=server_args.disable_fast_image_processor, ), role="tokenizer", ) diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index aa9af73ce56c..3c366806f0d5 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -137,7 +137,6 @@ _EXPOSED = { ("dllm/config.py", "max_running_requests"), ("dllm/config.py", "model_path"), - ("multimodal/processors/base_processor.py", "image_processor_backend"), ("speculative/spec_registry.py", "disable_overlap_schedule"), ("disaggregation/encoder/server.py", "model_loader_extra_config"), ("layers/moe/utils.py", "deepep_mode"), @@ -165,8 +164,6 @@ ("entrypoints/engine.py", "enable_symm_mem"), ("entrypoints/engine.py", "reasoning_parser"), ("entrypoints/engine.py", "tool_call_parser"), - ("eplb/eplb_manager.py", "ep_dispatch_algorithm"), - ("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"), ("layers/cp/base.py", "attn_cp_size"), ("layers/cp/base.py", "cp_strategy"), ("layers/cp/base.py", "enable_prefill_cp"), @@ -178,11 +175,9 @@ ("layers/moe/utils.py", "moe_runner_backend"), ("layers/moe/utils.py", "quantization"), ("layers/moe/utils.py", "speculative_moe_runner_backend"), - ("lora/lora_manager.py", "enable_lora_overlap_loading"), ("lora/marlin_lora_temp/policy.py", "lora_paths"), ("model_loader/expert_pack_runtime.py", "model_path"), ("model_loader/expert_pack_runtime.py", "tokenizer_path"), - ("multimodal/processors/base_processor.py", "image_processor_backend"), ("parser/template_detection.py", "model_path"), ("speculative/adaptive_spec_params.py", "speculative_algorithm"), ("speculative/adaptive_spec_params.py", "speculative_eagle_topk"), From 29c1b9eee22aaab202e9373e6e8771c6cea99f9d Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 11:24:15 +0000 Subject: [PATCH 07/26] config: the resolution hooks read the declaration stash, not the fields `declare_resolution` writes the field as it declares, so a resolver that reads the field afterwards sees the declared value -- and every one of the 1130 mid-resolution field reads in this tree depends on that write. Removing it (the last step of making `ServerArgs` hold only raw input) currently breaks resolution outright: 20 of 20 launch shapes die on `AttributeError: 'NoneType' object has no attribute 'prefill'`. This is the first half of that: the readers move to `resolving_view`, a live view that answers from the declaration stash and falls through to the field. While the immediate write is still there the two agree, so this is a no-op -- which is exactly what makes it checkable: 20 launch shapes, every field of the resolution result compared against the previous commit, zero differences. 337 reads in `arg_groups`: the speculative hook (146), the override passes (103), the DeepSeek-V4 and PD-disaggregation hooks, and the three small model hooks. `ResolvedView` stays for the post-process passes, which want the state snapshotted at their slot; `ResolvingConfig` is for a reader that outlives a declaration. --- .../sglang/srt/arg_groups/deepseek_v4_hook.py | 76 ++--- .../sglang/srt/arg_groups/expert_pack_hook.py | 29 +- python/sglang/srt/arg_groups/hisparse_hook.py | 7 +- python/sglang/srt/arg_groups/kimi_k3_hook.py | 26 +- python/sglang/srt/arg_groups/mega_moe_hook.py | 15 +- python/sglang/srt/arg_groups/overrides.py | 287 ++++++++++------- .../srt/arg_groups/pd_disaggregation_hook.py | 67 ++-- .../sglang/srt/arg_groups/speculative_hook.py | 298 +++++++++--------- .../srt/model_loader/expert_pack_runtime.py | 20 +- ...test_supplied_instance_exposure_ratchet.py | 3 - 10 files changed, 451 insertions(+), 377 deletions(-) diff --git a/python/sglang/srt/arg_groups/deepseek_v4_hook.py b/python/sglang/srt/arg_groups/deepseek_v4_hook.py index c2c78f46ffa5..0359d2a4c798 100644 --- a/python/sglang/srt/arg_groups/deepseek_v4_hook.py +++ b/python/sglang/srt/arg_groups/deepseek_v4_hook.py @@ -3,7 +3,10 @@ import logging from typing import TYPE_CHECKING -from sglang.srt.arg_groups.overrides import declare_resolution +from sglang.srt.arg_groups.overrides import ( + declare_resolution, + resolving_view, +) from sglang.srt.environ import envs if TYPE_CHECKING: @@ -16,18 +19,16 @@ def validate_deepseek_v4_mega_moe_token_budget( server_args: ServerArgs, ) -> None: """Ensure the DSV4 prefill budget fits MegaMoE's per-rank buffer.""" - mega_moe_enabled = server_args.moe_a2a_backend == "megamoe" - if not mega_moe_enabled or server_args.disaggregation_mode == "decode": + cfg = resolving_view(server_args) + mega_moe_enabled = cfg.moe_a2a_backend == "megamoe" + if not mega_moe_enabled or cfg.disaggregation_mode == "decode": # decode node will skip the check because decode bs is not relevant with --chunk-prefill-size return - if server_args.pp_size > 1 and server_args.enable_dynamic_chunking: + if cfg.pp_size > 1 and cfg.enable_dynamic_chunking: return - if ( - server_args.chunked_prefill_size is None - or server_args.chunked_prefill_size <= 0 - ): + if cfg.chunked_prefill_size is None or cfg.chunked_prefill_size <= 0: raise ValueError( "DeepSeekV4 with MegaMoE requires chunked prefill to be enabled. " "Set --chunked-prefill-size to a positive value; " @@ -35,40 +36,38 @@ def validate_deepseek_v4_mega_moe_token_budget( "token requirement would not have a strict prefill-forward bound." ) - if server_args.enable_prefill_cp: - token_partition_size = server_args.attn_cp_size + if cfg.enable_prefill_cp: + token_partition_size = cfg.attn_cp_size token_partition_name = "attn_cp_size" token_alignment = 1 local_chunked_prefill_size = ( - server_args.chunked_prefill_size + token_partition_size - 1 + cfg.chunked_prefill_size + token_partition_size - 1 ) // token_partition_size - elif server_args.enable_dp_attention: - token_partition_size = server_args.dp_size + elif cfg.enable_dp_attention: + token_partition_size = cfg.dp_size token_partition_name = "dp_size" token_alignment = max( - server_args.tp_size // server_args.dp_size // server_args.attn_cp_size, + cfg.tp_size // cfg.dp_size // cfg.attn_cp_size, 1, ) - local_chunked_prefill_size = ( - server_args.chunked_prefill_size // token_partition_size - ) + local_chunked_prefill_size = cfg.chunked_prefill_size // token_partition_size else: # Pure TP and PP with static chunking are handled here. token_partition_size = 1 token_partition_name = "none" # global_num_tokens will ceil_align to attn_tp_size so the validation needs to do alignment as well token_alignment = max( - server_args.tp_size // token_partition_size // server_args.attn_cp_size, + cfg.tp_size // token_partition_size // cfg.attn_cp_size, 1, ) - local_chunked_prefill_size = server_args.chunked_prefill_size + local_chunked_prefill_size = cfg.chunked_prefill_size if local_chunked_prefill_size <= 0: raise ValueError( "DeepSeekV4 with MegaMoE requires a positive effective per-rank " "chunked prefill size. " f"Current values: chunked_prefill_size=" - f"{server_args.chunked_prefill_size}, " + f"{cfg.chunked_prefill_size}, " f"token_partition={token_partition_name}, " f"token_partition_size={token_partition_size}." ) @@ -87,7 +86,7 @@ def validate_deepseek_v4_mega_moe_token_budget( "SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK to " "cover each rank's effective prefill token budget. " f"Current values: chunked_prefill_size=" - f"{server_args.chunked_prefill_size}, " + f"{cfg.chunked_prefill_size}, " f"token_partition={token_partition_name}, " f"token_partition_size={token_partition_size}, " f"token_alignment={token_alignment}, " @@ -112,6 +111,7 @@ def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None max_running_requests fill (the speculative hook is a later writer of that field) and the validations. """ + cfg = resolving_view(server_args) from sglang.srt.utils import is_hip # FlashMLA sparse prefill (SGLANG_OPT_FLASHMLA_SPARSE_PREFILL, default on) @@ -136,36 +136,36 @@ def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None run_post_process_pass(server_args, _deepseek_v4_kv_cache_dtype) - if server_args.max_running_requests is None: + if cfg.max_running_requests is None: declare_resolution( server_args, "apply_deepseek_v4_defaults", max_running_requests=256, ) logger.warning( - f"Setting max_running_requests to {server_args.max_running_requests} for {model_arch}." + f"Setting max_running_requests to {cfg.max_running_requests} for {model_arch}." ) - if server_args.speculative_algorithm is not None: - assert server_args.speculative_algorithm in ( + if cfg.speculative_algorithm is not None: + assert cfg.speculative_algorithm in ( "EAGLE", "DSPARK", ), f"Only EAGLE and DSPARK speculative algorithms are supported for {model_arch}" - if server_args.speculative_algorithm == "EAGLE": + if cfg.speculative_algorithm == "EAGLE": assert ( - server_args.speculative_eagle_topk == 1 + cfg.speculative_eagle_topk == 1 ), f"Only EAGLE speculative algorithm with topk == 1 is supported for {model_arch}" def validate_deepseek_v4_cp(server_args: ServerArgs) -> None: """Validate DeepSeek V4 context-parallel configuration.""" - if not server_args.enable_prefill_cp: + cfg = resolving_view(server_args) + if not cfg.enable_prefill_cp: return - if server_args.cp_strategy != "interleave": + if cfg.cp_strategy != "interleave": raise ValueError( - "DeepSeekV4 only supports interleave CP strategy, " - f"got {server_args.cp_strategy}" + "DeepSeekV4 only supports interleave CP strategy, " f"got {cfg.cp_strategy}" ) declare_resolution( @@ -196,19 +196,19 @@ def validate_deepseek_v4_cp(server_args: ServerArgs) -> None: declare_resolution( server_args, "validate_deepseek_v4_cp", - attn_cp_size=server_args.tp_size // server_args.dp_size, + attn_cp_size=cfg.tp_size // cfg.dp_size, ) assert ( - server_args.dp_size == 1 + cfg.dp_size == 1 ), "For round-robin split mode, dp attention is not supported." assert ( - server_args.tp_size <= 8 + cfg.tp_size <= 8 ), "Context parallel only supports single machine (tp_size <= 8). Cross-machine CP has precision issues." - if server_args.moe_a2a_backend not in ("none", "deepep", "megamoe"): + if cfg.moe_a2a_backend not in ("none", "deepep", "megamoe"): raise ValueError( "DeepSeekV4 CP supports moe_a2a_backend in " "('none', 'deepep', 'megamoe'), " - f"got {server_args.moe_a2a_backend!r}." + f"got {cfg.moe_a2a_backend!r}." ) logger.warning( "Disabling SGLANG_OPT_FLASHMLA_SPARSE_PREFILL because DeepSeekV4 " @@ -217,6 +217,6 @@ def validate_deepseek_v4_cp(server_args: ServerArgs) -> None: envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.set(False) logger.warning( f"Enable Context Parallel for DeepSeekV4, " - f"dp_size={server_args.dp_size}, moe_dense_tp_size={server_args.moe_dense_tp_size}, " - f"attn_cp_size={server_args.attn_cp_size}, ep_size={server_args.ep_size}, tp_size={server_args.tp_size}" + f"dp_size={cfg.dp_size}, moe_dense_tp_size={cfg.moe_dense_tp_size}, " + f"attn_cp_size={cfg.attn_cp_size}, ep_size={cfg.ep_size}, tp_size={cfg.tp_size}" ) diff --git a/python/sglang/srt/arg_groups/expert_pack_hook.py b/python/sglang/srt/arg_groups/expert_pack_hook.py index b7ee90484a07..4908c4355194 100644 --- a/python/sglang/srt/arg_groups/expert_pack_hook.py +++ b/python/sglang/srt/arg_groups/expert_pack_hook.py @@ -9,7 +9,7 @@ from pathlib import Path from typing import Any -from sglang.srt.arg_groups.overrides import declare_resolution +from sglang.srt.arg_groups.overrides import declare_resolution, resolving_view from sglang.srt.environ import envs from sglang.srt.model_executor.cuda_graph_config import ( Backend, @@ -27,31 +27,32 @@ def handle_expert_pack(server_args: Any) -> None: """Normalize expert-pack settings and report all startup errors together.""" - if server_args.load_format != "expert_pack": + cfg = resolving_view(server_args) + if cfg.load_format != "expert_pack": return errors = [] parallelism = ( - ("tensor", "--tp-size", server_args.tp_size), - ("data", "--dp-size", server_args.dp_size), - ("expert", "--ep-size", server_args.ep_size), + ("tensor", "--tp-size", cfg.tp_size), + ("data", "--dp-size", cfg.dp_size), + ("expert", "--ep-size", cfg.ep_size), ) for label, option, size in parallelism: if size != 1: errors.append(f"{label} parallelism ({option}) must be 1, got {size}") - if server_args.enforce_shared_experts_fusion: + if cfg.enforce_shared_experts_fusion: errors.append( "--enforce-shared-experts-fusion is incompatible with expert_pack" ) - if server_args.enable_waterfill: + if cfg.enable_waterfill: errors.append("--enable-waterfill is incompatible with expert_pack") explicit_cuda_graph_backends = { - Phase.DECODE: server_args.cuda_graph_backend_decode, - Phase.PREFILL: server_args.cuda_graph_backend_prefill, + Phase.DECODE: cfg.cuda_graph_backend_decode, + Phase.PREFILL: cfg.cuda_graph_backend_prefill, } - raw_cuda_graph_config = server_args.cuda_graph_config + raw_cuda_graph_config = cfg.cuda_graph_config if isinstance(raw_cuda_graph_config, CudaGraphConfig): raw_cuda_graph_config = raw_cuda_graph_config.to_dict() for phase in Phase.ALL: @@ -69,7 +70,7 @@ def handle_expert_pack(server_args: Any) -> None: f"disabled, got {explicit_backend!r}" ) - loader_config = server_args.model_loader_extra_config or {} + loader_config = cfg.model_loader_extra_config or {} if isinstance(loader_config, str): try: loader_config = json.loads(loader_config) @@ -82,7 +83,7 @@ def handle_expert_pack(server_args: Any) -> None: # A raw GGUF path is the public input form. Preparation is performed once # here, before model-config parsing and before the loader is constructed. - raw_model_path = Path(server_args.model_path).expanduser() + raw_model_path = Path(cfg.model_path).expanduser() raw_preparation_failed = False if not errors and raw_model_path.is_file(): try: @@ -136,12 +137,12 @@ def parse_path(label: str, value: Any) -> Path | None: errors.append(f"expert-pack file does not exist: {pack_path}") model_kind = None - model_path = parse_path("--model-path", server_args.model_path) + model_path = parse_path("--model-path", cfg.model_path) if not raw_preparation_failed: if model_path is None or not model_path.is_dir(): errors.append( "--model-path must be a local GGUF shard or tokenizer/config " - f"directory for expert_pack, got {server_args.model_path!r}" + f"directory for expert_pack, got {cfg.model_path!r}" ) else: try: diff --git a/python/sglang/srt/arg_groups/hisparse_hook.py b/python/sglang/srt/arg_groups/hisparse_hook.py index a5cc2661ef62..99b2b448ce06 100644 --- a/python/sglang/srt/arg_groups/hisparse_hook.py +++ b/python/sglang/srt/arg_groups/hisparse_hook.py @@ -3,6 +3,8 @@ import logging from typing import TYPE_CHECKING +from sglang.srt.arg_groups.overrides import resolving_view + if TYPE_CHECKING: from sglang.srt.server_args import ServerArgs @@ -79,7 +81,8 @@ def validate_hisparse_kv_cache_dtype(server_args: ServerArgs) -> None: def validate_hisparse(server_args: ServerArgs) -> None: """Validate --enable-hisparse constraints (model class, radix cache, DSA backend).""" - if not server_args.enable_hisparse: + cfg = resolving_view(server_args) + if not cfg.enable_hisparse: return from sglang.srt.configs.model_config import ( @@ -96,7 +99,7 @@ def validate_hisparse(server_args: ServerArgs) -> None: ) assert ( - server_args.disable_radix_cache + cfg.disable_radix_cache ), "Hierarchical sparse attention currently requires --disable-radix-cache." # DSv4 hisparse handles its own dtype/backend pairing elsewhere; the dtype- diff --git a/python/sglang/srt/arg_groups/kimi_k3_hook.py b/python/sglang/srt/arg_groups/kimi_k3_hook.py index 9b42917f8004..464df72d08e5 100644 --- a/python/sglang/srt/arg_groups/kimi_k3_hook.py +++ b/python/sglang/srt/arg_groups/kimi_k3_hook.py @@ -3,7 +3,10 @@ import logging from typing import TYPE_CHECKING -from sglang.srt.arg_groups.overrides import declare_resolution +from sglang.srt.arg_groups.overrides import ( + declare_resolution, + resolving_view, +) if TYPE_CHECKING: from sglang.srt.server_args import ServerArgs @@ -13,15 +16,16 @@ def apply_kimi_k3_spec_backend_defaults(server_args: ServerArgs) -> None: """Apply speculative backend defaults for Kimi hybrid models.""" + cfg = resolving_view(server_args) from sglang.srt.utils import is_sm100_supported - if server_args.speculative_algorithm is None: + if cfg.speculative_algorithm is None: return # Use the fused Kimi-K3/DSPARK CuTeDSL kernel for KDA target verification. # Decode is left free (its bf16-ssm SM100+ flashinfer default is fine -- the # target only verifies under spec); the verify backend is pinned directly. - if server_args.linear_attn_verify_backend is None: + if cfg.linear_attn_verify_backend is None: declare_resolution( server_args, "apply_kimi_k3_spec_backend_defaults", @@ -36,8 +40,8 @@ def apply_kimi_k3_spec_backend_defaults(server_args: ServerArgs) -> None: # dspark's draft is dense MQA; trtllm_mha avoids flashinfer's blocking # per-step host plan. DSPARK-only: other spec algos use MLA-family drafts. if ( - server_args.speculative_algorithm == "DSPARK" - and server_args.speculative_draft_attention_backend is None + cfg.speculative_algorithm == "DSPARK" + and cfg.speculative_draft_attention_backend is None and is_sm100_supported() ): declare_resolution( @@ -63,20 +67,21 @@ def disable_kimi_k3_symm_mem(server_args: ServerArgs) -> None: Gates on the arch itself: this runs from cuda-graph resolution, which is earlier than the model-specific hook block. """ + cfg = resolving_view(server_args) from sglang.srt.connector import ConnectorType from sglang.srt.model_executor.cuda_graph_config import Backend from sglang.srt.utils import parse_connector_type - if not server_args.enable_symm_mem: + if not cfg.enable_symm_mem: return - if parse_connector_type(server_args.model_path) == ConnectorType.INSTANCE: + if parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE: return if server_args.get_model_config().hf_config.architectures[0] not in ( "KimiLinearForCausalLM", "KimiK3ForConditionalGeneration", ): return - graph = server_args.cuda_graph_config + graph = cfg.cuda_graph_config if ( graph.decode.backend == Backend.DISABLED and graph.prefill.backend == Backend.DISABLED @@ -100,14 +105,15 @@ def disable_kimi_k3_symm_mem(server_args: ServerArgs) -> None: def apply_kimi_k3_linear_attn_defaults(server_args: ServerArgs) -> None: """KDA decode-fallback default for Kimi hybrid models (spec-independent).""" + cfg = resolving_view(server_args) from sglang.srt.utils import is_sm100_supported # Preempts the generic SM100+bf16 flashinfer switch (a GDN default): on # KDA shapes the triton packed decode measures ~35% faster than # recurrent_kda across bs 1-256, and ReplaySSM requires triton. if ( - server_args.linear_attn_decode_backend is None - and server_args.mamba_ssm_dtype == "bfloat16" + cfg.linear_attn_decode_backend is None + and cfg.mamba_ssm_dtype == "bfloat16" and is_sm100_supported() ): declare_resolution( diff --git a/python/sglang/srt/arg_groups/mega_moe_hook.py b/python/sglang/srt/arg_groups/mega_moe_hook.py index 0c1806e69619..99777ce15843 100644 --- a/python/sglang/srt/arg_groups/mega_moe_hook.py +++ b/python/sglang/srt/arg_groups/mega_moe_hook.py @@ -7,7 +7,10 @@ if TYPE_CHECKING: from sglang.srt.server_args import ServerArgs -from sglang.srt.arg_groups.overrides import declare_resolution +from sglang.srt.arg_groups.overrides import ( + declare_resolution, + resolving_view, +) logger = logging.getLogger(__name__) @@ -18,15 +21,16 @@ def handle_mega_moe(server_args: ServerArgs) -> None: def handle_moe_runner_backend_alias(server_args: ServerArgs) -> None: - if server_args.moe_runner_backend != "megamoe": + cfg = resolving_view(server_args) + if cfg.moe_runner_backend != "megamoe": return - if server_args.moe_a2a_backend not in ("none", "megamoe"): + if cfg.moe_a2a_backend not in ("none", "megamoe"): logger.warning( "--moe-runner-backend megamoe is an alias for " "--moe-a2a-backend megamoe; overriding " "--moe-a2a-backend %s.", - server_args.moe_a2a_backend, + cfg.moe_a2a_backend, ) declare_resolution( server_args, @@ -37,7 +41,8 @@ def handle_moe_runner_backend_alias(server_args: ServerArgs) -> None: def handle_w4a4_mxfp4_megamoe_env(server_args: ServerArgs) -> None: - if not server_args.enable_w4a4_mxfp4_megamoe: + cfg = resolving_view(server_args) + if not cfg.enable_w4a4_mxfp4_megamoe: return os.environ["DG_USE_FP4_ACTS"] = "1" diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index b0d7838812c2..a52b66995ccb 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -150,6 +150,42 @@ def __setattr__(self, name: str, value: Any) -> None: ) +class ResolvingConfig: + """Live read view of the resolution result: the declaration stash over the + record's fields, looked up per read. + + ``ResolvedView`` snapshots the overlay when it is built, which is what a + post-process pass wants -- it reads the state at its slot. A resolver that + reads *after* declaring, or after calling something that declares, needs the + current answer instead, so this one walks the stash on every read. It falls + through to the field, which is where the raw input lives. + """ + + __slots__ = ("_server_args",) + + def __init__(self, server_args: Any): + object.__setattr__(self, "_server_args", server_args) + + def __getattr__(self, name: str) -> Any: + server_args = object.__getattribute__(self, "_server_args") + for _source, declared in reversed( + getattr(server_args, "_resolved_overrides", None) or () + ): + if name in declared: + return declared[name] + return getattr(server_args, name) + + def __setattr__(self, name: str, value: Any) -> None: + raise AttributeError( + "ResolvingConfig is read-only; resolution writes through declarations" + ) + + +def resolving_view(server_args: Any) -> ResolvingConfig: + """A live read view of what resolution has decided so far.""" + return ResolvingConfig(server_args) + + # Ordered post-process passes (the normalization stage). List order is the # end-state execution order and mirrors today's handler call sequence in # __post_init__; during the transition each pass is invoked from its legacy @@ -581,16 +617,17 @@ def _require_kimi_k3_cutedsl_dcp_support() -> None: @_register_for("KimiK3ForConditionalGeneration") def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: - if server_args.dcp_size > 1: + cfg = resolving_view(server_args) + if cfg.dcp_size > 1: overrides = {} - if server_args.enable_symm_mem: + if cfg.enable_symm_mem: logger.warning( "Kimi-K3 DCP disables --enable-symm-mem due to decode CUDA " "graph correctness issues." ) overrides["enable_symm_mem"] = False - if server_args.speculative_algorithm == "DSPARK": + if cfg.speculative_algorithm == "DSPARK": from sglang.srt.speculative.ragged_verify import ( RaggedVerifyMode, read_ragged_verify_mode, @@ -630,7 +667,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: ) logger.info( "Kimi-K3 DCP with tokenspeed mla backend overrides KV cache dtype: " - f"{server_args.kv_cache_dtype!r} -> 'fp8_e4m3'." + f"{cfg.kv_cache_dtype!r} -> 'fp8_e4m3'." ) overrides.update( prefill_attention_backend="tokenspeed_mla", @@ -642,7 +679,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: f"Decode attention backend for Kimi-K3 DCP must be 'cutedsl_mla' or 'tokenspeed_mla', got {decode_backend!r}." ) - if server_args.dcp_replicate_q_proj is None: + if cfg.dcp_replicate_q_proj is None: logger.info("Kimi-K3 DCP enables replicated Q projection by default.") overrides["dcp_replicate_q_proj"] = True @@ -650,7 +687,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: dcp_comm_backend = "fi_a2a" if is_mnnvl_fabric_device() else "a2a" logger.info( "Kimi-K3 DCP selects communication backend on " - f"{device_name!r}: {server_args.dcp_comm_backend!r} -> " + f"{device_name!r}: {cfg.dcp_comm_backend!r} -> " f"{dcp_comm_backend!r}." ) overrides["dcp_comm_backend"] = dcp_comm_backend @@ -659,7 +696,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: if not (is_sm100_supported() and get_device_sm() in (100, 103)): return {} backends_unset = server_args.is_attention_backend_not_set() - if server_args.speculative_algorithm != "DSPARK": + if cfg.speculative_algorithm != "DSPARK": if not backends_unset: return {} logger.info( @@ -673,9 +710,9 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: # DSPARK: verify runs on the decode backend (mode=decode below), so this # picks the verify kernel -- mode=prefill routes it to flashinfer, which is # slow and syncs, while plain decode is cold under dspark. - q_len = server_args.speculative_num_draft_tokens or ( - server_args.speculative_dspark_block_size + 1 - if server_args.speculative_dspark_block_size is not None + q_len = cfg.speculative_num_draft_tokens or ( + cfg.speculative_dspark_block_size + 1 + if cfg.speculative_dspark_block_size is not None # Checkpoint auto-infer happens after overrides; K3 draft uses block 7. else 8 ) @@ -689,7 +726,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: # that still needs declaring -- else verify stays on the prefill backend, # whose host-side plan (flashinfer by default) forces a per-step D2H. _, backend = attention_backends_of(server_args) - if _dspark_verify_on_decode_backend(backend, q_len, server_args.kv_cache_dtype): + if _dspark_verify_on_decode_backend(backend, q_len, cfg.kv_cache_dtype): overrides["speculative_attention_mode"] = "decode" logger.info( "Kimi-K3 DSPARK on SM100/SM103: decode/verify attention backend " @@ -727,7 +764,8 @@ def _kimi_k3_moe_runner_overrides(server_args: Any, hf_config: Any) -> dict: # (M=bs) and the target-verify (M=bs*(gamma+1)) regimes on SM100/SM103. # SM107 uses the same packed-MXFP4 runner; leaving auto unresolved falls # back to BF16 weight materialization during model loading. - if server_args.moe_runner_backend != "auto": + cfg = resolving_view(server_args) + if cfg.moe_runner_backend != "auto": return {} if not (is_sm100_supported() and get_device_sm() in (100, 103, 107)): return {} @@ -757,6 +795,7 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict: writers), the kv-cache/split-backend defaults, the quant/moe block (read before it by _set_default_dsa_kv_cache_dtype) and the env writes stay in the branch.""" + cfg = resolving_view(server_args) from sglang.srt.configs.model_config import is_deepseek_dsa overrides: Dict[str, Any] = {} @@ -767,39 +806,39 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict: overrides["attention_backend"] = "dsa" logger.info("Use dsa attention backend for DeepSeek with DSA.") if not is_npu() and not is_xpu(): # CUDA or ROCm GPU - if server_args.enable_prefill_cp: + if cfg.enable_prefill_cp: logger.warning( "Context parallel feature is still under experiment. It has only been verified on Hopper platform." ) overrides["enable_dp_attention"] = True overrides["moe_dense_tp_size"] = 1 - if server_args.cp_strategy == "zigzag": + if cfg.cp_strategy == "zigzag": overrides["moe_a2a_backend"] = "deepep" - overrides["ep_size"] = server_args.tp_size + overrides["ep_size"] = cfg.tp_size logger.warning( "zigzag DSA CP requires moe_dense_tp_size=1, " "moe_a2a_backend=deepep, ep_size=tp_size, batch_size=1." ) else: assert ( - server_args.dp_size == 1 + cfg.dp_size == 1 ), "interleave DSA CP does not support DP attention." assert ( - server_args.tp_size <= 8 + cfg.tp_size <= 8 ), "Context parallel only supports single machine (tp_size <= 8). Cross-machine CP has precision issues." # Note(kpham-sgl): Keep attn_tp_size == 1 under DSA CP. # DSACPLayerCommunicator does not all-reduce attention-TP # partial o_proj outputs before replicated dense FFNs. - attn_cp_size = server_args.tp_size // server_args.dp_size + attn_cp_size = cfg.tp_size // cfg.dp_size overrides["attn_cp_size"] = attn_cp_size logger.warning( "Enabled DSA context parallel: " - f"strategy={server_args.cp_strategy}, dp_size={server_args.dp_size}, " + f"strategy={cfg.cp_strategy}, dp_size={cfg.dp_size}, " f"moe_dense_tp_size={overrides['moe_dense_tp_size']}, " - f"ep_size={overrides.get('ep_size', server_args.ep_size)}, tp_size={server_args.tp_size}, " + f"ep_size={overrides.get('ep_size', cfg.ep_size)}, tp_size={cfg.tp_size}, " f"attn_cp_size={attn_cp_size}, " - f"kv_cache_dtype={server_args.kv_cache_dtype}, " - f"moe_a2a_backend={overrides.get('moe_a2a_backend', server_args.moe_a2a_backend)}, " + f"kv_cache_dtype={cfg.kv_cache_dtype}, " + f"moe_a2a_backend={overrides.get('moe_a2a_backend', cfg.moe_a2a_backend)}, " f"cuda_graph_config[prefill].backend=disabled" ) @@ -826,9 +865,9 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict: # DeepSeek V3/R1/V3.1 if is_sm100_supported(): if ( - server_args.attention_backend is None - and server_args.prefill_attention_backend is None - and server_args.decode_attention_backend is None + cfg.attention_backend is None + and cfg.prefill_attention_backend is None + and cfg.decode_attention_backend is None ): overrides["attention_backend"] = "trtllm_mla" logger.info( @@ -836,7 +875,7 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict: ) # MLA prefill CP auto-config. Mirrors the NSA CP block above # (minus the in-seq/round-robin mode split, which MLA CP does not support) - if server_args.enable_prefill_cp and server_args.use_mla_backend(): + if cfg.enable_prefill_cp and server_args.use_mla_backend(): logger.warning( "MLA prefill context parallel is still experimental. " "Verified on Hopper with the fa3 backend." @@ -845,22 +884,22 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict: # TODO(kpham-sgl) Supports moe_dense_tp_size != 1. overrides["moe_dense_tp_size"] = 1 overrides["moe_a2a_backend"] = "deepep" - overrides["ep_size"] = server_args.tp_size + overrides["ep_size"] = cfg.tp_size logger.warning( "For MLA CP, we have the following restrictions: moe_dense_tp_size == 1, moe_a2a_backend == deepep, ep_size == tp_size, batch_size == 1" ) # FIXME(kpham-sgl): Keep attn_tp_size == 1 under MLA CP. # DSACPLayerCommunicator does not all-reduce attention-TP # partial o_proj outputs before replicated dense FFNs. - attn_cp_size = server_args.tp_size // server_args.dp_size + attn_cp_size = cfg.tp_size // cfg.dp_size overrides["attn_cp_size"] = attn_cp_size logger.warning( f"Enable Context Parallel opt for MLA, " - f"Setting dp_size == {server_args.dp_size} and " + f"Setting dp_size == {cfg.dp_size} and " f"attn_cp_size == {attn_cp_size}, " f"moe_dense_tp_size == {overrides['moe_dense_tp_size']}, " f"ep_size == {overrides['ep_size']}, " - f"tp_size == {server_args.tp_size}, " + f"tp_size == {cfg.tp_size}, " f"moe_a2a_backend {overrides['moe_a2a_backend']}, " f"cuda_graph_config[prefill].backend=disabled" ) @@ -870,8 +909,9 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict: # Keep in sync with MIMO_V2_MODEL_ARCHS (server_args.py / configs/hf_config.py). @_register_for("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM") def _mimo_v2_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) overrides: Dict[str, Any] = {} - if server_args.speculative_algorithm == "EAGLE": + if cfg.speculative_algorithm == "EAGLE": logger.info("Enable multi-layer EAGLE speculative decoding for MiMoV2 model.") overrides["enable_multi_layer_eagle"] = True @@ -879,7 +919,7 @@ def _mimo_v2_overrides(server_args: Any, hf_config: Any) -> dict: # slower at bs=1 decode. FP4 checkpoints use flashinfer_mxfp4 instead. if ( is_sm100_supported() - and server_args.moe_runner_backend == "auto" + and cfg.moe_runner_backend == "auto" and get_quantization_config(hf_config) == "fp8" ): overrides["moe_runner_backend"] = "flashinfer_trtllm" @@ -889,13 +929,14 @@ def _mimo_v2_overrides(server_args: Any, hf_config: Any) -> dict: @_register_for("MiniMaxM2ForCausalLM") def _minimax_m2_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) overrides = {"enable_tf32_matmul": True} logger.info( "Enable TF32 matmul for MiniMaxM2ForCausalLM model to improve gate gemm performance." ) if ( is_sm100_supported() - and server_args.moe_runner_backend == "auto" + and cfg.moe_runner_backend == "auto" and server_args.get_model_config().quantization == "modelopt_fp4" ): overrides["moe_runner_backend"] = "flashinfer_trtllm_routed" @@ -909,10 +950,11 @@ def _minimax_m2_overrides(server_args: Any, hf_config: Any) -> dict: @_register_for("MiniMaxM3SparseForCausalLM", "MiniMaxM3SparseForConditionalGeneration") def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) overrides: Dict[str, Any] = {} quant_method = get_quantization_config(hf_config) - quant_resolved = server_args.quantization + quant_resolved = cfg.quantization if ( quant_resolved is None and not server_args._quantization_explicitly_unset @@ -924,16 +966,12 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict: if is_hip(): if server_args.is_attention_backend_not_set(): overrides["attention_backend"] = "triton" - if server_args.moe_runner_backend == "auto" and quant_resolved == "mxfp8": + if cfg.moe_runner_backend == "auto" and quant_resolved == "mxfp8": overrides["moe_runner_backend"] = "triton" if not envs.USE_ROCM_AITER_ROPE_BACKEND.is_set(): envs.USE_ROCM_AITER_ROPE_BACKEND.set("0") - aiter_fusion_resolved = server_args.enable_aiter_allreduce_fusion - if ( - server_args.ep_size > 1 - and server_args.moe_a2a_backend == "none" - and aiter_fusion_resolved - ): + aiter_fusion_resolved = cfg.enable_aiter_allreduce_fusion + if cfg.ep_size > 1 and cfg.moe_a2a_backend == "none" and aiter_fusion_resolved: logger.warning( "Disable --enable-aiter-allreduce-fusion for MiniMax-M3 " "standard EP on ROCm because the deferred fused all-reduce " @@ -951,7 +989,7 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict: elif is_sm100_supported(): if server_args.is_attention_backend_not_set(): if ( - server_args.kv_cache_dtype == "fp8_e4m3" + cfg.kv_cache_dtype == "fp8_e4m3" and not envs.SGLANG_DISABLE_M3_FP8_ATTN_GEMM.get() ): # fp8 attention GEMMs activate whenever possible @@ -962,42 +1000,36 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict: overrides["attention_backend"] = "trtllm_mha" else: overrides["attention_backend"] = "fa4" - backend_resolved = overrides.get( - "attention_backend", server_args.attention_backend - ) - page_resolved = server_args.page_size + backend_resolved = overrides.get("attention_backend", cfg.attention_backend) + page_resolved = cfg.page_size # fa4 (fmha_sm100) and trtllm_mha both allow the page_size == 128 # sparse block MSA needs (trtllm_mha via trtllm-gen's dynamic # tokens-per-page kernels). if page_resolved is None and backend_resolved in ("fa4", "trtllm_mha"): overrides["page_size"] = 128 page_resolved = 128 - if server_args.moe_runner_backend == "auto" and quant_resolved == "mxfp8": + if cfg.moe_runner_backend == "auto" and quant_resolved == "mxfp8": overrides["moe_runner_backend"] = "deep_gemm" - elif ( - server_args.moe_runner_backend == "auto" - and quant_resolved == "modelopt_mixed" - ): + elif cfg.moe_runner_backend == "auto" and quant_resolved == "modelopt_mixed": overrides["moe_runner_backend"] = "flashinfer_trtllm_routed" logger.info( "MiniMax-M3 on SM100: attention_backend=" - f"{overrides.get('attention_backend', server_args.attention_backend)}, page_size={page_resolved}, " - f"moe_runner_backend={overrides.get('moe_runner_backend', server_args.moe_runner_backend)}." + f"{overrides.get('attention_backend', cfg.attention_backend)}, page_size={page_resolved}, " + f"moe_runner_backend={overrides.get('moe_runner_backend', cfg.moe_runner_backend)}." ) elif is_sm90_supported(): if server_args.is_attention_backend_not_set(): overrides["attention_backend"] = "fa3" - page_resolved = server_args.page_size + page_resolved = cfg.page_size if ( page_resolved is None - and overrides.get("attention_backend", server_args.attention_backend) - == "fa3" + and overrides.get("attention_backend", cfg.attention_backend) == "fa3" ): overrides["page_size"] = 128 page_resolved = 128 logger.info( "MiniMax-M3 on Hopper: attention_backend=" - f"{overrides.get('attention_backend', server_args.attention_backend)}, page_size={page_resolved} " + f"{overrides.get('attention_backend', cfg.attention_backend)}, page_size={page_resolved} " "(MSA is SM100-only; sparse attention runs on the Triton path)." ) @@ -1008,7 +1040,7 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict: # silently dispatch the e4m3 kernel, so e5m2 stays on the widening Triton # path), log when the fp8 GEMM mode is active, and log when the # SGLANG_DISABLE_M3_FP8_ATTN_GEMM kill switch suppresses it. - if server_args.kv_cache_dtype == "fp8_e5m2": + if cfg.kv_cache_dtype == "fp8_e5m2": logger.warning( "MiniMax-M3 with kv_cache_dtype fp8_e5m2: fp8 attention GEMMs stay " "DISABLED (fmha_sm100's variant lookup would silently dispatch the " @@ -1016,9 +1048,8 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict: "Triton path. Use --kv-cache-dtype fp8_e4m3 for fp8 attention GEMMs." ) elif ( - server_args.kv_cache_dtype == "fp8_e4m3" - and overrides.get("attention_backend", server_args.attention_backend) - == "trtllm_mha" + cfg.kv_cache_dtype == "fp8_e4m3" + and overrides.get("attention_backend", cfg.attention_backend) == "trtllm_mha" and is_sm100_supported() ): if envs.SGLANG_DISABLE_M3_FP8_ATTN_GEMM.get(): @@ -1036,9 +1067,7 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict: "force the pre-fp8 numerics." ) - moe_runner_resolved = overrides.get( - "moe_runner_backend", server_args.moe_runner_backend - ) + moe_runner_resolved = overrides.get("moe_runner_backend", cfg.moe_runner_backend) if quant_resolved is None and moe_runner_resolved in ("auto", "deep_gemm"): if moe_runner_resolved == "deep_gemm": logger.warning( @@ -1078,6 +1107,7 @@ def _exaone_overrides(server_args: Any, hf_config: Any) -> dict: @_register_for("GptOssForCausalLM") def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) overrides: Dict[str, Any] = {} # Set attention backend for GPT-OSS if server_args.is_attention_backend_not_set(): @@ -1100,14 +1130,14 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict: # Check for bf16 dtype on Intel XPU. Reads the pristine dtype request, # which equals the legacy mid-branch read: dtype had no earlier writer # for this arch. - if server_args.dtype == "auto": + if cfg.dtype == "auto": logger.warning( "GptOssForCausalLM on Intel XPU currently supports bfloat16 dtype only" ) - elif server_args.dtype not in ["bfloat16"]: + elif cfg.dtype not in ["bfloat16"]: raise NotImplementedError( f"GptOssForCausalLM on Intel XPU only supports bfloat16 dtype, " - f"but got '{server_args.dtype}'. Please use --dtype bfloat16 or remove --dtype to use auto." + f"but got '{cfg.dtype}'. Please use --dtype bfloat16 or remove --dtype to use auto." ) quantization_config = getattr(hf_config, "quantization_config", None) is_mxfp4_quant_format = ( @@ -1117,7 +1147,7 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict: if is_mxfp4_quant_format: # use bf16 for mxfp4 triton kernels overrides["dtype"] = "bfloat16" - if server_args.moe_runner_backend == "auto": + if cfg.moe_runner_backend == "auto": if is_sm100_supported() and is_mxfp4_quant_format: overrides["moe_runner_backend"] = "flashinfer_mxfp4" @@ -1155,9 +1185,9 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict: "Detected MUSA with SGLANG_DEEPEP_BF16_DISPATCH for bf16 model, using deep_gemm kernel." ) elif ( - server_args.ep_size == 1 + cfg.ep_size == 1 and is_triton_kernels_available() - and server_args.quantization is None + and cfg.quantization is None and not (is_cpu() and cpu_has_amx_support()) ): # The triton_kernels package segfaults on Blackwell (B200) @@ -1179,18 +1209,19 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict: # Keep in sync with LLAMA4_MODEL_ARCHS (server_args.py). @_register_for("Llama4ForConditionalGeneration", "Llama4ForCausalLM") def _llama4_overrides(server_args: Any, hf_config: Any) -> dict: - if server_args.device == "cpu": + cfg = resolving_view(server_args) + if cfg.device == "cpu": return {} overrides: Dict[str, Any] = {} # Auto-select attention backend for Llama4 if not specified - if server_args.attention_backend is None: + if cfg.attention_backend is None: if is_sm100_supported(): backend, platform = "trtllm_mha", "sm100" elif is_sm90_supported(): backend, platform = "fa3", "sm90" elif is_hip(): backend, platform = "aiter", "hip" - elif server_args.device == "xpu": + elif cfg.device == "xpu": backend, platform = "intel_xpu", "xpu" else: backend, platform = "triton", "other platforms" @@ -1198,8 +1229,8 @@ def _llama4_overrides(server_args: Any, hf_config: Any) -> dict: f"Use {backend} as attention backend on {platform} for Llama4 model" ) overrides["attention_backend"] = backend - if is_sm100_supported() and server_args.moe_runner_backend == "auto": - if server_args.quantization in {"fp8", "modelopt_fp8"}: + if is_sm100_supported() and cfg.moe_runner_backend == "auto": + if cfg.quantization in {"fp8", "modelopt_fp8"}: overrides["moe_runner_backend"] = "flashinfer_trtllm" logger.info( "Use flashinfer_trtllm as MoE runner backend on SM100 for Llama4" @@ -1213,6 +1244,7 @@ def _llama4_overrides(server_args: Any, hf_config: Any) -> dict: "Gemma4UnifiedForConditionalGeneration", ) def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) overrides: Dict[str, Any] = {} default_attention_backend = "trtllm_mha" if is_sm100_supported() else "triton" if server_args.is_attention_backend_not_set(): @@ -1223,9 +1255,9 @@ def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict: # If only one split backend is set, keep the other side on a # Gemma4-compatible fallback instead of letting generic backend selection # choose an unsupported backend later. - elif server_args.attention_backend is None: + elif cfg.attention_backend is None: overrides["attention_backend"] = default_attention_backend - if is_sm100_supported() and server_args.moe_runner_backend == "auto": + if is_sm100_supported() and cfg.moe_runner_backend == "auto": if server_args.get_model_config().quantization == "modelopt_fp4": overrides["quantization"] = "modelopt_fp4" overrides["moe_runner_backend"] = "flashinfer_trtllm" @@ -1306,7 +1338,8 @@ def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict: @_register_for("MiniCPMV4_6ForConditionalGeneration") def _minicpm_v4_6_overrides(server_args: Any, hf_config: Any) -> dict: - if is_sm100_supported() and server_args.attention_backend is None: + cfg = resolving_view(server_args) + if is_sm100_supported() and cfg.attention_backend is None: return {"attention_backend": "triton"} return {} @@ -1315,24 +1348,27 @@ def _minicpm_v4_6_overrides(server_args: Any, hf_config: Any) -> dict: "FalconH1ForCausalLM", "JetNemotronForCausalLM", "JetVLMForConditionalGeneration" ) def _falcon_h1_jet_overrides(server_args: Any, hf_config: Any) -> dict: - if is_sm100_supported() and server_args.attention_backend is None: + cfg = resolving_view(server_args) + if is_sm100_supported() and cfg.attention_backend is None: return {"attention_backend": "triton"} return {} @_register_for("GraniteMoeHybridForCausalLM") def _granite_moe_hybrid_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) has_mamba = any( layer_type == "mamba" for layer_type in getattr(hf_config, "layer_types", []) ) - if has_mamba and is_sm100_supported() and server_args.attention_backend is None: + if has_mamba and is_sm100_supported() and cfg.attention_backend is None: return {"attention_backend": "flashinfer"} return {} @_register_for("Lfm2ForCausalLM", "Lfm2MoeForCausalLM") def _lfm2_overrides(server_args: Any, hf_config: Any) -> dict: - if is_sm100_supported() and server_args.attention_backend is None: + cfg = resolving_view(server_args) + if is_sm100_supported() and cfg.attention_backend is None: return {"attention_backend": "flashinfer"} return {} @@ -1343,13 +1379,14 @@ def _deepseek_v4_overrides(server_args: Any, hf_config: Any) -> dict: arg_groups/deepseek_v4_hook.py). The kv-cache dtype and NPU split-backend writes, the max_running_requests fill and the validations stay in the hook at its legacy slot.""" + cfg = resolving_view(server_args) from sglang.srt.server_args import ServerArgs model_arch = hf_config.architectures[0] overrides: Dict[str, Any] = {"attention_backend": "dsv4"} page_size = 256 - if server_args.device == "npu": + if cfg.device == "npu": # NPU keeps the device-aware "dsv4" backend (the registry routes it to # the Ascend V4 subclass); only the pool geometry / dtype differ. # set_default_server_args() pins all three backends to "ascend" for @@ -1363,11 +1400,11 @@ def _deepseek_v4_overrides(server_args: Any, hf_config: Any) -> dict: f"Use dsv4 attention backend for {model_arch}, setting page_size to {page_size}." ) - if server_args.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio: + if cfg.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio: overrides["swa_full_tokens_ratio"] = 0.1 logger.info(f"Setting swa_full_tokens_ratio to 0.1 for {model_arch}.") - if server_args.moe_runner_backend == "auto": + if cfg.moe_runner_backend == "auto": model_config = server_args.get_model_config() # nvidia/DeepSeek-V4-Pro-NVFP4 uses the routed TRT-LLM runner. if model_config.nvfp4_moe_meta is not None: @@ -1377,9 +1414,9 @@ def _deepseek_v4_overrides(server_args: Any, hf_config: Any) -> dict: f"{model_arch} hybrid FP8+NVFP4 checkpoint." ) elif ( - server_args.device == "cuda" + cfg.device == "cuda" and not is_hip() - and server_args.moe_a2a_backend == "none" + and cfg.moe_a2a_backend == "none" and not envs.SGLANG_DSV4_FP4_DEQUANT.get() and model_config.is_fp4_experts and (is_sm90_supported() or is_sm100_supported() or is_sm120_supported()) @@ -1407,6 +1444,7 @@ def _inkling_overrides(server_args: Any, hf_config: Any) -> dict: prefill.backend, and an explicit --cuda-graph-backend-prefill / --disable-prefill-cuda-graph still wins. The unified-radix env write follows the MiniMax-M3 handler precedent (env is not a resolvable server-arg).""" + cfg = resolving_view(server_args) from sglang.srt.server_args import ServerArgs overrides: Dict[str, Any] = {} @@ -1415,14 +1453,14 @@ def _inkling_overrides(server_args: Any, hf_config: Any) -> dict: # cuda_graph_backend_prefill declared here lands too late (the breakable # default would already have been auto-disabled for this multimodal arch). # It is set inline before _handle_cuda_graph_config instead. - if server_args.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio: + if cfg.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio: overrides["swa_full_tokens_ratio"] = 0.1 - if server_args.mamba_full_memory_ratio == ServerArgs.mamba_full_memory_ratio: + if cfg.mamba_full_memory_ratio == ServerArgs.mamba_full_memory_ratio: overrides["mamba_full_memory_ratio"] = 0.1 # Inkling requires the extra-buffer mamba strategy (inkling.py asserts # enable_mamba_extra_buffer()); the generic "auto" resolution does not cover # Inkling, so pin it here. Yields to an explicit --mamba-scheduler-strategy. - if server_args.mamba_radix_cache_strategy == ServerArgs.mamba_radix_cache_strategy: + if cfg.mamba_radix_cache_strategy == ServerArgs.mamba_radix_cache_strategy: overrides["mamba_radix_cache_strategy"] = "extra_buffer" # Inkling attention runs only on the fa4 (Blackwell) or triton backends -- # models/inkling_common/attn.py asserts attention_backend in {fa4, triton}. @@ -1447,6 +1485,7 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict: """NemotronH quantization / MoE runner / attention backend defaults (absorbed from the retired arg_groups/nemotron_h_hook.py; the mamba radix cache handling and the triton-backend assert stay in the arch branch).""" + cfg = resolving_view(server_args) model_arch = hf_config.architectures[0] model_config = server_args.get_model_config() overrides: Dict[str, Any] = {} @@ -1457,7 +1496,7 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict: "modelopt_fp4", "modelopt_mixed", ] - quantization = server_args.quantization + quantization = cfg.quantization if is_modelopt: assert model_config.hf_config.mlp_hidden_act == "relu2" if model_config.quantization == "modelopt": @@ -1482,22 +1521,22 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict: ) if has_w4a16_moe_layers: - if server_args.moe_a2a_backend != "none": + if cfg.moe_a2a_backend != "none": raise ValueError("W4A16_NVFP4 MoE layers require --moe-a2a-backend=none.") - if server_args.moe_runner_backend not in ("auto", "marlin"): + if cfg.moe_runner_backend not in ("auto", "marlin"): raise ValueError( "W4A16_NVFP4 MoE layers require --moe-runner-backend=marlin." ) - if server_args.moe_runner_backend == "auto": + if cfg.moe_runner_backend == "auto": overrides["moe_runner_backend"] = "marlin" logger.info( "Use marlin as MoE runner backend for " f"{model_arch} with W4A16_NVFP4 MoE layers" ) elif (is_modelopt or model_config.quantization is None) and ( - server_args.moe_runner_backend == "auto" + cfg.moe_runner_backend == "auto" ): - if is_sm100_supported() and server_args.moe_a2a_backend == "none": + if is_sm100_supported() and cfg.moe_a2a_backend == "none": overrides["moe_runner_backend"] = "flashinfer_trtllm" logger.info( f"Use flashinfer_trtllm as MoE runner backend on sm100 for {model_arch}" @@ -1518,27 +1557,27 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict: else: overrides["moe_runner_backend"] = "flashinfer_cutlass" - if is_blackwell_supported() and server_args.is_attention_backend_not_set(): - if server_args.speculative_algorithm is not None: - speculative_algorithm = server_args.speculative_algorithm.upper() - if is_sm100_supported() and server_args.speculative_eagle_topk in ( + if is_blackwell_supported() and cfg.is_attention_backend_not_set(): + if cfg.speculative_algorithm is not None: + speculative_algorithm = cfg.speculative_algorithm.upper() + if is_sm100_supported() and cfg.speculative_eagle_topk in ( None, 1, ): overrides["attention_backend"] = "trtllm_mha" - if server_args.page_size is None: + if cfg.page_size is None: overrides["page_size"] = 64 - if server_args.mamba_radix_cache_strategy == "auto": + if cfg.mamba_radix_cache_strategy == "auto": overrides["mamba_radix_cache_strategy"] = "extra_buffer" if ( - server_args.speculative_draft_attention_backend is None + cfg.speculative_draft_attention_backend is None and speculative_algorithm in ("EAGLE", "NEXTN", "DSPARK") ): overrides["speculative_draft_attention_backend"] = "trtllm_mha" else: overrides["attention_backend"] = "triton" if ( - server_args.speculative_draft_attention_backend is None + cfg.speculative_draft_attention_backend is None and speculative_algorithm in ("EAGLE", "NEXTN", "DFLASH", "DSPARK") ): overrides["speculative_draft_attention_backend"] = "flashinfer" @@ -1555,7 +1594,8 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict: "Qwen3_5ForConditionalGeneration", ) def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict: - if not is_sm100_supported() or server_args.attention_backend is not None: + cfg = resolving_view(server_args) + if not is_sm100_supported() or cfg.attention_backend is not None: return {} sm100_default_attn_backend = "triton" # trtllm_mha requires speculative_eagle_topk == 1 and page_size > 1. @@ -1572,8 +1612,8 @@ def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict: # already-written field here). if default_attn_backend == "trtllm_mha" and not ( not mamba_extra_buffer_of(resolved_view(server_args)) - and not server_args.disable_radix_cache - and server_args.speculative_algorithm is None + and not cfg.disable_radix_cache + and cfg.speculative_algorithm is None ): sm100_default_attn_backend = "trtllm_mha" return { @@ -1585,7 +1625,8 @@ def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict: @_register_for("InternS2MobiusForConditionalGeneration") def _interns2_mobius_baseline_overrides(server_args: Any, hf_config: Any) -> dict: """Select the only MoE runner validated for the 2,560-expert baseline.""" - if server_args.moe_runner_backend == "auto": + cfg = resolving_view(server_args) + if cfg.moe_runner_backend == "auto": return {"moe_runner_backend": "triton_kernel"} return {} @@ -1593,11 +1634,8 @@ def _interns2_mobius_baseline_overrides(server_args: Any, hf_config: Any) -> dic @_register_for("Qwen3VLForConditionalGeneration") def _qwen3vl_overrides(server_args: Any, hf_config: Any) -> dict: - if ( - is_hip() - and envs.SGLANG_USE_AITER_UNIFIED_ATTN.get() - and server_args.page_size is None - ): + cfg = resolving_view(server_args) + if is_hip() and envs.SGLANG_USE_AITER_UNIFIED_ATTN.get() and cfg.page_size is None: logger.info( "Setting page_size=16 for aiter unified attention on Qwen3VLForConditionalGeneration." ) @@ -1614,10 +1652,11 @@ def _qwen3vl_overrides(server_args: Any, hf_config: Any) -> dict: "Qwen3_5ForConditionalGeneration", ) def _qwen3_moe_family_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) overrides: Dict[str, Any] = {} if is_sm100_supported(): quant_method = get_quantization_config(hf_config) - quantization = server_args.quantization + quantization = cfg.quantization if ( quantization is None and not server_args._quantization_explicitly_unset @@ -1627,8 +1666,8 @@ def _qwen3_moe_family_overrides(server_args: Any, hf_config: Any) -> dict: quantization = quant_method if ( (quantization in ("fp8", "modelopt_fp4") or quantization is None) - and server_args.moe_a2a_backend == "none" - and server_args.moe_runner_backend == "auto" + and cfg.moe_a2a_backend == "none" + and cfg.moe_runner_backend == "auto" ): overrides["moe_runner_backend"] = "flashinfer_trtllm" logger.info( @@ -1640,6 +1679,7 @@ def _qwen3_moe_family_overrides(server_args: Any, hf_config: Any) -> dict: @_register_for("Glm4MoeForCausalLM") def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) overrides: Dict[str, Any] = {} if is_sm100_supported(): quantization_config = getattr(hf_config, "quantization_config", None) @@ -1648,7 +1688,7 @@ def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict: if quantization_config is not None else None ) - quantization = server_args.quantization + quantization = cfg.quantization if ( quantization is None and not server_args._quantization_explicitly_unset @@ -1658,8 +1698,8 @@ def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict: quantization = quant_method if ( quantization in {"modelopt_fp4", None} - and server_args.moe_a2a_backend == "none" - and server_args.moe_runner_backend == "auto" + and cfg.moe_a2a_backend == "none" + and cfg.moe_runner_backend == "auto" ): overrides["moe_runner_backend"] = "flashinfer_trtllm" logger.info( @@ -1674,13 +1714,14 @@ def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict: @_register_for("Olmo2ForCausalLM") def _olmo2_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) overrides: Dict[str, Any] = {} # FIXME: https://github.com/sgl-project/sglang/pull/7367 is not compatible with Olmo3 model. logger.warning( f"Disabling hybrid SWA memory for {hf_config.architectures[0]} as it is not yet supported." ) overrides["disable_hybrid_swa_memory"] = True - if server_args.attention_backend is None: + if cfg.attention_backend is None: if is_cuda() and is_sm100_supported(): overrides["attention_backend"] = "trtllm_mha" elif is_cuda() and get_device_sm() >= 80: @@ -1695,6 +1736,7 @@ def _olmo2_overrides(server_args: Any, hf_config: Any) -> dict: or "Step3p7ForConditionalGeneration" in arch ) def _step3p_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) overrides: Dict[str, Any] = {} if server_args.is_attention_backend_not_set(): if is_blackwell_supported(): @@ -1703,12 +1745,12 @@ def _step3p_overrides(server_args: Any, hf_config: Any) -> dict: elif is_sm90_supported(): logger.info("Auto-select fa3 attention backend for Step3p7 on Hopper.") overrides["attention_backend"] = "fa3" - if server_args.speculative_algorithm == "EAGLE": + if cfg.speculative_algorithm == "EAGLE": logger.info( "Enable multi-layer EAGLE speculative decoding for Step3p5ForCausalLM model." ) overrides["enable_multi_layer_eagle"] = True - if server_args.enable_hierarchical_cache: + if cfg.enable_hierarchical_cache: logger.warning( "Reset swa_full_tokens_ratio to 1.0 for Step3p5ForCausalLM model with hierarchical cache" ) @@ -2170,7 +2212,8 @@ def _deepseek_v4_kv_cache_dtype(view: Any) -> dict: @_register_for("MuseGlimmerForConditionalGeneration", "MuseGlimmerForCausalLM") def _muse_glimmer_fp4_gemm_runner_overrides(server_args: Any, hf_config: Any) -> dict: - if is_sm120_supported() and server_args.fp4_gemm_runner_backend == "auto": + cfg = resolving_view(server_args) + if is_sm120_supported() and cfg.fp4_gemm_runner_backend == "auto": logger.info("Use marlin as FP4 GEMM runner backend on SM120 for Muse Glimmer") return {"fp4_gemm_runner_backend": "marlin"} return {} diff --git a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py index ff2257d61eb3..7d45061baed1 100644 --- a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py +++ b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py @@ -5,7 +5,10 @@ import os from typing import TYPE_CHECKING -from sglang.srt.arg_groups.overrides import declare_resolution +from sglang.srt.arg_groups.overrides import ( + declare_resolution, + resolving_view, +) from sglang.srt.environ import envs if TYPE_CHECKING: @@ -16,10 +19,11 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None: """Validate and normalize PD-disaggregation server args.""" + cfg = resolving_view(server_args) # "mooncake_tcp" is mooncake with the TCP transport forced: set MC_FORCE_TCP # so mooncake installs TcpTransport instead of RDMA, rewrite the backend to # mooncake, and skip RDMA HCA selection. Must run before backend-name checks. - if server_args.disaggregation_transfer_backend == "mooncake_tcp": + if cfg.disaggregation_transfer_backend == "mooncake_tcp": os.environ.setdefault("MC_FORCE_TCP", "1") declare_resolution( server_args, @@ -36,17 +40,17 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None: "with MC_FORCE_TCP=1 (TCP transport, no RDMA)" ) - if server_args.disaggregation_mode == "prefill" and server_args.dcp_size > 1: + if cfg.disaggregation_mode == "prefill" and cfg.dcp_size > 1: logger.warning( "DCP on a PD prefill server is supported when prefill and decode " "use the same DCP layout, but it usually adds communication " "overhead without improving prefill performance." ) - if server_args.disaggregation_mode == "decode" and server_args.dcp_size > 1: + if cfg.disaggregation_mode == "decode" and cfg.dcp_size > 1: # Fake transfer moves no KV and is only used for synthetic decode # benchmarks, so it does not need the DCP relayout from Mooncake/NIXL. - if server_args.disaggregation_transfer_backend not in ( + if cfg.disaggregation_transfer_backend not in ( "mooncake", "nixl", "fake", @@ -54,36 +58,36 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None: raise ValueError( "PD decode DCP requires --disaggregation-transfer-backend " "mooncake, nixl, or fake for synthetic benchmarking, got " - f"{server_args.disaggregation_transfer_backend!r}." + f"{cfg.disaggregation_transfer_backend!r}." ) - if server_args.disaggregation_decode_enable_radix_cache: + if cfg.disaggregation_decode_enable_radix_cache: raise ValueError( "PD decode DCP currently requires chunk cache; " "--disaggregation-decode-enable-radix-cache is not supported." ) - if server_args.enable_hierarchical_cache: + if cfg.enable_hierarchical_cache: raise ValueError( "PD decode DCP currently requires chunk cache; " "--enable-hierarchical-cache is not supported." ) - if server_args.disaggregation_mode == "decode": - if server_args.disaggregation_decode_enable_radix_cache: - if server_args.enable_hisparse: + if cfg.disaggregation_mode == "decode": + if cfg.disaggregation_decode_enable_radix_cache: + if cfg.enable_hisparse: raise ValueError( "--disaggregation-decode-enable-radix-cache is incompatible " "with --enable-hisparse" ) - if server_args.disaggregation_transfer_backend == "fake": + if cfg.disaggregation_transfer_backend == "fake": raise ValueError( "--disaggregation-decode-enable-radix-cache is incompatible " "with --disaggregation-transfer-backend fake" ) - if server_args.speculative_algorithm is not None: + if cfg.speculative_algorithm is not None: raise ValueError( "--disaggregation-decode-enable-radix-cache is incompatible " "with speculative decoding " - f"(--speculative-algorithm {server_args.speculative_algorithm})" + f"(--speculative-algorithm {cfg.speculative_algorithm})" ) from sglang.srt.arg_groups.overrides import resolved_view @@ -110,12 +114,10 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None: # in-transfer (being-received-from-prefill) requests, on top of the # max_running_requests-derived pool. Large batches get none; small # per-worker batches reserve 2x the batch as cheap overlap headroom. - if server_args.disaggregation_decode_extra_slots is None: + if cfg.disaggregation_decode_extra_slots is None: extra_slots = 0 - if server_args.max_running_requests is not None: - per_worker = server_args.max_running_requests // max( - 1, server_args.dp_size - ) + if cfg.max_running_requests is not None: + per_worker = cfg.max_running_requests // max(1, cfg.dp_size) if per_worker <= 32: extra_slots = per_worker * 2 declare_resolution( @@ -124,23 +126,23 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None: disaggregation_decode_extra_slots=extra_slots, ) - elif server_args.disaggregation_mode == "prefill": + elif cfg.disaggregation_mode == "prefill": assert ( - server_args.disaggregation_transfer_backend != "fake" + cfg.disaggregation_transfer_backend != "fake" ), "Prefill server does not support 'fake' as the transfer backend" if envs.SGLANG_RUST_SERVER.get(): _alias_bootstrap_port_to_api_port(server_args) - if server_args.disaggregation_mode in ("prefill", "decode"): + if cfg.disaggregation_mode in ("prefill", "decode"): if ( envs.SGLANG_DISAGG_STAGING_BUFFER.get() - and server_args.disaggregation_transfer_backend not in ("mooncake", "nixl") + and cfg.disaggregation_transfer_backend not in ("mooncake", "nixl") ): raise ValueError( f"SGLANG_DISAGG_STAGING_BUFFER requires " f"disaggregation_transfer_backend='mooncake' or 'nixl', " - f"got '{server_args.disaggregation_transfer_backend}'." + f"got '{cfg.disaggregation_transfer_backend}'." ) @@ -151,31 +153,32 @@ def _alias_bootstrap_port_to_api_port(server_args: ServerArgs) -> None: field and agrees automatically. Decode is untouched: there the field names the PREFILL side's bootstrap port and must stay as the operator set it. """ + cfg = resolving_view(server_args) default_port = next( f.default for f in dataclasses.fields(server_args) if f.name == "disaggregation_bootstrap_port" ) - if server_args.disaggregation_bootstrap_port not in ( + if cfg.disaggregation_bootstrap_port not in ( default_port, - server_args.port, + cfg.port, ): raise ValueError( "SGLANG_RUST_SERVER serves the PD KV bootstrap registry on the api " "port itself; --disaggregation-bootstrap-port " - f"{server_args.disaggregation_bootstrap_port} conflicts with --port " - f"{server_args.port}. Drop --disaggregation-bootstrap-port (decode " + f"{cfg.disaggregation_bootstrap_port} conflicts with --port " + f"{cfg.port}. Drop --disaggregation-bootstrap-port (decode " "nodes and the PD router must then target the prefill api port)." ) - if server_args.disaggregation_bootstrap_port != server_args.port: + if cfg.disaggregation_bootstrap_port != cfg.port: logger.info( "SGLANG_RUST_SERVER: KV bootstrap registry is served on the api " "port; disaggregation_bootstrap_port %d -> %d", - server_args.disaggregation_bootstrap_port, - server_args.port, + cfg.disaggregation_bootstrap_port, + cfg.port, ) declare_resolution( server_args, "_alias_bootstrap_port_to_api_port", - disaggregation_bootstrap_port=server_args.port, + disaggregation_bootstrap_port=cfg.port, ) diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py index 905dbcb546f1..05cb2cbe517a 100644 --- a/python/sglang/srt/arg_groups/speculative_hook.py +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -8,6 +8,7 @@ from sglang.srt.arg_groups.overrides import ( declare_direct_writes, declare_resolution, + resolving_view, ) if TYPE_CHECKING: @@ -17,7 +18,8 @@ def _disable_overlap_schedule_for_cpu(server_args: ServerArgs) -> None: - if server_args.device != "cpu" or server_args.disable_overlap_schedule: + cfg = resolving_view(server_args) + if cfg.device != "cpu" or cfg.disable_overlap_schedule: return declare_resolution( @@ -71,9 +73,10 @@ def _resolve_speculative_algorithm_alias( def handle_speculative_decoding(server_args: ServerArgs) -> None: + cfg = resolving_view(server_args) if ( - server_args.speculative_draft_model_path is not None - and server_args.speculative_draft_model_revision is None + cfg.speculative_draft_model_path is not None + and cfg.speculative_draft_model_revision is None ): declare_resolution( server_args, @@ -90,11 +93,11 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None: run_post_process_pass(server_args, _speculative_moe_runner_default) - if server_args.speculative_algorithm is not None: + if cfg.speculative_algorithm is not None: declare_resolution( server_args, "handle_speculative_decoding", - speculative_algorithm=server_args.speculative_algorithm.upper(), + speculative_algorithm=cfg.speculative_algorithm.upper(), ) # Removal notice for the retired env var; raw os.getenv on purpose -- the @@ -108,7 +111,7 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None: kwargs = {} - override_config_file = server_args.decrypted_draft_config_file + override_config_file = cfg.decrypted_draft_config_file if override_config_file and override_config_file.strip(): kwargs["_configuration_file"] = override_config_file.strip() @@ -116,17 +119,17 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None: server_args, "handle_speculative_decoding", speculative_algorithm=_resolve_speculative_algorithm_alias( - server_args.speculative_algorithm, - server_args.speculative_draft_model_path, - trust_remote_code=server_args.trust_remote_code, + cfg.speculative_algorithm, + cfg.speculative_draft_model_path, + trust_remote_code=cfg.trust_remote_code, kwargs=kwargs, ), ) # Validate --speculative-draft-window-size once, regardless of algorithm. # Consumed by DFLASH (compact draft KV cache) and Llama EAGLE-3 (drafter attention SWA). - if server_args.speculative_draft_window_size is not None: - window_size = int(server_args.speculative_draft_window_size) + if cfg.speculative_draft_window_size is not None: + window_size = int(cfg.speculative_draft_window_size) if window_size <= 0: raise ValueError( f"--speculative-draft-window-size must be positive, got {window_size}." @@ -136,19 +139,19 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None: "handle_speculative_decoding", speculative_draft_window_size=window_size, ) - if server_args.speculative_algorithm not in ("EAGLE3", "DFLASH"): + if cfg.speculative_algorithm not in ("EAGLE3", "DFLASH"): logger.warning( "--speculative-draft-window-size has no effect with " "speculative_algorithm=%s (honored by Llama EAGLE-3 and DFLASH only).", - server_args.speculative_algorithm, + cfg.speculative_algorithm, ) algo = None - if server_args.speculative_algorithm is not None: + if cfg.speculative_algorithm is not None: from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_registry import CustomSpecAlgo - algo = SpeculativeAlgorithm.from_string(server_args.speculative_algorithm) + algo = SpeculativeAlgorithm.from_string(cfg.speculative_algorithm) # TODO: move the per-algorithm validation below into spec module hooks. if isinstance(algo, CustomSpecAlgo) and algo.validate_server_args is not None: @@ -158,15 +161,15 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None: algo.validate_server_args, ) - if server_args.speculative_skip_dp_mlp_sync: - assert server_args.speculative_algorithm == "EAGLE", ( + if cfg.speculative_skip_dp_mlp_sync: + assert cfg.speculative_algorithm == "EAGLE", ( "--speculative-skip-dp-mlp-sync is only supported with " - f"speculative_algorithm == EAGLE, got {server_args.speculative_algorithm}." + f"speculative_algorithm == EAGLE, got {cfg.speculative_algorithm}." ) - if server_args.speculative_adaptive: + if cfg.speculative_adaptive: _maybe_disable_adaptive(server_args) - if server_args.speculative_adaptive: + if cfg.speculative_adaptive: _init_adaptive_speculative_params(server_args) if algo is not None: @@ -180,9 +183,10 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None: def _handle_dflash(server_args: ServerArgs) -> None: + cfg = resolving_view(server_args) from sglang.srt.arg_groups.overrides import resolved_view - if not (server_args.device.startswith("cuda") or server_args.device == "npu"): + if not (cfg.device.startswith("cuda") or cfg.device == "npu"): raise ValueError( "DFLASH speculative decoding only supports CUDA and NPU devices." ) @@ -192,12 +196,12 @@ def _handle_dflash(server_args: ServerArgs) -> None: "Currently DFLASH speculative decoding does not support dp attention." ) - if server_args.pp_size != 1: + if cfg.pp_size != 1: raise ValueError( "Currently DFLASH speculative decoding only supports pp_size == 1." ) - if server_args.speculative_draft_model_path is None: + if cfg.speculative_draft_model_path is None: raise ValueError( "DFLASH speculative decoding requires setting --speculative-draft-model-path." ) @@ -207,16 +211,16 @@ def _handle_dflash(server_args: ServerArgs) -> None: # RoPE reservation). Force them to 1 to avoid surprising memory behavior. # # For DFlash, the natural unit is `block_size` (verify window length). - if server_args.speculative_num_steps is None: + if cfg.speculative_num_steps is None: declare_resolution( server_args, "_handle_dflash", speculative_num_steps=1, ) - elif int(server_args.speculative_num_steps) != 1: + elif int(cfg.speculative_num_steps) != 1: logger.warning( "DFLASH only supports speculative_num_steps == 1; overriding speculative_num_steps=%s to 1.", - server_args.speculative_num_steps, + cfg.speculative_num_steps, ) declare_resolution( server_args, @@ -224,16 +228,16 @@ def _handle_dflash(server_args: ServerArgs) -> None: speculative_num_steps=1, ) - if server_args.speculative_eagle_topk is None: + if cfg.speculative_eagle_topk is None: declare_resolution( server_args, "_handle_dflash", speculative_eagle_topk=1, ) - elif int(server_args.speculative_eagle_topk) != 1: + elif int(cfg.speculative_eagle_topk) != 1: logger.warning( "DFLASH only supports speculative_eagle_topk == 1; overriding speculative_eagle_topk=%s to 1.", - server_args.speculative_eagle_topk, + cfg.speculative_eagle_topk, ) declare_resolution( server_args, @@ -241,41 +245,41 @@ def _handle_dflash(server_args: ServerArgs) -> None: speculative_eagle_topk=1, ) - if server_args.speculative_dflash_block_size is not None: - if int(server_args.speculative_dflash_block_size) <= 0: + if cfg.speculative_dflash_block_size is not None: + if int(cfg.speculative_dflash_block_size) <= 0: raise ValueError( "DFLASH requires --speculative-dflash-block-size to be positive, " - f"got {server_args.speculative_dflash_block_size}." + f"got {cfg.speculative_dflash_block_size}." ) - if server_args.speculative_num_draft_tokens is not None and int( - server_args.speculative_num_draft_tokens - ) != int(server_args.speculative_dflash_block_size): + if cfg.speculative_num_draft_tokens is not None and int( + cfg.speculative_num_draft_tokens + ) != int(cfg.speculative_dflash_block_size): raise ValueError( "Both --speculative-num-draft-tokens and --speculative-dflash-block-size are set " "but they differ. For DFLASH they must match. " - f"speculative_num_draft_tokens={server_args.speculative_num_draft_tokens}, " - f"speculative_dflash_block_size={server_args.speculative_dflash_block_size}." + f"speculative_num_draft_tokens={cfg.speculative_num_draft_tokens}, " + f"speculative_dflash_block_size={cfg.speculative_dflash_block_size}." ) declare_resolution( server_args, "_handle_dflash", - speculative_num_draft_tokens=int(server_args.speculative_dflash_block_size), + speculative_num_draft_tokens=int(cfg.speculative_dflash_block_size), ) - if server_args.speculative_num_draft_tokens is None: + if cfg.speculative_num_draft_tokens is None: from sglang.srt.speculative.dflash_utils import ( parse_dflash_draft_config, ) - model_override_args = json.loads(server_args.json_model_override_args) + model_override_args = json.loads(cfg.json_model_override_args) inferred_block_size = None try: from sglang.srt.utils.hf_transformers_utils import get_config draft_hf_config = get_config( - server_args.speculative_draft_model_path, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.speculative_draft_model_revision, + cfg.speculative_draft_model_path, + trust_remote_code=cfg.trust_remote_code, + revision=cfg.speculative_draft_model_revision, model_override_args=model_override_args, ) inferred_block_size = parse_dflash_draft_config( @@ -300,18 +304,18 @@ def _handle_dflash(server_args: ServerArgs) -> None: speculative_num_draft_tokens=inferred_block_size, ) - if server_args.speculative_draft_window_size is not None: - draft_tokens = int(server_args.speculative_num_draft_tokens) - if server_args.speculative_draft_window_size < draft_tokens: + if cfg.speculative_draft_window_size is not None: + draft_tokens = int(cfg.speculative_num_draft_tokens) + if cfg.speculative_draft_window_size < draft_tokens: raise ValueError( "--speculative-draft-window-size must be >= " "--speculative-num-draft-tokens (block_size). " - f"window_size={server_args.speculative_draft_window_size}, block_size={draft_tokens}." + f"window_size={cfg.speculative_draft_window_size}, block_size={draft_tokens}." ) _resolve_dflash_draft_attention_backend(server_args) - if server_args.max_running_requests is None: + if cfg.max_running_requests is None: declare_resolution( server_args, "_handle_dflash", @@ -321,7 +325,7 @@ def _handle_dflash(server_args: ServerArgs) -> None: "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." ) - if server_args.enable_mixed_chunk: + if cfg.enable_mixed_chunk: declare_resolution( server_args, "_handle_dflash", @@ -341,23 +345,24 @@ def _target_checkpoint_bundles_dspark_draft(server_args: ServerArgs) -> bool: def _handle_dspark(server_args: ServerArgs) -> None: - _is_npu = server_args.device.startswith("npu") - if not server_args.device.startswith(("cuda", "npu")): + cfg = resolving_view(server_args) + _is_npu = cfg.device.startswith("npu") + if not cfg.device.startswith(("cuda", "npu")): raise ValueError( "DSpark speculative decoding only supports CUDA or NPU device." ) # dp_size==1 with dp_attention is a degenerate flag under DSV4 CP; skip DP-only checks. - if server_args.enable_dp_attention and server_args.dp_size > 1: - if not server_args.enable_dp_lm_head: + if cfg.enable_dp_attention and cfg.dp_size > 1: + if not cfg.enable_dp_lm_head: raise ValueError("DSpark with dp attention requires --enable-dp-lm-head.") - if not _is_npu and server_args.moe_a2a_backend not in ("none", "megamoe"): + if not _is_npu and cfg.moe_a2a_backend not in ("none", "megamoe"): raise ValueError( "DSpark with dp attention supports moe_a2a_backend 'none' " "(built-in TP MoE) or 'megamoe', got " - f"{server_args.moe_a2a_backend!r}." + f"{cfg.moe_a2a_backend!r}." ) - if not _is_npu and server_args.moe_a2a_backend != "none": + if not _is_npu and cfg.moe_a2a_backend != "none": from sglang.srt.speculative.ragged_verify import ( RaggedVerifyMode, read_ragged_verify_mode, @@ -366,46 +371,46 @@ def _handle_dspark(server_args: ServerArgs) -> None: if read_ragged_verify_mode() is not RaggedVerifyMode.STATIC: raise ValueError( "DSpark with dp attention + " - f"moe_a2a_backend={server_args.moe_a2a_backend!r} requires " + f"moe_a2a_backend={cfg.moe_a2a_backend!r} requires " "SGLANG_RAGGED_VERIFY_MODE=static." ) - if server_args.attn_cp_size > 1: + if cfg.attn_cp_size > 1: raise ValueError( "DSpark with dp attention does not support context parallel " - f"(attn_cp_size={server_args.attn_cp_size})." + f"(attn_cp_size={cfg.attn_cp_size})." ) if ( not _is_npu - and server_args.speculative_moe_a2a_backend is not None - and server_args.speculative_moe_a2a_backend != server_args.moe_a2a_backend + and cfg.speculative_moe_a2a_backend is not None + and cfg.speculative_moe_a2a_backend != cfg.moe_a2a_backend ): raise ValueError( "DSpark ignores --speculative-moe-a2a-backend; with dp attention it " - f"must match the target moe_a2a_backend={server_args.moe_a2a_backend!r} " - f"(got {server_args.speculative_moe_a2a_backend!r})." + f"must match the target moe_a2a_backend={cfg.moe_a2a_backend!r} " + f"(got {cfg.speculative_moe_a2a_backend!r})." ) - if server_args.pp_size != 1: + if cfg.pp_size != 1: raise ValueError( "Currently DSpark speculative decoding only supports pp_size == 1." ) - if server_args.speculative_draft_model_path is None: + if cfg.speculative_draft_model_path is None: if _target_checkpoint_bundles_dspark_draft(server_args): declare_resolution( server_args, "_handle_dspark", - speculative_draft_model_path=server_args.model_path, + speculative_draft_model_path=cfg.model_path, ) declare_resolution( server_args, "_handle_dspark", - speculative_draft_model_revision=server_args.revision, + speculative_draft_model_revision=cfg.revision, ) logger.info( "DSpark draft weights are bundled in the target checkpoint; " "defaulting --speculative-draft-model-path to --model-path (%s).", - server_args.model_path, + cfg.model_path, ) else: raise ValueError( @@ -413,16 +418,16 @@ def _handle_dspark(server_args: ServerArgs) -> None: "--speculative-draft-model-path." ) - if server_args.speculative_num_steps is None: + if cfg.speculative_num_steps is None: declare_resolution( server_args, "_handle_dspark", speculative_num_steps=1, ) - elif int(server_args.speculative_num_steps) != 1: + elif int(cfg.speculative_num_steps) != 1: logger.warning( "DSpark only supports speculative_num_steps == 1; overriding speculative_num_steps=%s to 1.", - server_args.speculative_num_steps, + cfg.speculative_num_steps, ) declare_resolution( server_args, @@ -430,16 +435,16 @@ def _handle_dspark(server_args: ServerArgs) -> None: speculative_num_steps=1, ) - if server_args.speculative_eagle_topk is None: + if cfg.speculative_eagle_topk is None: declare_resolution( server_args, "_handle_dspark", speculative_eagle_topk=1, ) - elif int(server_args.speculative_eagle_topk) != 1: + elif int(cfg.speculative_eagle_topk) != 1: logger.warning( "DSpark only supports speculative_eagle_topk == 1; overriding speculative_eagle_topk=%s to 1.", - server_args.speculative_eagle_topk, + cfg.speculative_eagle_topk, ) declare_resolution( server_args, @@ -463,17 +468,17 @@ def _handle_dspark(server_args: ServerArgs) -> None: ) gamma: Optional[int] = None - if server_args.speculative_dspark_block_size is not None: - if int(server_args.speculative_dspark_block_size) <= 0: + if cfg.speculative_dspark_block_size is not None: + if int(cfg.speculative_dspark_block_size) <= 0: raise ValueError( "DSpark requires --speculative-dspark-block-size to be positive, " - f"got {server_args.speculative_dspark_block_size}." + f"got {cfg.speculative_dspark_block_size}." ) - gamma = int(server_args.speculative_dspark_block_size) + gamma = int(cfg.speculative_dspark_block_size) else: if draft_config is not None: gamma = draft_config.resolve_gamma(default=None) - if gamma is None and server_args.speculative_num_draft_tokens is None: + if gamma is None and cfg.speculative_num_draft_tokens is None: gamma = DEFAULT_DSPARK_GAMMA logger.warning( "DSpark gamma is not set; defaulting to %d.", @@ -483,13 +488,13 @@ def _handle_dspark(server_args: ServerArgs) -> None: if gamma is not None: verify_window = int(gamma) + 1 if ( - server_args.speculative_num_draft_tokens is not None - and int(server_args.speculative_num_draft_tokens) != verify_window + cfg.speculative_num_draft_tokens is not None + and int(cfg.speculative_num_draft_tokens) != verify_window ): raise ValueError( "DSpark speculative_num_draft_tokens must equal gamma + 1 " f"(= {verify_window} for gamma={gamma}), but got " - f"speculative_num_draft_tokens={server_args.speculative_num_draft_tokens}." + f"speculative_num_draft_tokens={cfg.speculative_num_draft_tokens}." ) declare_resolution( server_args, @@ -497,18 +502,18 @@ def _handle_dspark(server_args: ServerArgs) -> None: speculative_num_draft_tokens=verify_window, ) - if server_args.speculative_num_draft_tokens is None: + if cfg.speculative_num_draft_tokens is None: raise ValueError( "DSpark could not resolve speculative_num_draft_tokens; set " "--speculative-dspark-block-size (= gamma)." ) - if int(server_args.speculative_num_draft_tokens) < 2: + if int(cfg.speculative_num_draft_tokens) < 2: raise ValueError( "DSpark speculative_num_draft_tokens must be >= 2 (= gamma + 1), " - f"got {server_args.speculative_num_draft_tokens}." + f"got {cfg.speculative_num_draft_tokens}." ) - if server_args.max_running_requests is None: + if cfg.max_running_requests is None: declare_resolution( server_args, "_handle_dspark", @@ -518,7 +523,7 @@ def _handle_dspark(server_args: ServerArgs) -> None: "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." ) - if server_args.enable_mixed_chunk: + if cfg.enable_mixed_chunk: declare_resolution( server_args, "_handle_dspark", @@ -535,7 +540,7 @@ def _handle_dspark(server_args: ServerArgs) -> None: ragged_mode = read_ragged_verify_mode() if ( - server_args.speculative_dspark_align_verify_tokens_to_graph_tier + cfg.speculative_dspark_align_verify_tokens_to_graph_tier and ragged_mode is not RaggedVerifyMode.COMPACT ): logger.warning( @@ -544,10 +549,7 @@ def _handle_dspark(server_args: ServerArgs) -> None: "a no-op.", ragged_mode.value, ) - if ( - server_args.speculative_dspark_sps_table_path - and ragged_mode is RaggedVerifyMode.STATIC - ): + if cfg.speculative_dspark_sps_table_path and ragged_mode is RaggedVerifyMode.STATIC: logger.warning( "--speculative-dspark-sps-table-path feeds the ragged-verify budget " "scheduler, which is off under SGLANG_RAGGED_VERIFY_MODE=static; it " @@ -561,6 +563,7 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None: Consumed by ModelRunner's `is_draft_worker` override (one backend for all draft modes). """ + cfg = resolving_view(server_args) from sglang.srt.utils import is_hip supported_draft_backends = ( @@ -574,7 +577,7 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None: # Use triton on ROCm (no FlashInfer), flashinfer on CUDA. fallback_backend = "triton" if is_hip() else "flashinfer" - draft_backend = server_args.speculative_draft_attention_backend + draft_backend = cfg.speculative_draft_attention_backend if draft_backend is None: from sglang.srt.arg_groups.overrides import ( attention_backends_of, @@ -589,10 +592,10 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None: from sglang.srt.utils.hf_transformers_utils import get_config draft_hf_config = get_config( - server_args.speculative_draft_model_path, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.speculative_draft_model_revision, - model_override_args=json.loads(server_args.json_model_override_args), + cfg.speculative_draft_model_path, + trust_remote_code=cfg.trust_remote_code, + revision=cfg.speculative_draft_model_revision, + model_override_args=json.loads(cfg.json_model_override_args), ) draft_text_config = ( getattr(draft_hf_config, "text_config", None) or draft_hf_config @@ -635,7 +638,8 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None: def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None: - if server_args.max_running_requests is None: + cfg = resolving_view(server_args) + if cfg.max_running_requests is None: declare_resolution( server_args, "_handle_frozen_kv_mtp", @@ -645,7 +649,7 @@ def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None: "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." ) - if server_args.enable_mixed_chunk: + if cfg.enable_mixed_chunk: declare_resolution( server_args, "_handle_frozen_kv_mtp", @@ -658,13 +662,14 @@ def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None: def _handle_eagle_family(server_args: ServerArgs) -> None: + cfg = resolving_view(server_args) from sglang.srt.arg_groups.overrides import ( attention_backends_of, resolved_view, ) if ( - server_args.speculative_algorithm == "STANDALONE" + cfg.speculative_algorithm == "STANDALONE" and resolved_view(server_args).enable_dp_attention ): # TODO: support dp attention for standalone speculative decoding @@ -672,7 +677,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: "Currently standalone speculative decoding does not support dp attention." ) - if server_args.max_running_requests is None: + if cfg.max_running_requests is None: declare_resolution( server_args, "_handle_eagle_family", @@ -690,7 +695,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: "speculative decoding." ) - if server_args.enable_mixed_chunk: + if cfg.enable_mixed_chunk: declare_resolution( server_args, "_handle_eagle_family", @@ -716,16 +721,16 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: "PixtralForConditionalGeneration", "HYV3ForCausalLM", ]: - if server_args.speculative_draft_model_path is None: + if cfg.speculative_draft_model_path is None: declare_resolution( server_args, "_handle_eagle_family", - speculative_draft_model_path=server_args.model_path, + speculative_draft_model_path=cfg.model_path, ) declare_resolution( server_args, "_handle_eagle_family", - speculative_draft_model_revision=server_args.revision, + speculative_draft_model_revision=cfg.revision, ) else: if model_arch not in [ @@ -736,13 +741,10 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: "DeepSeek MTP does not require setting speculative_draft_model_path." ) - if ( - not server_args.speculative_adaptive - and server_args.speculative_num_steps is None - ): + if not cfg.speculative_adaptive and cfg.speculative_num_steps is None: assert ( - server_args.speculative_eagle_topk is None - and server_args.speculative_num_draft_tokens is None + cfg.speculative_eagle_topk is None + and cfg.speculative_num_draft_tokens is None ) steps, topk, draft_tokens = _auto_choose_speculative_params( @@ -757,29 +759,29 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: ) if "trtllm_mha" in attention_backends_of(resolved_view(server_args)): - if server_args.speculative_eagle_topk > 1: + if cfg.speculative_eagle_topk > 1: raise ValueError( "trtllm_mha backend only supports topk = 1 for speculative decoding." ) - if server_args.speculative_use_rejection_sampling: + if cfg.speculative_use_rejection_sampling: # Resolved alias by now: NEXTN -> EAGLE, Gemma4 draft -> FROZEN_KV_MTP. # Only the EAGLE/EAGLE3 draft workers emit a target-vocab proposal that # the rejection-sampling kernel consumes; everything else (STANDALONE, # FROZEN_KV_MTP, NGRAM, DFLASH) is unsupported. - if server_args.speculative_algorithm not in ("EAGLE", "EAGLE3"): + if cfg.speculative_algorithm not in ("EAGLE", "EAGLE3"): raise NotImplementedError( "--speculative-use-rejection-sampling is only supported for " "EAGLE / EAGLE3 / NEXTN, not " - f"speculative_algorithm={server_args.speculative_algorithm}." + f"speculative_algorithm={cfg.speculative_algorithm}." ) - if server_args.speculative_eagle_topk != 1: + if cfg.speculative_eagle_topk != 1: raise ValueError( "--speculative-use-rejection-sampling requires --speculative-eagle-topk=1." ) if ( - server_args.speculative_accept_threshold_single != 1.0 - or server_args.speculative_accept_threshold_acc != 1.0 + cfg.speculative_accept_threshold_single != 1.0 + or cfg.speculative_accept_threshold_acc != 1.0 ): raise ValueError( "--speculative-use-rejection-sampling is incompatible with " @@ -787,7 +789,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: "--speculative-accept-threshold-acc; rejection sampling ignores " "the accept thresholds." ) - if server_args.enable_deterministic_inference: + if cfg.enable_deterministic_inference: raise ValueError( "--speculative-use-rejection-sampling is incompatible with " "--enable-deterministic-inference; the sampling kernel draws " @@ -798,7 +800,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: if ( resolved_view(server_args).enable_multi_layer_eagle - and server_args.speculative_eagle_topk != 1 + and cfg.speculative_eagle_topk != 1 ): raise ValueError( "--speculative-use-rejection-sampling with multi-layer EAGLE " @@ -811,9 +813,8 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: ) if ( - server_args.speculative_eagle_topk == 1 - and server_args.speculative_num_draft_tokens - != server_args.speculative_num_steps + 1 + cfg.speculative_eagle_topk == 1 + and cfg.speculative_num_draft_tokens != cfg.speculative_num_steps + 1 ): logger.warning( "speculative_num_draft_tokens is adjusted to speculative_num_steps + 1 when speculative_eagle_topk == 1" @@ -821,7 +822,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: declare_resolution( server_args, "_handle_eagle_family", - speculative_num_draft_tokens=server_args.speculative_num_steps + 1, + speculative_num_draft_tokens=cfg.speculative_num_steps + 1, ) # topk > 1 + page_size > 1 needs the two-pass cascade draft-decode (shared prefix @@ -830,7 +831,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: _PAGE_TREE_SPEC_BACKENDS = ("flashinfer", "fa3", "triton") view = resolved_view(server_args) if ( - server_args.speculative_eagle_topk > 1 + cfg.speculative_eagle_topk > 1 and view.page_size > 1 and view.attention_backend not in _PAGE_TREE_SPEC_BACKENDS ): @@ -842,14 +843,15 @@ def _handle_eagle_family(server_args: ServerArgs) -> None: def _handle_ngram(server_args: ServerArgs) -> None: - if server_args.device not in ("cuda", "cpu"): + cfg = resolving_view(server_args) + if cfg.device not in ("cuda", "cpu"): raise ValueError( "Ngram speculative decoding only supports CUDA or CPU devices." ) _disable_overlap_schedule_for_cpu(server_args) - if server_args.max_running_requests is None: + if cfg.max_running_requests is None: declare_resolution( server_args, "_handle_ngram", @@ -867,9 +869,9 @@ def _handle_ngram(server_args: ServerArgs) -> None: declare_resolution( server_args, "_handle_ngram", - speculative_eagle_topk=server_args.speculative_ngram_max_bfs_breadth, + speculative_eagle_topk=cfg.speculative_ngram_max_bfs_breadth, ) - if server_args.speculative_num_draft_tokens is None: + if cfg.speculative_num_draft_tokens is None: declare_resolution( server_args, "_handle_ngram", @@ -879,31 +881,31 @@ def _handle_ngram(server_args: ServerArgs) -> None: "speculative_num_draft_tokens is set to 12 by default for ngram speculative decoding. " "You can override this by explicitly setting --speculative-num-draft-tokens." ) - if server_args.speculative_num_steps is None: + if cfg.speculative_num_steps is None: declare_resolution( server_args, "_handle_ngram", - speculative_num_steps=server_args.speculative_num_draft_tokens - // server_args.speculative_eagle_topk, + speculative_num_steps=cfg.speculative_num_draft_tokens + // cfg.speculative_eagle_topk, ) - if server_args.speculative_ngram_external_corpus_path is not None: - if server_args.speculative_ngram_external_sam_budget <= 0: + if cfg.speculative_ngram_external_corpus_path is not None: + if cfg.speculative_ngram_external_sam_budget <= 0: raise ValueError( "--speculative-ngram-external-sam-budget must be positive when " "--speculative-ngram-external-corpus-path is set." ) - if server_args.speculative_ngram_external_corpus_max_tokens <= 0: + if cfg.speculative_ngram_external_corpus_max_tokens <= 0: raise ValueError( "--speculative-ngram-external-corpus-max-tokens must be positive when " "--speculative-ngram-external-corpus-path is set." ) if ( - server_args.speculative_ngram_external_sam_budget - > server_args.speculative_num_draft_tokens - 1 + cfg.speculative_ngram_external_sam_budget + > cfg.speculative_num_draft_tokens - 1 ): raise ValueError( "speculative_ngram_external_sam_budget must be less than or equal to " - f"speculative_num_draft_tokens - 1 ({server_args.speculative_num_draft_tokens - 1})." + f"speculative_num_draft_tokens - 1 ({cfg.speculative_num_draft_tokens - 1})." ) logger.warning( "The mixed chunked prefill are disabled because of " @@ -914,12 +916,12 @@ def _handle_ngram(server_args: ServerArgs) -> None: view = resolved_view(server_args) if ( - server_args.speculative_eagle_topk > 1 + cfg.speculative_eagle_topk > 1 and view.page_size > 1 and view.attention_backend != "flashinfer" ): raise ValueError( - f"speculative_eagle_topk({server_args.speculative_eagle_topk}) > 1 " + f"speculative_eagle_topk({cfg.speculative_eagle_topk}) > 1 " f"with page_size({view.page_size}) > 1 is unstable " "and produces incorrect results for paged attention backends. " "This combination is only supported for the 'flashinfer' backend." @@ -950,31 +952,32 @@ def _maybe_disable_adaptive(server_args: ServerArgs) -> None: def _init_adaptive_speculative_params(server_args: ServerArgs) -> None: + cfg = resolving_view(server_args) from sglang.srt.speculative.adaptive_spec_params import ( resolve_candidate_steps_from_config, ) candidate_steps = resolve_candidate_steps_from_config( - cfg_path=server_args.speculative_adaptive_config, + cfg_path=cfg.speculative_adaptive_config, ) - if server_args.speculative_eagle_topk is None: + if cfg.speculative_eagle_topk is None: declare_resolution( server_args, "_init_adaptive_speculative_params", speculative_eagle_topk=1, ) - if server_args.speculative_num_steps is None: + if cfg.speculative_num_steps is None: declare_resolution( server_args, "_init_adaptive_speculative_params", speculative_num_steps=candidate_steps[len(candidate_steps) // 2], ) - if server_args.speculative_num_steps not in candidate_steps: + if cfg.speculative_num_steps not in candidate_steps: raise ValueError( - f"--speculative-num-steps={server_args.speculative_num_steps} " + f"--speculative-num-steps={cfg.speculative_num_steps} " f"is not in the adaptive config candidate_steps {candidate_steps}. " "Pass one of those values." ) @@ -982,7 +985,7 @@ def _init_adaptive_speculative_params(server_args: ServerArgs) -> None: declare_resolution( server_args, "_init_adaptive_speculative_params", - speculative_num_draft_tokens=server_args.speculative_num_steps + 1, + speculative_num_draft_tokens=cfg.speculative_num_steps + 1, ) @@ -992,7 +995,8 @@ def _auto_choose_speculative_params(server_args: ServerArgs, model_arch: str) -> You can tune them on your own models and prompts with scripts/playground/bench_speculative.py """ - if server_args.speculative_algorithm == "STANDALONE": + cfg = resolving_view(server_args) + if cfg.speculative_algorithm == "STANDALONE": return (3, 1, 4) if model_arch in ["LlamaForCausalLM"]: return (5, 4, 8) diff --git a/python/sglang/srt/model_loader/expert_pack_runtime.py b/python/sglang/srt/model_loader/expert_pack_runtime.py index 173d890b0c15..cb7fda46b694 100644 --- a/python/sglang/srt/model_loader/expert_pack_runtime.py +++ b/python/sglang/srt/model_loader/expert_pack_runtime.py @@ -196,10 +196,14 @@ def prepare_raw_kimi_server_args( server_args: Any, loader_config: dict[str, Any] ) -> None: """Resolve a raw GGUF model path into the normal loader inputs.""" - model_path = Path(server_args.model_path).expanduser() + + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) + model_path = Path(cfg.model_path).expanduser() if not model_path.is_file() or model_path.suffix.lower() != ".gguf": return - tokenizer_path = server_args.tokenizer_path + tokenizer_path = cfg.tokenizer_path if tokenizer_path and Path(tokenizer_path).expanduser() == model_path: tokenizer_path = None assets = ensure_kimi_assets( @@ -503,7 +507,11 @@ def prepare_raw_deepseek_server_args( server_args: Any, loader_config: dict[str, Any] ) -> None: """Resolve a raw DeepSeek V4 GGUF into metadata and Expert Pack inputs.""" - source = Path(server_args.model_path).expanduser().resolve(strict=True) + + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) + source = Path(cfg.model_path).expanduser().resolve(strict=True) if not source.is_file(): return repo = _repo_root() @@ -546,7 +554,11 @@ def prepare_raw_expert_pack_server_args( server_args: Any, loader_config: dict[str, Any] ) -> None: """Dispatch a raw GGUF to the model-specific expert-pack preparation path.""" - source = Path(server_args.model_path).expanduser() + + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) + source = Path(cfg.model_path).expanduser() if not source.is_file(): return name = source.name.upper() diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index 3c366806f0d5..466e622b3a42 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -176,8 +176,6 @@ ("layers/moe/utils.py", "quantization"), ("layers/moe/utils.py", "speculative_moe_runner_backend"), ("lora/marlin_lora_temp/policy.py", "lora_paths"), - ("model_loader/expert_pack_runtime.py", "model_path"), - ("model_loader/expert_pack_runtime.py", "tokenizer_path"), ("parser/template_detection.py", "model_path"), ("speculative/adaptive_spec_params.py", "speculative_algorithm"), ("speculative/adaptive_spec_params.py", "speculative_eagle_topk"), @@ -221,7 +219,6 @@ ("dllm/config.py", "model_path"), ("entrypoints/engine.py", "reasoning_parser"), ("entrypoints/engine.py", "tool_call_parser"), - ("model_loader/expert_pack_runtime.py", "model_path"), ("weight_cache/daemon.py", "dp_size"), ("weight_cache/daemon.py", "dtype"), ("weight_cache/daemon.py", "ep_size"), From 79e67aeed7ff9d1d6d2c07d9826656e22994687b Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 11:33:55 +0000 Subject: [PATCH 08/26] config: the resolution handlers read the declaration stash, not the fields The other half of the readers: 779 `self.` reads across the 88 resolution handlers reachable from the dispatcher now go through `resolving_view(self)`. Same argument as the hooks -- while `declare_resolution` still writes the field as it declares, the view and the field agree, so this changes nothing and can be checked exactly: 20 launch shapes, every field of the resolution result compared against the previous commit, zero differences. What it buys is that the pipeline no longer depends on that write to see its own decisions. Removing it -- so the record holds the raw input and nothing else -- needs the readers outside `arg_groups` and `ServerArgs` that resolution calls with the record (the platform defaults, the spec-algo hook, `ModelConfig`, the CP strategy) to move as well; those are next. --- python/sglang/srt/server_args.py | 1516 ++++++++++++++++-------------- 1 file changed, 795 insertions(+), 721 deletions(-) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index d8e483f7b7e7..de4a5cc44535 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -48,6 +48,7 @@ mamba_extra_buffer_of, remote_instance_transfer_engine_of, resolved_view, + resolving_view, ) from sglang.srt.configs.embedding_model_spec import BCGPrefillPolicy from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_spec_by_arch @@ -4029,12 +4030,13 @@ def _run_resolution_pipeline(self): materialize_declarations(self) def _handle_return_hidden_states_mode(self): - if self.return_hidden_states_mode not in (None, "last", "full"): + cfg = resolving_view(self) + if cfg.return_hidden_states_mode not in (None, "last", "full"): raise ValueError( "return_hidden_states_mode must be one of: None, 'last', or 'full'." ) - if self.return_hidden_states_mode is None: - if self.enable_return_hidden_states: + if cfg.return_hidden_states_mode is None: + if cfg.enable_return_hidden_states: self._declare( "_handle_return_hidden_states_mode", return_hidden_states_mode="full", @@ -4046,7 +4048,8 @@ def _handle_return_hidden_states_mode(self): ) def _handle_model_capability_adjustments(self): - if parse_connector_type(self.model_path) == ConnectorType.INSTANCE: + cfg = resolving_view(self) + if parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE: return from sglang.srt.arg_groups.overrides import ( _hrm_text_attention_force, @@ -4083,8 +4086,8 @@ def _handle_model_capability_adjustments(self): ) # cuda_graph_config was already parsed from the legacy boolean, so # flipping the boolean alone would not stop graph capture. - self.cuda_graph_config.decode.backend = Backend.DISABLED - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED logger.warning( "HRM-Text (prefix_lm) detected: forcing --attention-backend " "triton, --chunked-prefill-size -1, --disable-radix-cache, and " @@ -4110,7 +4113,7 @@ def _handle_model_capability_adjustments(self): if ( embedding_model_spec is not None and embedding_model_spec.auto_enable_embedding - and not self.is_embedding + and not cfg.is_embedding ): self._declare( "_handle_model_capability_adjustments", @@ -4151,7 +4154,7 @@ def _handle_model_capability_adjustments(self): enable_tokenizer_batch_encode=True, ) requested_prefill_backend = ( - self.prefill_attention_backend or self.attention_backend + cfg.prefill_attention_backend or cfg.attention_backend ) if ( is_cuda() @@ -4167,15 +4170,15 @@ def _handle_model_capability_adjustments(self): prefill_only_disable_kv_cache=True, ) self._validate_prefill_only_disable_kv_cache_args() - self.cuda_graph_config.decode.backend = Backend.DISABLED - if is_cuda() and self.cuda_graph_config.prefill.backend != Backend.DISABLED: - self.cuda_graph_config.prefill.backend = Backend.BREAKABLE + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + if is_cuda() and cfg.cuda_graph_config.prefill.backend != Backend.DISABLED: + cfg.cuda_graph_config.prefill.backend = Backend.BREAKABLE # CUDA-graph sizing has already run by this point and derives # its generic maximum from the 8K chunked-prefill default. # On the Hopper/Blackwell FA raw-K/V path, raise the unlocked # default to a full eight-way 2K embedding batch; callers can # still override this for larger aggregate prefills. - prefill_config = self.cuda_graph_config.prefill + prefill_config = cfg.cuda_graph_config.prefill # Unit-level capability tests may invoke this hook without # running the full CUDA-graph configuration parser, which is # where this internal lock set is normally initialized. @@ -4199,7 +4202,7 @@ def _handle_model_capability_adjustments(self): elif not is_cuda(): # BCG is CUDA-only. Other graph backends do not support this # encoder-style prefill, so retain the eager Triton path. - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED logger.info( "EmbeddingGemma detected: disabling radix cache and chunked " "prefill; using breakable CUDA graph for CUDA prefill." @@ -4220,13 +4223,14 @@ def _handle_model_capability_adjustments(self): def _handle_model_source_paths(self): """Prepare metadata for model paths backed by remote object stores.""" + cfg = resolving_view(self) self._resolve_hf_gguf_model_path() seen_paths = set() for model_path in ( - self.model_path, - self.tokenizer_path, - self.speculative_draft_model_path, + cfg.model_path, + cfg.tokenizer_path, + cfg.speculative_draft_model_path, ): if ( model_path is not None @@ -4244,20 +4248,21 @@ def _handle_pd_disaggregation(self): handle_pd_disaggregation(self) def _handle_dcp_validation(self): - if self.dcp_size < 1: + cfg = resolving_view(self) + if cfg.dcp_size < 1: raise ValueError( "Decode context parallel size (--dcp-size / " "--decode-context-parallel-size) must be >= 1, but got " - f"dcp_size={self.dcp_size}." + f"dcp_size={cfg.dcp_size}." ) - if self.dcp_comm_backend in ("a2a", "fi_a2a") and self.dcp_size <= 1: + if cfg.dcp_comm_backend in ("a2a", "fi_a2a") and cfg.dcp_size <= 1: raise ValueError( - f"--dcp-comm-backend {self.dcp_comm_backend} only affects the " + f"--dcp-comm-backend {cfg.dcp_comm_backend} only affects the " "decode context-parallel attention reduction and therefore " "requires --dcp-size / --decode-context-parallel-size > 1, but " - f"got dcp_size={self.dcp_size}." + f"got dcp_size={cfg.dcp_size}." ) - if self.dcp_comm_backend == "fi_a2a" and not is_cuda(): + if cfg.dcp_comm_backend == "fi_a2a" and not is_cuda(): raise ValueError( "--dcp-comm-backend fi_a2a delegates the exchange to FlashInfer's " "MNNVL All-to-All kernel, which requires an NVIDIA CUDA platform " @@ -4265,23 +4270,22 @@ def _handle_dcp_validation(self): "authoritative fabric probe runs at model-runner init; use 'a2a' " "or 'ag_rs' on clusters without MNNVL." ) - if self.dcp_replicate_q_proj: - if self.dcp_size <= 1: + if cfg.dcp_replicate_q_proj: + if cfg.dcp_size <= 1: raise ValueError("--dcp-replicate-q-proj requires --dcp-size > 1.") - if self.dcp_comm_backend not in ("a2a", "fi_a2a"): + if cfg.dcp_comm_backend not in ("a2a", "fi_a2a"): raise ValueError( "--dcp-replicate-q-proj only applies to the a2a/fi_a2a DCP " "communication backend (it removes the head-dim Q all-gather); " - f"got --dcp-comm-backend={self.dcp_comm_backend}." + f"got --dcp-comm-backend={cfg.dcp_comm_backend}." ) def _handle_load_balance_method(self): - if self.disaggregation_mode not in ("null", "prefill", "decode"): - raise ValueError( - f"Invalid disaggregation_mode={self.disaggregation_mode!r}" - ) + cfg = resolving_view(self) + if cfg.disaggregation_mode not in ("null", "prefill", "decode"): + raise ValueError(f"Invalid disaggregation_mode={cfg.disaggregation_mode!r}") - if self.load_balance_method == "auto": + if cfg.load_balance_method == "auto": # Default behavior: # - non-PD: round_robin # - PD prefill: follow_bootstrap_room @@ -4290,7 +4294,7 @@ def _handle_load_balance_method(self): "_handle_load_balance_method", load_balance_method=( "follow_bootstrap_room" - if self.disaggregation_mode == "prefill" + if cfg.disaggregation_mode == "prefill" else "round_robin" ), ) @@ -4298,47 +4302,48 @@ def _handle_load_balance_method(self): def _handle_ssl_validation(self): """Ensure SSL arguments are consistent and referenced files exist.""" - if self.ssl_keyfile and not self.ssl_certfile: + cfg = resolving_view(self) + if cfg.ssl_keyfile and not cfg.ssl_certfile: raise ValueError( "--ssl-keyfile requires --ssl-certfile to be specified as well." ) - if self.ssl_certfile and not self.ssl_keyfile: + if cfg.ssl_certfile and not cfg.ssl_keyfile: raise ValueError( "--ssl-certfile requires --ssl-keyfile to be specified as well." ) - if not self.ssl_certfile and not self.ssl_keyfile: - if self.ssl_ca_certs: + if not cfg.ssl_certfile and not cfg.ssl_keyfile: + if cfg.ssl_ca_certs: raise ValueError( "--ssl-ca-certs has no effect without --ssl-certfile and --ssl-keyfile." ) - if self.ssl_keyfile_password: + if cfg.ssl_keyfile_password: raise ValueError( "--ssl-keyfile-password has no effect without --ssl-certfile and --ssl-keyfile." ) # Validate files exist early to avoid late failures after model loading. - if self.ssl_keyfile and not os.path.isfile(self.ssl_keyfile): + if cfg.ssl_keyfile and not os.path.isfile(cfg.ssl_keyfile): raise ValueError( - f"SSL key file not found: '{self.ssl_keyfile}'. " + f"SSL key file not found: '{cfg.ssl_keyfile}'. " f"Please check the --ssl-keyfile path." ) - if self.ssl_certfile and not os.path.isfile(self.ssl_certfile): + if cfg.ssl_certfile and not os.path.isfile(cfg.ssl_certfile): raise ValueError( - f"SSL certificate file not found: '{self.ssl_certfile}'. " + f"SSL certificate file not found: '{cfg.ssl_certfile}'. " f"Please check the --ssl-certfile path." ) - if self.ssl_ca_certs and not os.path.isfile(self.ssl_ca_certs): + if cfg.ssl_ca_certs and not os.path.isfile(cfg.ssl_ca_certs): raise ValueError( - f"SSL CA certificates file not found: '{self.ssl_ca_certs}'. " + f"SSL CA certificates file not found: '{cfg.ssl_ca_certs}'. " f"Please check the --ssl-ca-certs path." ) - if self.enable_ssl_refresh and not (self.ssl_certfile and self.ssl_keyfile): + if cfg.enable_ssl_refresh and not (cfg.ssl_certfile and cfg.ssl_keyfile): raise ValueError( "--enable-ssl-refresh requires --ssl-certfile and --ssl-keyfile " "to be specified." ) - if self.enable_http2: - if not 0 < self.http2_max_concurrent_streams < 2**32: + if cfg.enable_http2: + if not 0 < cfg.http2_max_concurrent_streams < 2**32: raise ValueError( "--http2-max-concurrent-streams must be between 1 and " "4294967295." @@ -4352,7 +4357,7 @@ def _handle_ssl_validation(self): 'Install it with: pip install "sglang[http2]"' ) - if self.enable_ssl_refresh: + if cfg.enable_ssl_refresh: raise ValueError( "--enable-ssl-refresh is not supported with --enable-http2. " "Granian does not support SSL certificate hot-reloading. " @@ -4361,42 +4366,45 @@ def _handle_ssl_validation(self): def _handle_multimodal(self): """Validate mm_process_config structure before model loading.""" + cfg = resolving_view(self) if ( - self.mm_preprocess_cache_size_mb is not None - and self.mm_preprocess_cache_size_mb < 0 + cfg.mm_preprocess_cache_size_mb is not None + and cfg.mm_preprocess_cache_size_mb < 0 ): raise ValueError("mm_preprocess_cache_size_mb must be non-negative") - if self.mm_process_config is not None: - if not isinstance(self.mm_process_config, dict): + if cfg.mm_process_config is not None: + if not isinstance(cfg.mm_process_config, dict): raise TypeError( f"mm_process_config must be a dict, " - f"but got {type(self.mm_process_config)}" + f"but got {type(cfg.mm_process_config)}" ) for key in ("image", "video", "audio"): - if key in self.mm_process_config and not isinstance( - self.mm_process_config[key], dict + if key in cfg.mm_process_config and not isinstance( + cfg.mm_process_config[key], dict ): raise TypeError( f"mm_process_config['{key}'] must be a dict, " - f"but got {type(self.mm_process_config[key])}" + f"but got {type(cfg.mm_process_config[key])}" ) def _handle_media_url_security(self): """Normalize and publish the media URL policy before workers start.""" + cfg = resolving_view(self) self._declare( "_handle_media_url_security", allowed_media_domains=configure_media_url_security( - self.allowed_media_domains, - self.media_url_max_file_size_mb, + cfg.allowed_media_domains, + cfg.media_url_max_file_size_mb, ), ) def _handle_deprecated_args(self): - if self.disable_fast_image_processor: - if self.image_processor_backend not in {"auto", "pil"}: + cfg = resolving_view(self) + if cfg.disable_fast_image_processor: + if cfg.image_processor_backend not in {"auto", "pil"}: raise ValueError( "--disable-fast-image-processor conflicts with " - f"--image-processor-backend={self.image_processor_backend}." + f"--image-processor-backend={cfg.image_processor_backend}." ) logger.warning( "--disable-fast-image-processor is deprecated; use " @@ -4406,19 +4414,19 @@ def _handle_deprecated_args(self): # Handle deprecated tool call parsers deprecated_tool_call_parsers = {"qwen25": "qwen", "glm45": "glm"} - if self.tool_call_parser in deprecated_tool_call_parsers: + if cfg.tool_call_parser in deprecated_tool_call_parsers: logger.warning( - f"The tool_call_parser '{self.tool_call_parser}' is deprecated. Please use '{deprecated_tool_call_parsers[self.tool_call_parser]}' instead." + f"The tool_call_parser '{cfg.tool_call_parser}' is deprecated. Please use '{deprecated_tool_call_parsers[cfg.tool_call_parser]}' instead." ) self._declare( "_handle_deprecated_args", - tool_call_parser=deprecated_tool_call_parsers[self.tool_call_parser], + tool_call_parser=deprecated_tool_call_parsers[cfg.tool_call_parser], ) # When user passes --enable-flashinfer-allreduce-fusion, enable with auto backend if ( - self.enable_flashinfer_allreduce_fusion - and self.flashinfer_allreduce_fusion_backend is None + cfg.enable_flashinfer_allreduce_fusion + and cfg.flashinfer_allreduce_fusion_backend is None ): logger.warning( "--enable-flashinfer-allreduce-fusion is deprecated. " @@ -4450,7 +4458,7 @@ def _handle_deprecated_args(self): self._declare("_handle_deprecated_args", **renamed) # --grpc-mode is a deprecated alias for --smg-grpc-mode. - if self.grpc_mode and not self.smg_grpc_mode: + if cfg.grpc_mode and not cfg.smg_grpc_mode: logger.warning( "--grpc-mode is deprecated and will be removed in a future " "version. Use --smg-grpc-mode for the legacy SMG gRPC server, " @@ -4466,7 +4474,7 @@ def _handle_deprecated_args(self): self.grpc_worker_threads = envs.SGLANG_GRPC_WORKER_THREADS.get() grpc_port_env = envs.SGLANG_GRPC_PORT.get() - if self.grpc_port is None and grpc_port_env is not None: + if cfg.grpc_port is None and grpc_port_env is not None: self._declare( "_handle_deprecated_args", grpc_port=grpc_port_env, @@ -4474,18 +4482,18 @@ def _handle_deprecated_args(self): # Legacy SMG defaults its port to --port + 10000. Derive/validate only # when gRPC is in use, so HTTP-only high ports don't fail validation. - legacy_grpc = self.smg_grpc_mode or self.grpc_mode - if legacy_grpc and self.grpc_port is None: + legacy_grpc = cfg.smg_grpc_mode or cfg.grpc_mode + if legacy_grpc and cfg.grpc_port is None: self._declare( "_handle_deprecated_args", - grpc_port=self.port + 10000, + grpc_port=cfg.port + 10000, ) - if self.grpc_port is not None: - if not (1 <= self.grpc_port <= 65535): + if cfg.grpc_port is not None: + if not (1 <= cfg.grpc_port <= 65535): raise ValueError( "--grpc-port / SGLANG_GRPC_PORT " - f"({self.grpc_port}) must be between 1 and 65535" + f"({cfg.grpc_port}) must be between 1 and 65535" ) if self.grpc_worker_threads < 1: raise ValueError( @@ -4495,41 +4503,41 @@ def _handle_deprecated_args(self): # Native gRPC is incompatible with launch paths it doesn't wire into. # Legacy takes precedence over grpc_port, keeping re-runs idempotent. - native_grpc = self.grpc_port is not None and not legacy_grpc - if self.sidecar_args is not None: - if self.sidecar is None: + native_grpc = cfg.grpc_port is not None and not legacy_grpc + if cfg.sidecar_args is not None: + if cfg.sidecar is None: raise ValueError("--sidecar-args requires --sidecar.") - if not isinstance(self.sidecar_args, list) or not all( - isinstance(arg, str) for arg in self.sidecar_args + if not isinstance(cfg.sidecar_args, list) or not all( + isinstance(arg, str) for arg in cfg.sidecar_args ): raise ValueError("--sidecar-args must be a JSON array of strings.") - if self.sidecar is not None: - if not self.sidecar.strip(): + if cfg.sidecar is not None: + if not cfg.sidecar.strip(): raise ValueError("--sidecar must not be empty.") if legacy_grpc: raise ValueError( "--sidecar requires SGLang's native gRPC server; " "it cannot be combined with --smg-grpc-mode/--grpc-mode." ) - if self.grpc_port is None: + if cfg.grpc_port is None: raise ValueError("--sidecar requires --grpc-port or SGLANG_GRPC_PORT.") if native_grpc: - if self.use_ray: + if cfg.use_ray: raise ValueError( "--grpc-port is not supported with --use-ray: the Ray " "serve launch path does not start the native gRPC server." ) - if self.encoder_only: + if cfg.encoder_only: raise ValueError( "--grpc-port is not supported with --encoder-only: " "encoder disaggregation uses its own server." ) - if self.tokenizer_worker_num > 1: + if cfg.tokenizer_worker_num > 1: raise ValueError( "Native gRPC does not yet support --tokenizer-worker-num > 1. " "Unset --grpc-port or set --tokenizer-worker-num 1." ) - if self.api_key or self.admin_api_key: + if cfg.api_key or cfg.admin_api_key: raise ValueError( "--grpc-port is incompatible with --api-key/--admin-api-key: " "the native gRPC listener bypasses HTTP auth middleware." @@ -4553,17 +4561,18 @@ def _handle_prefill_delayer_env_compat(self): ) def _handle_missing_default_values(self): - if self.tokenizer_path is None: + cfg = resolving_view(self) + if cfg.tokenizer_path is None: self._declare( "_handle_missing_default_values", - tokenizer_path=self.model_path, + tokenizer_path=cfg.model_path, ) - if self.served_model_name is None: + if cfg.served_model_name is None: self._declare( "_handle_missing_default_values", - served_model_name=self.model_path, + served_model_name=cfg.model_path, ) - if self.device is None: + if cfg.device is None: self._declare( "_handle_missing_default_values", device=get_device(), @@ -4571,14 +4580,14 @@ def _handle_missing_default_values(self): # strip device index from user if any (e.g. "cuda:0" -> "cuda") self._declare( "_handle_missing_default_values", - device=self.device.split(":")[0], + device=cfg.device.split(":")[0], ) - if self.random_seed is None: + if cfg.random_seed is None: self._declare( "_handle_missing_default_values", random_seed=random.randint(0, 1 << 30), ) - if self.mm_process_config is None: + if cfg.mm_process_config is None: self._declare("_handle_missing_default_values", mm_process_config={}) # Handle ModelScope model downloads @@ -4588,22 +4597,22 @@ def _handle_missing_default_values(self): # In speculative scenario: # - If `speculative_draft_model_quantization` is specified, the draft model uses this quantization method. # - Otherwise, the draft model defaults to the same quantization as the target model. - if self._speculative_draft_quantization_explicitly_set is None: + if cfg._speculative_draft_quantization_explicitly_set is None: self._declare( "_handle_missing_default_values", - _speculative_draft_quantization_explicitly_set=self.speculative_draft_model_quantization + _speculative_draft_quantization_explicitly_set=cfg.speculative_draft_model_quantization is not None, ) - if self.speculative_draft_model_quantization is None: + if cfg.speculative_draft_model_quantization is None: self._declare( "_handle_missing_default_values", - speculative_draft_model_quantization=self.quantization, + speculative_draft_model_quantization=cfg.quantization, ) # Resolve --quantization unquant before model config validation. Record # the explicit opt-out so later auto-detection does not re-enable # quantization. - if self.quantization == "unquant": + if cfg.quantization == "unquant": self._declare( "_handle_missing_default_values", quantization=None, @@ -4611,7 +4620,7 @@ def _handle_missing_default_values(self): self._quantization_explicitly_unset = True else: self._quantization_explicitly_unset = False - if self.speculative_draft_model_quantization == "unquant": + if cfg.speculative_draft_model_quantization == "unquant": self._declare( "_handle_missing_default_values", speculative_draft_model_quantization=None, @@ -4627,6 +4636,7 @@ def _handle_modelscope_paths(self): plain repo ID. That resolution lives in :func:`sglang.srt.speculative.spec_utils.load_token_map`. """ + cfg = resolving_view(self) ms_root = None ms_snapshot_download = None @@ -4656,42 +4666,43 @@ def _resolve_or_download( if os.path.exists(cached): return cached # Check user-specified download dir - if self.download_dir: - alt = os.path.join(self.download_dir, path) + if cfg.download_dir: + alt = os.path.join(cfg.download_dir, path) if os.path.exists(alt): return alt # Cache miss — download from ModelScope hub return ms_snapshot_download( path, - cache_dir=self.download_dir, + cache_dir=cfg.download_dir, revision=revision, **({"ignore_patterns": ignore_patterns} if ignore_patterns else {}), ) self._declare( "_handle_modelscope_paths", - model_path=_resolve_or_download(self.model_path, revision=self.revision), + model_path=_resolve_or_download(cfg.model_path, revision=cfg.revision), ) self._declare( "_handle_modelscope_paths", tokenizer_path=_resolve_or_download( - self.tokenizer_path, + cfg.tokenizer_path, ignore_patterns=["*.bin", "*.safetensors"], - revision=self.revision, + revision=cfg.revision, ), ) - if self.speculative_draft_model_path: + if cfg.speculative_draft_model_path: self._declare( "_handle_modelscope_paths", speculative_draft_model_path=_resolve_or_download( - self.speculative_draft_model_path, - revision=self.speculative_draft_model_revision or "main", + cfg.speculative_draft_model_path, + revision=cfg.speculative_draft_model_revision or "main", ), ) def _handle_hpu_backends(self): - if self.device == "hpu": + cfg = resolving_view(self) + if cfg.device == "hpu": self._declare( "_handle_hpu_backends", attention_backend="torch_native", @@ -4702,8 +4713,9 @@ def _handle_hpu_backends(self): ) def _handle_cpu_backends(self): - if self.device == "cpu": - if self.attention_backend is None: + cfg = resolving_view(self) + if cfg.device == "cpu": + if cfg.attention_backend is None: self._declare( "_handle_cpu_backends", attention_backend=( @@ -4723,21 +4735,23 @@ def _handle_hardware_runtime_validation(self): use_mlx() def _handle_npu_backends(self): - if self.device == "npu": + cfg = resolving_view(self) + if cfg.device == "npu": from sglang.srt.hardware_backend.npu.utils import set_default_server_args set_default_server_args(self) - current = self.cuda_graph_config.prefill.tc_compiler + current = cfg.cuda_graph_config.prefill.tc_compiler if current is not None and current != "eager": logger.warning( "At this moment Ascend platform only support prefill graph compilation with " "cuda_graph_config[prefill].tc_compiler='eager'." ) - self.cuda_graph_config.prefill.tc_compiler = "eager" + cfg.cuda_graph_config.prefill.tc_compiler = "eager" def _handle_mps_backends(self): - if self.device == "mps": + cfg = resolving_view(self) + if cfg.device == "mps": if not use_mlx(): self._declare( "_handle_mps_backends", @@ -4745,22 +4759,23 @@ def _handle_mps_backends(self): ) def _handle_xpu_backends(self): - if self.device == "xpu": + cfg = resolving_view(self) + if cfg.device == "xpu": # Decode graph is opt-in on XPU: unless the user explicitly set # --cuda-graph-backend-decode (or --cuda-graph-config), keep it # disabled so the default startup doesn't require graph capture. if (Phase.DECODE, "backend") not in self._cuda_graph_config_locked: - self.cuda_graph_config.decode.backend = Backend.DISABLED - elif self.cuda_graph_config.decode.backend not in ( + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + elif cfg.cuda_graph_config.decode.backend not in ( Backend.DISABLED, Backend.FULL, ): logger.warning( "XPU platform only supports decode backend 'full'; " "disabling unsupported decode backend '%s'.", - self.cuda_graph_config.decode.backend, + cfg.cuda_graph_config.decode.backend, ) - self.cuda_graph_config.decode.backend = Backend.DISABLED + cfg.cuda_graph_config.decode.backend = Backend.DISABLED # ------------------------------------------------------------------ # CUDA graph configuration resolution @@ -4771,10 +4786,11 @@ def _apply_inkling_prefill_cuda_graph_default(self): auto-disabled for this multimodal arch, and declarative model overrides materialize too late to steer cuda-graph resolution. Honors an explicit --cuda-graph-backend-prefill / --disable-prefill-cuda-graph.""" + cfg = resolving_view(self) if ( - self.cuda_graph_backend_prefill is not None - or self.disable_prefill_cuda_graph - or parse_connector_type(self.model_path) == ConnectorType.INSTANCE + cfg.cuda_graph_backend_prefill is not None + or cfg.disable_prefill_cuda_graph + or parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE ): return arch = self.get_model_config().hf_config.architectures[0] @@ -4788,9 +4804,10 @@ def _apply_inkling_prefill_cuda_graph_default(self): ) def _apply_muse_glimmer_prefill_cuda_graph_max_bs_default(self): + cfg = resolving_view(self) if ( - self.cuda_graph_max_bs_prefill is not None - or parse_connector_type(self.model_path) == ConnectorType.INSTANCE + cfg.cuda_graph_max_bs_prefill is not None + or parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE ): return arch = self.get_model_config().hf_config.architectures[0] @@ -4801,6 +4818,7 @@ def _apply_muse_glimmer_prefill_cuda_graph_max_bs_default(self): ) def _handle_cuda_graph_config(self): + cfg = resolving_view(self) from sglang.srt.arg_groups.kimi_k3_hook import disable_kimi_k3_symm_mem self._parse_cuda_graph_config() @@ -4814,7 +4832,7 @@ def _handle_cuda_graph_config(self): # Warn on the final resolved config (not inside the compat cascade — # that path is skipped when the user explicitly sets the backend, # which is the only way to get 'full' for prefill today). - if self.cuda_graph_config.prefill.backend == Backend.FULL: + if cfg.cuda_graph_config.prefill.backend == Backend.FULL: logger.warning( "cuda_graph_config[prefill].backend='full' is experimental. " "Use breakable or tc_piecewise for production workloads." @@ -4822,16 +4840,17 @@ def _handle_cuda_graph_config(self): def _apply_deepep_adjustments(self): """Config adjustments required by the DeepEP a2a backend.""" + cfg = resolving_view(self) if resolved_view(self).moe_a2a_backend != "deepep": return # Non-multiple-of-8 prefill buckets can hang DeepEP a2a capture under # breakable CUDA graph - if self.cuda_graph_config.prefill.backend == Backend.BREAKABLE: - bs = self.cuda_graph_config.prefill.bs + if cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE: + bs = cfg.cuda_graph_config.prefill.bs if bs is None: # 2048 = documented prefill default; max_bs unresolved here. - max_bs = self.cuda_graph_config.prefill.max_bs or 2048 + max_bs = cfg.cuda_graph_config.prefill.max_bs or 2048 bs = self._generate_prefill_cuda_graph_batch_sizes(max_bs) aligned = sorted({((b + 7) // 8) * 8 for b in bs}) if aligned != sorted(bs): @@ -4841,8 +4860,8 @@ def _apply_deepep_adjustments(self): sorted(bs), aligned, ) - self.cuda_graph_config.prefill.bs = aligned - self.cuda_graph_config.prefill.max_bs = aligned[-1] + cfg.cuda_graph_config.prefill.bs = aligned + cfg.cuda_graph_config.prefill.max_bs = aligned[-1] def _parse_cuda_graph_config(self): """Resolve cuda_graph_config from explicit JSON, per-phase @@ -4853,7 +4872,8 @@ def _parse_cuda_graph_config(self): auto-disable cascade respects this lock (the old --enforce-piecewise-cuda-graph semantics generalized). """ - raw_input = self.cuda_graph_config + cfg = resolving_view(self) + raw_input = cfg.cuda_graph_config if isinstance(raw_input, CudaGraphConfig): explicit_input = raw_input.to_dict() else: @@ -4866,36 +4886,36 @@ def _set(phase: str, key: str, value: Any) -> None: locked.add((phase, key)) # ---- Legacy global flags (lowest precedence above defaults) ---- - if self.disable_cuda_graph: + if cfg.disable_cuda_graph: _set(Phase.DECODE, "backend", Backend.DISABLED) _set(Phase.PREFILL, "backend", Backend.DISABLED) # ---- Boolean per-phase off-switches ---- # Below the explicit backend selectors so --cuda-graph-backend-* # wins if both are given. - if self.disable_prefill_cuda_graph: + if cfg.disable_prefill_cuda_graph: _set(Phase.PREFILL, "backend", Backend.DISABLED) - if self.disable_decode_cuda_graph: + if cfg.disable_decode_cuda_graph: _set(Phase.DECODE, "backend", Backend.DISABLED) # ---- Per-phase convenience flags ---- - if self.cuda_graph_backend_decode is not None: - _set(Phase.DECODE, "backend", self.cuda_graph_backend_decode) - if self.cuda_graph_backend_prefill is not None: - _set(Phase.PREFILL, "backend", self.cuda_graph_backend_prefill) - if self.cuda_graph_max_bs_decode is not None: - _set(Phase.DECODE, "max_bs", self.cuda_graph_max_bs_decode) - if self.cuda_graph_max_bs_prefill is not None: - _set(Phase.PREFILL, "max_bs", self.cuda_graph_max_bs_prefill) - if self.cuda_graph_bs_decode is not None: - _set(Phase.DECODE, "bs", self.cuda_graph_bs_decode) - if self.cuda_graph_bs_prefill is not None: - _set(Phase.PREFILL, "bs", self.cuda_graph_bs_prefill) - if self.cuda_graph_tc_compiler is not None: + if cfg.cuda_graph_backend_decode is not None: + _set(Phase.DECODE, "backend", cfg.cuda_graph_backend_decode) + if cfg.cuda_graph_backend_prefill is not None: + _set(Phase.PREFILL, "backend", cfg.cuda_graph_backend_prefill) + if cfg.cuda_graph_max_bs_decode is not None: + _set(Phase.DECODE, "max_bs", cfg.cuda_graph_max_bs_decode) + if cfg.cuda_graph_max_bs_prefill is not None: + _set(Phase.PREFILL, "max_bs", cfg.cuda_graph_max_bs_prefill) + if cfg.cuda_graph_bs_decode is not None: + _set(Phase.DECODE, "bs", cfg.cuda_graph_bs_decode) + if cfg.cuda_graph_bs_prefill is not None: + _set(Phase.PREFILL, "bs", cfg.cuda_graph_bs_prefill) + if cfg.cuda_graph_tc_compiler is not None: # Written to both phases so the value is in place when TC_PIECEWISE # decode is implemented; today decode ignores it. - _set(Phase.DECODE, "tc_compiler", self.cuda_graph_tc_compiler) - _set(Phase.PREFILL, "tc_compiler", self.cuda_graph_tc_compiler) + _set(Phase.DECODE, "tc_compiler", cfg.cuda_graph_tc_compiler) + _set(Phase.PREFILL, "tc_compiler", cfg.cuda_graph_tc_compiler) # ---- Explicit JSON config (highest precedence) ---- for phase, phase_config in explicit_input.items(): @@ -4917,6 +4937,7 @@ def _apply_cuda_graph_compatibility(self): prefill backend (this folds in the old --enforce-piecewise-cuda-graph contract). """ + cfg = resolving_view(self) if (Phase.PREFILL, "backend") in self._cuda_graph_config_locked: return @@ -4925,7 +4946,7 @@ def _apply_cuda_graph_compatibility(self): # there instead. Archs also on the breakable allowlist keep it -- # this runs first, so piecewise would otherwise silently win. if ( - self.cuda_graph_config.prefill.backend == Backend.BREAKABLE + cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE and self.get_model_config().is_multimodal_piecewise_cuda_graph_supported and not self.get_model_config().is_multimodal_breakable_cuda_graph_supported # Keep trtllm_mla on the preferred breakable path, which now serves @@ -4936,27 +4957,29 @@ def _apply_cuda_graph_compatibility(self): "Using tc_piecewise CUDA graph for validated multimodal " "decoder prefill." ) - self.cuda_graph_config.prefill.backend = Backend.TC_PIECEWISE + cfg.cuda_graph_config.prefill.backend = Backend.TC_PIECEWISE - if self.cuda_graph_config.prefill.backend == Backend.TC_PIECEWISE: + if cfg.cuda_graph_config.prefill.backend == Backend.TC_PIECEWISE: self._disable_tc_piecewise_cudagraph_if_incompatible() - elif self.cuda_graph_config.prefill.backend == Backend.BREAKABLE: + elif cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE: self._disable_breakable_cudagraph_if_incompatible() - elif self.cuda_graph_config.prefill.backend == Backend.FULL: + elif cfg.cuda_graph_config.prefill.backend == Backend.FULL: self._disable_full_prefill_cudagraph_if_incompatible() def _apply_cuda_graph_disaggregation_roles(self): - if self.disaggregation_mode == "prefill": + cfg = resolving_view(self) + if cfg.disaggregation_mode == "prefill": if (Phase.DECODE, "backend") not in self._cuda_graph_config_locked: - self.cuda_graph_config.decode.backend = Backend.DISABLED - elif self.disaggregation_mode == "decode": + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + elif cfg.disaggregation_mode == "decode": if (Phase.PREFILL, "backend") not in self._cuda_graph_config_locked: - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED def _disable_tc_piecewise_cudagraph_if_incompatible(self): """TcPiecewise (torch.compile + piecewise) is incompatible with these configurations. Most are torch.compile / dynamo limitations. """ + cfg = resolving_view(self) rules = [ ( @@ -4964,8 +4987,8 @@ def _disable_tc_piecewise_cudagraph_if_incompatible(self): lambda: self.get_model_config().is_piecewise_cuda_graph_disabled_model, ), ("DP attention", lambda: self._resolved().enable_dp_attention), - ("full torch.compile mode", lambda: self.enable_torch_compile), - ("pipeline parallelism (pp_size > 1)", lambda: self.pp_size > 1), + ("full torch.compile mode", lambda: cfg.enable_torch_compile), + ("pipeline parallelism (pp_size > 1)", lambda: cfg.pp_size > 1), ( "non-CUDA hardware (HIP/NPU/CPU/MPS/XPU)", lambda: is_hip() or is_npu() or is_cpu() or is_mps() or is_xpu(), @@ -4981,7 +5004,7 @@ def _disable_tc_piecewise_cudagraph_if_incompatible(self): ), # Dynamo blocks LoRA under tc_piecewise (per-batch LoRABatchInfo # rebinds break guards); breakable/full support LoRA. - ("LoRA", lambda: bool(self.lora_paths) or self.enable_lora), + ("LoRA", lambda: bool(cfg.lora_paths) or cfg.enable_lora), ( "multimodal model", lambda: self.get_model_config().is_multimodal @@ -4989,50 +5012,51 @@ def _disable_tc_piecewise_cudagraph_if_incompatible(self): ), ( "GGUF quantization", - lambda: self.load_format == "gguf" + lambda: cfg.load_format == "gguf" or resolved_view(self).quantization == "gguf" - or check_gguf_file(self.model_path), + or check_gguf_file(cfg.model_path), ), - ("DLLM (diffusion LLM)", lambda: self.dllm_algorithm is not None), + ("DLLM (diffusion LLM)", lambda: cfg.dllm_algorithm is not None), ( "CPU offload / hierarchical cache", - lambda: self.cpu_offload_gb > 0 or self.enable_hierarchical_cache, + lambda: cfg.cpu_offload_gb > 0 or cfg.enable_hierarchical_cache, ), ( "deterministic inference", - lambda: self.enable_deterministic_inference, + lambda: cfg.enable_deterministic_inference, ), - ("PD disaggregation", lambda: self.disaggregation_mode != "null"), - ("symmetric memory", lambda: self.enable_symm_mem), + ("PD disaggregation", lambda: cfg.disaggregation_mode != "null"), + ("symmetric memory", lambda: cfg.enable_symm_mem), ( "expert distribution recorder", - lambda: self.enable_eplb - or self.expert_distribution_recorder_mode is not None, + lambda: cfg.enable_eplb + or cfg.expert_distribution_recorder_mode is not None, ), ( "context parallel (attn_cp_size > 1)", lambda: self._resolved().attn_cp_size > 1, ), - ("CUDA graph debug mode", lambda: self.debug_cuda_graph), + ("CUDA graph debug mode", lambda: cfg.debug_cuda_graph), ( "DSA prefill context parallelism", - lambda: self.enable_dsa_prefill_context_parallel, + lambda: cfg.enable_dsa_prefill_context_parallel, ), # Capture builds a dummy extend forward with attn_dcp_metadata=None. ( "decode context parallel (dcp_size > 1)", - lambda: self.dcp_size > 1, + lambda: cfg.dcp_size > 1, ), ] for _name, predicate in rules: if predicate(): - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED def _disable_breakable_cudagraph_if_incompatible(self): """Breakable (segmented capture, no torch.compile). Breakable enforces memory-saver rejection in its own __init__; config-time rules can be added here as they're discovered. """ + cfg = resolving_view(self) from sglang.srt.configs.model_config import is_deepseek_v4 from sglang.srt.layers.cp.bcg import supports_prefill_cp_bcg @@ -5052,12 +5076,12 @@ def _disable_breakable_cudagraph_if_incompatible(self): # Capture builds a dummy extend forward with attn_dcp_metadata=None. ( "decode context parallel (dcp_size > 1)", - lambda: self.dcp_size > 1, + lambda: cfg.dcp_size > 1, ), # TBO capture is unsupported. ( "two-batch overlap", - lambda: self.enable_two_batch_overlap, + lambda: cfg.enable_two_batch_overlap, ), ( "unvalidated a2a backend", @@ -5078,11 +5102,12 @@ def _disable_breakable_cudagraph_if_incompatible(self): "disabling prefill CUDA graph.", name, ) - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED return def _disable_full_prefill_cudagraph_if_incompatible(self): """Full prefill CG: empty rule list today; see the experimental warning.""" + cfg = resolving_view(self) rules = [] for name, predicate in rules: if predicate(): @@ -5091,7 +5116,7 @@ def _disable_full_prefill_cudagraph_if_incompatible(self): "disabling prefill CUDA graph.", name, ) - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED return def _disable_prefill_cuda_graph_for_deepseek_trtllm_mla(self): @@ -5100,10 +5125,11 @@ def _disable_prefill_cuda_graph_for_deepseek_trtllm_mla(self): breakable) trtllm_mla falls back to FlashAttention for prefill and regresses performance, so disable whichever prefill graph backend is in effect. """ + cfg = resolving_view(self) if (Phase.PREFILL, "backend") in self._cuda_graph_config_locked: return - if self.cuda_graph_config.prefill.backend == Backend.DISABLED: + if cfg.cuda_graph_config.prefill.backend == Backend.DISABLED: return if ( "DeepseekV3ForCausalLM" @@ -5118,15 +5144,16 @@ def _disable_prefill_cuda_graph_for_deepseek_trtllm_mla(self): "the trtllm_mla attention backend (a captured prefill graph forces a " "FlashAttention fallback that regresses prefill). Set the prefill cuda graph " "backend explicitly (e.g. --cuda-graph-backend-prefill tc_piecewise) to override.", - self.cuda_graph_config.prefill.backend, + cfg.cuda_graph_config.prefill.backend, ) - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED def _validate_cuda_graph_config(self): - if self.cuda_graph_config is None: + cfg = resolving_view(self) + if cfg.cuda_graph_config is None: return for phase in Phase.ALL: - backend = getattr(self.cuda_graph_config, phase).backend + backend = getattr(cfg.cuda_graph_config, phase).backend if backend not in ALLOWED_BACKENDS_PER_PHASE[phase]: raise ValueError( f"--cuda-graph-config[{phase}].backend={backend!r} not allowed; " @@ -5141,22 +5168,23 @@ def _handle_multi_item_scoring(self): changing it silently could surprise users who intentionally picked a non-flashinfer backend. """ - if not self.enable_mis: + cfg = resolving_view(self) + if not cfg.enable_mis: return - if self.cuda_graph_config.decode.backend != Backend.DISABLED: + if cfg.cuda_graph_config.decode.backend != Backend.DISABLED: logger.warning("CUDA graph is disabled because --enable-mis is set.") - self.cuda_graph_config.decode.backend = Backend.DISABLED - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED - if not self.disable_radix_cache: + if not cfg.disable_radix_cache: logger.warning("Radix cache is disabled because --enable-mis is set.") self._declare( "_handle_multi_item_scoring", disable_radix_cache=True, ) - if self.chunked_prefill_size != -1: + if cfg.chunked_prefill_size != -1: logger.warning("Chunked prefill is disabled because --enable-mis is set.") self._declare( "_handle_multi_item_scoring", @@ -5194,14 +5222,15 @@ def _handle_gpu_memory_settings(self, gpu_mem): The coefficient 1.5 is a heuristic value, in the future, we can do better estimation by looking at the model types, hidden sizes or even do a dummy run. """ - decode_cuda_graph_config = self.cuda_graph_config.decode - prefill_cuda_graph_config = self.cuda_graph_config.prefill + cfg = resolving_view(self) + decode_cuda_graph_config = cfg.cuda_graph_config.decode + prefill_cuda_graph_config = cfg.cuda_graph_config.prefill if gpu_mem is not None: if gpu_mem < 20 * 1024: # T4, 4080 # (chunked_prefill_size 2k, max_bs 8) - if self.chunked_prefill_size is None: + if cfg.chunked_prefill_size is None: self._declare( "_handle_gpu_memory_settings", chunked_prefill_size=2048, @@ -5211,59 +5240,59 @@ def _handle_gpu_memory_settings(self, gpu_mem): elif gpu_mem < 35 * 1024: # A10, 4090, 5090 # (chunked_prefill_size 2k, max_bs 24 if tp < 4 else 80) - if self.chunked_prefill_size is None: + if cfg.chunked_prefill_size is None: self._declare( "_handle_gpu_memory_settings", chunked_prefill_size=2048, ) if decode_cuda_graph_config.max_bs is None: - if self.tp_size < 4: + if cfg.tp_size < 4: decode_cuda_graph_config.max_bs = 24 else: decode_cuda_graph_config.max_bs = 80 elif gpu_mem < 60 * 1024: # A100 (40GB), L40, # (chunked_prefill_size 4k, max_bs 32 if tp < 4 else 160) - if self.chunked_prefill_size is None: + if cfg.chunked_prefill_size is None: self._declare( "_handle_gpu_memory_settings", chunked_prefill_size=4096, ) if decode_cuda_graph_config.max_bs is None: - if self.tp_size < 4: + if cfg.tp_size < 4: decode_cuda_graph_config.max_bs = 32 else: decode_cuda_graph_config.max_bs = 160 elif gpu_mem < 90 * 1024: # H100, A100 # (chunked_prefill_size 8k, max_bs 256 if tp < 4 else 512) - if self.chunked_prefill_size is None: + if cfg.chunked_prefill_size is None: self._declare( "_handle_gpu_memory_settings", chunked_prefill_size=8192, ) if decode_cuda_graph_config.max_bs is None: - if self.tp_size < 4: + if cfg.tp_size < 4: decode_cuda_graph_config.max_bs = 256 else: decode_cuda_graph_config.max_bs = 512 elif gpu_mem < 160 * 1024: # H20, H200 # (chunked_prefill_size 8k, max_bs 256 if tp < 4 else 512) - if self.chunked_prefill_size is None: + if cfg.chunked_prefill_size is None: self._declare( "_handle_gpu_memory_settings", chunked_prefill_size=8192, ) if decode_cuda_graph_config.max_bs is None: - if self.tp_size < 4: + if cfg.tp_size < 4: decode_cuda_graph_config.max_bs = 256 else: decode_cuda_graph_config.max_bs = 512 else: # B200, MI300 # (chunked_prefill_size 16k, max_bs 512) - if self.chunked_prefill_size is None: + if cfg.chunked_prefill_size is None: self._declare( "_handle_gpu_memory_settings", chunked_prefill_size=16384, @@ -5272,7 +5301,7 @@ def _handle_gpu_memory_settings(self, gpu_mem): decode_cuda_graph_config.max_bs = 512 else: # Fallback defaults when gpu_mem is None - if self.chunked_prefill_size is None: + if cfg.chunked_prefill_size is None: self._declare( "_handle_gpu_memory_settings", chunked_prefill_size=4096, @@ -5281,7 +5310,7 @@ def _handle_gpu_memory_settings(self, gpu_mem): decode_cuda_graph_config.max_bs = 160 # Set cuda graph batch sizes - if self.device != "cpu": + if cfg.device != "cpu": if decode_cuda_graph_config.bs is None: decode_cuda_graph_config.bs = ( self._generate_decode_cuda_graph_batch_sizes( @@ -5303,34 +5332,34 @@ def _handle_gpu_memory_settings(self, gpu_mem): # to generate decode_cuda_graph_config.bs self._declare( "_handle_gpu_memory_settings", - torch_compile_max_bs=self.torch_compile_max_bs + torch_compile_max_bs=cfg.torch_compile_max_bs or decode_cuda_graph_config.max_bs, ) decode_cuda_graph_config.bs = self._generate_cpu_graph_batch_sizes() assert ( - self.torch_compile_max_bs > 0 + cfg.torch_compile_max_bs > 0 ), "cuda_graph_config[decode].bs should contain positive batch sizes" - decode_cuda_graph_config.max_bs = self.torch_compile_max_bs + decode_cuda_graph_config.max_bs = cfg.torch_compile_max_bs if prefill_cuda_graph_config.max_bs is None: # Refer to pr #15927, by default we set the prefill max_bs to the chunked prefill size. # For MLA backend, the introduction of piecewise cuda graph will influence the kernel dispatch difference compared to the original mode. # To avoid the performance regression, we set max_bs to 2048 by default. if not self.use_mla_backend(): - prefill_cuda_graph_config.max_bs = self.chunked_prefill_size + prefill_cuda_graph_config.max_bs = cfg.chunked_prefill_size else: prefill_cuda_graph_config.max_bs = 2048 # If max_total_tokens is set, cap prefill max_bs to not exceed max_total_tokens. - if self.max_total_tokens is not None: + if cfg.max_total_tokens is not None: prefill_cuda_graph_config.max_bs = min( - prefill_cuda_graph_config.max_bs, self.max_total_tokens + prefill_cuda_graph_config.max_bs, cfg.max_total_tokens ) # For Llama2 series models, max_bs is limited to 4096. # TODO(yuwei): remove this after the issue is fixed - if "llama-2" in self.model_path.lower(): + if "llama-2" in cfg.model_path.lower(): prefill_cuda_graph_config.max_bs = min( prefill_cuda_graph_config.max_bs, 4096 ) @@ -5342,31 +5371,29 @@ def _handle_gpu_memory_settings(self, gpu_mem): ) ) - if self.mem_fraction_static is None: + if cfg.mem_fraction_static is None: if self.post_capture_kv_sizing_planned(): # Post-capture sizing measures free memory after graph capture, so # skip the graph/activation reserve; keep only the floor + parallel slack. reserved_mem = 1536 - reserved_mem += self.tp_size * self.pp_size / 8 * 1024 + reserved_mem += cfg.tp_size * cfg.pp_size / 8 * 1024 else: # Tokens the activation working set scales with (per serving mode). - if self.disaggregation_mode == "decode": + if cfg.disaggregation_mode == "decode": running_requests = ( - self.max_running_requests - or decode_cuda_graph_config.max_bs - or 1 + cfg.max_running_requests or decode_cuda_graph_config.max_bs or 1 ) - draft_tokens = self.speculative_num_draft_tokens or 1 + draft_tokens = cfg.speculative_num_draft_tokens or 1 activation_tokens = max(running_requests * draft_tokens, 2048) - elif self.chunked_prefill_size > 0: - activation_tokens = max(self.chunked_prefill_size, 2048) + elif cfg.chunked_prefill_size > 0: + activation_tokens = max(cfg.chunked_prefill_size, 2048) else: - activation_tokens = max(self.max_prefill_tokens, 2048) + activation_tokens = max(cfg.max_prefill_tokens, 2048) # Constant meta data (e.g., from attention backend) + activation slack. reserved_mem = 512 reserved_mem += activation_tokens * 1.5 # Some adjustments for large parallel size - reserved_mem += self.tp_size * self.pp_size / 8 * 1024 + reserved_mem += cfg.tp_size * cfg.pp_size / 8 * 1024 reserved_mem += self.reserve_for_graph_mb() if gpu_mem is not None and gpu_mem > 60 * 1024: reserved_mem = max(reserved_mem, 10 * 1024) @@ -5389,14 +5416,14 @@ def _handle_gpu_memory_settings(self, gpu_mem): model_config = self.get_model_config() if ( model_config.is_multimodal - and not self.language_only - and not self.language_model_only - and self.disaggregation_mode != "decode" + and not cfg.language_only + and not cfg.language_model_only + and cfg.disaggregation_mode != "decode" ): self.adjust_mem_fraction_for_vlm(model_config) # If symm mem is enabled and prealloc size is not set, set it to 4GB - if self.enable_symm_mem and not envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.is_set(): + if cfg.enable_symm_mem and not envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.is_set(): envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.set(4) logger.warning( "Symmetric memory is enabled, setting symmetric memory prealloc size to 4GB as default." @@ -5407,41 +5434,42 @@ def post_capture_kv_sizing_planned(self) -> bool: """Whether the mem_fraction heuristic may skip the graph reserve; must be False for any config the runtime won't post-capture-size, else it gets an under-reserved fraction.""" + cfg = resolving_view(self) # use_mla_backend is a method at args time but ModelRunner overwrites it # with a bool on global_server_args (see the FIXME there) -- handle both. use_mla = self.use_mla_backend mla_enabled = use_mla() if callable(use_mla) else use_mla if not envs.SGLANG_ENABLE_POST_CAPTURE_KV_SIZING.get(): return False - if self.device != "cuda": + if cfg.device != "cuda": return False - if self.dcp_size != 1: + if cfg.dcp_size != 1: return False if mla_enabled: return False - if self.kv_cache_dtype == "fp4_e2m1": + if cfg.kv_cache_dtype == "fp4_e2m1": return False - if self.prefill_only_disable_kv_cache: + if cfg.prefill_only_disable_kv_cache: return False - if self.enable_memory_saver: + if cfg.enable_memory_saver: return False if envs.SGLANG_MOONCAKE_CUSTOM_MEM_POOL.get() is not None: return False if ( - self.disaggregation_mode != "prefill" - and self.cuda_graph_config.decode.backend == Backend.DISABLED + cfg.disaggregation_mode != "prefill" + and cfg.cuda_graph_config.decode.backend == Backend.DISABLED ): return False - if self.disaggregation_mode != "decode": - prefill_cfg = self.cuda_graph_config.prefill + if cfg.disaggregation_mode != "decode": + prefill_cfg = cfg.cuda_graph_config.prefill # We can only skip eager activation headroom when the largest # prefill forward batch size is already graph-captured. Otherwise, # an eager forward will need more memory and lead to OOM. if ( prefill_cfg.backend == Backend.DISABLED - or self.chunked_prefill_size <= 0 + or cfg.chunked_prefill_size <= 0 or self.max_prefill_buffer_tokens() > max(prefill_cfg.bs or (0,)) ): return False @@ -5476,28 +5504,29 @@ def pre_capture_activation_reserve_mb(self, gpu_mem: Optional[float]) -> float: return reserved_mem def reserve_for_graph_mb(self) -> float: - decode_cuda_graph_config = self.cuda_graph_config.decode - prefill_cuda_graph_config = self.cuda_graph_config.prefill + cfg = resolving_view(self) + decode_cuda_graph_config = cfg.cuda_graph_config.decode + prefill_cuda_graph_config = cfg.cuda_graph_config.prefill reserved_mem = 0.0 if ( - self.disaggregation_mode != "prefill" + cfg.disaggregation_mode != "prefill" and decode_cuda_graph_config.backend != Backend.DISABLED ): reserved_mem += decode_cuda_graph_config.max_bs * 2 if ( self._resolved().enable_dp_attention - and self.disaggregation_mode != "prefill" + and cfg.disaggregation_mode != "prefill" ): # DP attention needs more padding for some operations, and much more for large # cuda graph max bs (torch allocator / implementation inefficiencies). - reserved_mem += decode_cuda_graph_config.max_bs * self.dp_size * 3 + reserved_mem += decode_cuda_graph_config.max_bs * cfg.dp_size * 3 if decode_cuda_graph_config.max_bs > 300: - reserved_mem += decode_cuda_graph_config.max_bs * self.dp_size * 1.5 + reserved_mem += decode_cuda_graph_config.max_bs * cfg.dp_size * 1.5 if ( - self.disaggregation_mode != "decode" + cfg.disaggregation_mode != "decode" and prefill_cuda_graph_config.backend != Backend.DISABLED ): if not self.use_mla_backend(): @@ -5521,9 +5550,10 @@ def reserve_for_deepep_a2a_mb(self) -> float: # DeepEP all-to-all buffers captured in the decode graph are real extra # allocations, reserved on top of the floor. - decode_cuda_graph_config = self.cuda_graph_config.decode + cfg = resolving_view(self) + decode_cuda_graph_config = cfg.cuda_graph_config.decode if ( - self.disaggregation_mode != "prefill" + cfg.disaggregation_mode != "prefill" and decode_cuda_graph_config.backend != Backend.DISABLED and resolved_view(self).moe_a2a_backend == "deepep" ): @@ -5535,10 +5565,11 @@ def _generate_decode_cuda_graph_batch_sizes(self, max_bs: int): Generate the list of batch sizes for CUDA graph capture based on max_bs. This integrates the logic from cuda_graph_runner.py. """ + cfg = resolving_view(self) # Handle disable_cuda_graph_padding as the first condition for both spec and non-spec - if self.disable_cuda_graph_padding: + if cfg.disable_cuda_graph_padding: capture_bs = list(range(1, max_bs + 1)) - elif self.speculative_algorithm is None: + elif cfg.speculative_algorithm is None: # Normal case: capture_bs = ( [1, 2, 4, 8, 12] @@ -5567,19 +5598,20 @@ def _generate_cpu_graph_batch_sizes(self): """ Generate the list of batch sizes for CPU graph capture based on torch_compile_max_bs. """ - if self.disable_cuda_graph_padding: - capture_bs = list(range(1, self.torch_compile_max_bs + 1)) + cfg = resolving_view(self) + if cfg.disable_cuda_graph_padding: + capture_bs = list(range(1, cfg.torch_compile_max_bs + 1)) else: capture_bs = sorted( set().union( range(1, 17), range(18, 31, 2), range(32, 81, 4), - range(84, self.torch_compile_max_bs + 1, 8), - {self.torch_compile_max_bs}, + range(84, cfg.torch_compile_max_bs + 1, 8), + {cfg.torch_compile_max_bs}, ) ) - capture_bs = [bs for bs in capture_bs if bs <= self.torch_compile_max_bs] + capture_bs = [bs for bs in capture_bs if bs <= cfg.torch_compile_max_bs] return capture_bs @@ -5633,12 +5665,13 @@ def _validate_hisparse_kv_cache_dtype(self): validate_hisparse_kv_cache_dtype(self) def _handle_model_specific_adjustments(self): + cfg = resolving_view(self) from sglang.srt.configs.model_config import ( get_mimo_v2_fused_qkv_expected_tp_size, is_deepseek_dsa, ) - if self.enable_deterministic_inference: + if cfg.enable_deterministic_inference: self._declare( "_handle_model_specific_adjustments", enforce_disable_flashinfer_allreduce_fusion=True, @@ -5648,7 +5681,7 @@ def _handle_model_specific_adjustments(self): "_handle_model_specific_adjustments", uses_mamba_radix_cache=False, ) - if parse_connector_type(self.model_path) == ConnectorType.INSTANCE: + if parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE: # No model overrides for an instance connector: no hf_config to # key them on. return @@ -5659,22 +5692,22 @@ def _handle_model_specific_adjustments(self): if model_arch == "InternS2MobiusForConditionalGeneration": unsupported = [] - if self.pp_size != 1: + if cfg.pp_size != 1: unsupported.append("pipeline parallelism (--pp-size must be 1)") - if self.ep_size != 1: + if cfg.ep_size != 1: unsupported.append("expert parallelism (--ep-size must be 1)") if unsupported: raise ValueError( "Intern-S2-Mobius does not support: " + "; ".join(unsupported) + "." ) - if self.enable_dsa_cache_layer_split and not is_deepseek_dsa(hf_config): + if cfg.enable_dsa_cache_layer_split and not is_deepseek_dsa(hf_config): raise ValueError( "--enable-dsa-cache-layer-split is only supported for DSA " "(DeepSeek Sparse Attention) models." ) - if self.enable_cp_decode_attn_tp: + if cfg.enable_cp_decode_attn_tp: from sglang.srt.layers.cp.cp_decode_attn_tp import ( CP_DECODE_ATTN_TP_SUPPORTED_ARCHS, ) @@ -5759,7 +5792,7 @@ def _handle_model_specific_adjustments(self): index_topk_freq = getattr(hf_config, "index_topk_freq", 1) or 1 index_topk_pattern = getattr(hf_config, "index_topk_pattern", None) - if self.enable_two_batch_overlap and ( + if cfg.enable_two_batch_overlap and ( index_topk_freq > 1 or (index_topk_pattern is not None and "S" in index_topk_pattern) ): @@ -5772,17 +5805,17 @@ def _handle_model_specific_adjustments(self): ) if not is_npu() and not is_xpu(): # CUDA or ROCm GPU - if self.enable_prefill_cp: + if cfg.enable_prefill_cp: # The DSA CP field declarations moved to the override # registry (arg_groups/overrides.py: # _deepseek_family_overrides). - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED else: # Pure TP and partial DP Attention mode is active for DSA, logging a warning - if self.dp_size < self.tp_size: + if cfg.dp_size < cfg.tp_size: logger.warning( - f"DSA with TP mode is active, dp_size={self.dp_size}, tp_size={self.tp_size}, " - f"attn_tp_size={self.tp_size}, attention weights will be sharded across {self.tp_size} ranks." + f"DSA with TP mode is active, dp_size={cfg.dp_size}, tp_size={cfg.tp_size}, " + f"attn_tp_size={cfg.tp_size}, attention weights will be sharded across {cfg.tp_size} ranks." ) # The DSA page-size selection moved to the override registry @@ -5796,15 +5829,15 @@ def _handle_model_specific_adjustments(self): ) self._set_default_dsa_backends(major) - if self.enable_prefill_cp: + if cfg.enable_prefill_cp: assert ( - self.disaggregation_mode != "decode" + cfg.disaggregation_mode != "decode" ), "CP is only supported for prefill when PD disaggregation, please remove --enable-prefill-cp." if ( - self.enable_dsa_cache_layer_split - and self.disaggregation_mode != "prefill" + cfg.enable_dsa_cache_layer_split + and cfg.disaggregation_mode != "prefill" ): - if self.disaggregation_mode == "decode": + if cfg.disaggregation_mode == "decode": raise ValueError( "--enable-dsa-cache-layer-split is not supported on " "decode workers. This flag is a prefill-CP " @@ -5816,8 +5849,8 @@ def _handle_model_specific_adjustments(self): "prefill workers. Non-PD workers also run decode and " "require ordinary local decode cache semantics." ) - if self.enable_dsa_cache_layer_split and ( - not self.enable_prefill_cp or self.cp_strategy != "interleave" + if cfg.enable_dsa_cache_layer_split and ( + not cfg.enable_prefill_cp or cfg.cp_strategy != "interleave" ): raise ValueError( "--enable-dsa-cache-layer-split requires " @@ -5829,17 +5862,17 @@ def _handle_model_specific_adjustments(self): # transfer path. mori/nixl support is a temporary limitation # and will be added later by the community. if ( - self.enable_dsa_cache_layer_split - and self.disaggregation_transfer_backend != "mooncake" + cfg.enable_dsa_cache_layer_split + and cfg.disaggregation_transfer_backend != "mooncake" ): raise ValueError( "--enable-dsa-cache-layer-split currently only supports " "the mooncake transfer backend (mooncake / mooncake_tcp). " f"Got --disaggregation-transfer-backend " - f"{self.disaggregation_transfer_backend!r}. mori/nixl " + f"{cfg.disaggregation_transfer_backend!r}. mori/nixl " "support will be added later by the community." ) - if self.enable_dsa_cache_layer_split and self.pp_size > 1: + if cfg.enable_dsa_cache_layer_split and cfg.pp_size > 1: raise ValueError( "--enable-dsa-cache-layer-split is not supported with " "pipeline parallelism (pp_size > 1) yet. It requires " @@ -5849,7 +5882,7 @@ def _handle_model_specific_adjustments(self): else: # DeepSeek V3/R1/V3.1 - if self.cuda_graph_config.prefill.backend != Backend.DISABLED: + if cfg.cuda_graph_config.prefill.backend != Backend.DISABLED: logger.info("Piecewise CUDA graph is enabled, use MLA for prefill.") # The sm100 trtllm_mla fill moved to the override registry @@ -5858,8 +5891,8 @@ def _handle_model_specific_adjustments(self): # MLA prefill CP auto-config: the field declarations moved to # the override registry (arg_groups/overrides.py: # _deepseek_family_overrides). - if self.enable_prefill_cp and self.use_mla_backend(): - self.cuda_graph_config.prefill.backend = Backend.DISABLED + if cfg.enable_prefill_cp and self.use_mla_backend(): + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED # Set moe backend for DeepSeek: the sm100 quant/moe resolution # moved to the resolution pipeline (arg_groups/overrides.py: @@ -5884,7 +5917,7 @@ def _handle_model_specific_adjustments(self): # here for the rest of the DSA family (DeepSeek-V3.2 / # GLM-5.x) that shares the same decode top-k path. envs.SGLANG_OPT_USE_TOPK_V2.set(False) - if not self._resolved().enable_dp_attention and self.nnodes == 1: + if not self._resolved().enable_dp_attention and cfg.nnodes == 1: # TODO (Hubert): Put this back later # self.enable_aiter_allreduce_fusion = True logger.info( @@ -5970,7 +6003,7 @@ def _handle_model_specific_adjustments(self): is_mxfp4_quant_format = quant_method == "mxfp4" if ( not self._resolved().enable_dp_attention - and self.nnodes == 1 + and cfg.nnodes == 1 and is_hip() ): # TODO (Hubert): Put this back later @@ -5993,14 +6026,14 @@ def _handle_model_specific_adjustments(self): ), "Triton kernel MoE is only supported when ep_size == 1" elif model_arch in ("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM"): - if model_arch == "MiMoV2ForCausalLM" and not self.encoder_only: + if model_arch == "MiMoV2ForCausalLM" and not cfg.encoder_only: expected_attn_tp_size = get_mimo_v2_fused_qkv_expected_tp_size( hf_config ) view = self._resolved() - attn_dp_size = self.dp_size if view.enable_dp_attention else 1 + attn_dp_size = cfg.dp_size if view.enable_dp_attention else 1 effective_attn_tp_size = ( - self.tp_size // attn_dp_size // view.attn_cp_size + cfg.tp_size // attn_dp_size // view.attn_cp_size ) if ( expected_attn_tp_size is not None @@ -6012,7 +6045,7 @@ def _handle_model_specific_adjustments(self): "qkv_proj weights are " f"TP={expected_attn_tp_size}-interleaved; got " f"{effective_attn_tp_size} " - f"(tp_size={self.tp_size}, dp_size={self.dp_size}, " + f"(tp_size={cfg.tp_size}, dp_size={cfg.dp_size}, " f"enable_dp_attention={view.enable_dp_attention}, " f"attn_cp_size={view.attn_cp_size}). " "Set --tp, --dp, --enable-dp-attention, and " @@ -6038,7 +6071,7 @@ def _handle_model_specific_adjustments(self): pass elif ( model_arch in ("Llama4ForConditionalGeneration", "Llama4ForCausalLM") - and self.device != "cpu" + and cfg.device != "cpu" ): # Attention backend auto-select moved to the override registry # (arg_groups/overrides.py: _llama4_overrides). @@ -6328,6 +6361,7 @@ def _get_default_attn_backend(self, use_mla_backend: bool, model_config): return "triton" def _handle_attention_backend_compatibility(self): + cfg = resolving_view(self) model_config = self.get_model_config() # The attention_backend write clusters of this handler moved to the @@ -6353,24 +6387,24 @@ def _handle_attention_backend_compatibility(self): logger.warning( "Cuda graph is disabled because of using torch native attention backend" ) - self.cuda_graph_config.decode.backend = Backend.DISABLED - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED if attention_backend == "flex_attention": logger.warning( "Cuda graph is disabled because of using torch Flex Attention backend" ) - self.cuda_graph_config.decode.backend = Backend.DISABLED - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED assert ( - self.speculative_algorithm is None + cfg.speculative_algorithm is None ), "Speculative decoding is currently not supported with Flex Attention backend" # Whisper's encoder token padding conflicts with prefix caching. # Only disable for Whisper; other encoder-decoder models (e.g., mllama) use radix cache. if ( model_config.is_encoder_decoder - and not self.disable_radix_cache + and not cfg.disable_radix_cache and "WhisperForConditionalGeneration" in (model_config.hf_config.architectures or []) ): @@ -6413,7 +6447,7 @@ def _handle_attention_backend_compatibility(self): prefill_backend == "trtllm_mha" and is_sm120_supported() and ( - self.kv_cache_dtype == "fp8_e4m3" + cfg.kv_cache_dtype == "fp8_e4m3" or ( envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get() or 0.0 @@ -6434,7 +6468,7 @@ def _handle_attention_backend_compatibility(self): if ( prefill_backend == "trtllm_mha" and not is_sm100_supported() - and (self.enable_prefill_context_parallel or self.attn_cp_size > 1) + and (cfg.enable_prefill_context_parallel or cfg.attn_cp_size > 1) ): raise ValueError( "Prefill context parallelism with the TRTLLM MHA prefill backend " @@ -6451,7 +6485,7 @@ def _handle_attention_backend_compatibility(self): if model_config.context_len > 8192: self._declare( "_handle_attention_backend_compatibility", - mem_fraction_static=self.mem_fraction_static * 0.85, + mem_fraction_static=cfg.mem_fraction_static * 0.85, ) # Other platforms backends @@ -6482,7 +6516,8 @@ def _handle_attention_backend_compatibility(self): def _handle_mxfp8_kv_cache_compatibility(self): """MXFP8 KV cache uses operands available only on SM100+ (Blackwell).""" - if self.kv_cache_dtype != "mxfp8": + cfg = resolving_view(self) + if cfg.kv_cache_dtype != "mxfp8": return if not is_blackwell_supported(): raise ValueError( @@ -6492,8 +6527,9 @@ def _handle_mxfp8_kv_cache_compatibility(self): def _handle_kv4_compatibility(self): """Check FP4 KV cache compatibility with the attention backend""" + cfg = resolving_view(self) - if self.kv_cache_dtype not in ("nvfp4", "fp4_mx_block16"): + if cfg.kv_cache_dtype not in ("nvfp4", "fp4_mx_block16"): return use_mla_backend = self.use_mla_backend() @@ -6501,7 +6537,7 @@ def _handle_kv4_compatibility(self): attention_backend = resolved_view(self).attention_backend if is_cuda(): - if self.kv_cache_dtype == "nvfp4" and not ( + if cfg.kv_cache_dtype == "nvfp4" and not ( is_sm100_supported() or is_sm120_supported() ): raise RuntimeError( @@ -6579,7 +6615,8 @@ def _handle_amd_specifics(self): def _handle_nccl_pre_warm(self): # pre_warm_nccl is only used with CUDA or HIP hardware or NPU hardware - if self.pre_warm_nccl and not (is_cuda() or is_hip() or is_npu()): + cfg = resolving_view(self) + if cfg.pre_warm_nccl and not (is_cuda() or is_hip() or is_npu()): logger.warning( "pre_warm_nccl is only applicable for CUDA or HIP hardware or NPU hardware. " "Ignoring pre_warm_nccl setting on current hardware." @@ -6587,24 +6624,26 @@ def _handle_nccl_pre_warm(self): self._declare("_handle_nccl_pre_warm", pre_warm_nccl=False) def _handle_grammar_backend(self): - if self.grammar_backend is None: + cfg = resolving_view(self) + if cfg.grammar_backend is None: self._declare("_handle_grammar_backend", grammar_backend="xgrammar") def _handle_mamba_backend(self): - if self.mamba_cache_philox_rounds < 0: + cfg = resolving_view(self) + if cfg.mamba_cache_philox_rounds < 0: raise ValueError("--mamba-cache-philox-rounds must be non-negative.") - if self.mamba_max_states_per_path == 0 or self.mamba_max_states_per_path < -1: + if cfg.mamba_max_states_per_path == 0 or cfg.mamba_max_states_per_path < -1: raise ValueError( "--mamba-max-states-per-path must be -1 (unlimited) or a positive " - f"integer, got {self.mamba_max_states_per_path}." + f"integer, got {cfg.mamba_max_states_per_path}." ) - if self.enable_mamba_cache_stochastic_rounding: - if self.mamba_ssm_dtype != "float16": + if cfg.enable_mamba_cache_stochastic_rounding: + if cfg.mamba_ssm_dtype != "float16": raise ValueError( "Stochastic rounding for the Mamba SSM cache requires " - f"--mamba-ssm-dtype float16, got {self.mamba_ssm_dtype!r}. " + f"--mamba-ssm-dtype float16, got {cfg.mamba_ssm_dtype!r}. " "Run with --mamba-ssm-dtype float16 or disable " "--enable-mamba-cache-stochastic-rounding." ) @@ -6614,7 +6653,7 @@ def _handle_mamba_backend(self): "supported on NVIDIA CUDA platforms. Disable " "--enable-mamba-cache-stochastic-rounding on this platform." ) - if self.mamba_backend == "triton" and not is_sm100_supported(): + if cfg.mamba_backend == "triton" and not is_sm100_supported(): raise ValueError( "Stochastic rounding for the Mamba SSM cache with " "--mamba-backend triton requires SM100 with CUDA >= 12.8 " @@ -6624,12 +6663,12 @@ def _handle_mamba_backend(self): "--enable-mamba-cache-stochastic-rounding." ) - if self.mamba_backend == "flashinfer": + if cfg.mamba_backend == "flashinfer": flashinfer_error = ( "FlashInfer mamba module not available, please check the " "FlashInfer installation." ) - if self.enable_mamba_cache_stochastic_rounding: + if cfg.enable_mamba_cache_stochastic_rounding: flashinfer_error += ( " Stochastic rounding with --mamba-backend flashinfer " "requires FlashInfer Mamba and --mamba-ssm-dtype float16." @@ -6651,22 +6690,24 @@ def _handle_int8_mamba_checkpoint(self): # int8-aware: they would read int8 checkpoint slots as bf16 active slots # (wrong pool / out-of-range). Reject the combination up front rather than # silently corrupting state. - if not self.enable_int8_mamba_checkpoint: + cfg = resolving_view(self) + if not cfg.enable_int8_mamba_checkpoint: return - if self.enable_hierarchical_cache: + if cfg.enable_hierarchical_cache: raise ValueError( "--enable-int8-mamba-checkpoint is not supported together with " "--enable-hierarchical-cache: the host-offload path " "is not int8-aware. Disable one of them." ) - if self.radix_cache_backend is not None: + if cfg.radix_cache_backend is not None: raise ValueError( "--enable-int8-mamba-checkpoint only supports the built-in mamba " - f"radix cache; --radix-cache-backend={self.radix_cache_backend!r} " + f"radix cache; --radix-cache-backend={cfg.radix_cache_backend!r} " "is not int8-aware. Omit --radix-cache-backend." ) def _handle_linear_attn_backend(self): + cfg = resolving_view(self) import torch # SM100+: default to FlashInfer GDN decode (and MTP verify, via pool API) @@ -6674,10 +6715,10 @@ def _handle_linear_attn_backend(self): # mamba-ssm-dtype is bf16 (required by FlashInfer GDN on SM100+). # Fixed in FlashInfer v0.6.7: flashinfer-ai/flashinfer#2810 if ( - self.linear_attn_decode_backend is None - and self.linear_attn_backend != "helion" + cfg.linear_attn_decode_backend is None + and cfg.linear_attn_backend != "helion" and is_sm100_supported() - and self.mamba_ssm_dtype == "bfloat16" + and cfg.mamba_ssm_dtype == "bfloat16" # Stage 4: flashinfer's recurrent_kda compiles the state slot stride # as a free int64, so it reads the page-major/unified envelope-strided # state correctly — the unified-memory skip is no longer needed (the @@ -6693,7 +6734,7 @@ def _handle_linear_attn_backend(self): ) # SM100+ FlashInfer GDN decode requires bf16 state; SM90 uses float32. - decode = self.linear_attn_decode_backend or self.linear_attn_backend + decode = cfg.linear_attn_decode_backend or cfg.linear_attn_backend # FlashKDA is a prefill-only KDA kernel (no decode kernel) but shares the # backend choice list, so guard it from being selected for decode: error @@ -6701,7 +6742,7 @@ def _handle_linear_attn_backend(self): # triton decode when it was only inherited from base=flashkda (prefill # keeps FlashKDA). if decode == "flashkda": - if self.linear_attn_decode_backend == "flashkda": + if cfg.linear_attn_decode_backend == "flashkda": raise ValueError( "--linear-attn-decode-backend flashkda is not supported: " "FlashKDA is prefill-only. Use " @@ -6719,34 +6760,34 @@ def _handle_linear_attn_backend(self): if ( decode == "flashinfer" - and self.mamba_ssm_dtype != "bfloat16" + and cfg.mamba_ssm_dtype != "bfloat16" and is_cuda() and torch.cuda.get_device_capability()[0] >= 10 ): raise ValueError( "--linear-attn-decode-backend flashinfer on SM100+ requires " "--mamba-ssm-dtype bfloat16, " - f"got {self.mamba_ssm_dtype!r}" + f"got {cfg.mamba_ssm_dtype!r}" ) - verify = self.linear_attn_verify_backend + verify = cfg.linear_attn_verify_backend if verify is None and decode == "flashinfer": verify = "flashinfer" if ( verify == "flashinfer" - and self.mamba_ssm_dtype != "bfloat16" + and cfg.mamba_ssm_dtype != "bfloat16" and is_cuda() and torch.cuda.get_device_capability()[0] >= 10 ): raise ValueError( "--linear-attn-verify-backend flashinfer on SM100+ requires " "--mamba-ssm-dtype bfloat16, " - f"got {self.mamba_ssm_dtype!r}" + f"got {cfg.mamba_ssm_dtype!r}" ) # SM100+ FlashInfer GDN prefill requires CUDA 13+ (CuTe DSL kernel) # for correctness and best performance. - prefill = self.linear_attn_prefill_backend or self.linear_attn_backend + prefill = cfg.linear_attn_prefill_backend or cfg.linear_attn_backend cuda_version = torch.version.cuda cuda_major = int(cuda_version.split(".")[0]) if cuda_version is not None else 0 if ( @@ -6774,7 +6815,7 @@ def _handle_linear_attn_backend(self): # does NOT route through MambaPool.copy_from, so the ReplaySSM ring # cursor of the donated/kept slot would not be reset there. Handling # that donation path is a follow-up; for now require no_buffer. - if self.enable_linear_replayssm: + if cfg.enable_linear_replayssm: if decode not in {"triton", "helion"}: raise ValueError( "--enable-linear-replayssm requires Triton, or Helion for " @@ -6790,9 +6831,9 @@ def _handle_linear_attn_backend(self): "--enable-linear-replayssm requires --mamba-radix-cache-strategy " "no_buffer (the default); the extra_buffer ping-pong " "donation path is not yet supported (follow-up). Got " - f"--mamba-radix-cache-strategy={self.mamba_radix_cache_strategy!r}." + f"--mamba-radix-cache-strategy={cfg.mamba_radix_cache_strategy!r}." ) - if self.disaggregation_mode != "null": + if cfg.disaggregation_mode != "null": # The disaggregated decode pool (HybridMambaDecodeReqToTokenPool) # is not wired for the ReplaySSM ring, so the flag would silently # no-op there; disagg also runs a different cache/coordination @@ -6800,12 +6841,12 @@ def _handle_linear_attn_backend(self): raise ValueError( "--enable-linear-replayssm is not supported under PD " "disaggregation yet (follow-up). Got " - f"--disaggregation-mode={self.disaggregation_mode!r}." + f"--disaggregation-mode={cfg.disaggregation_mode!r}." ) - if self.linear_replayssm_cache_len < 1: + if cfg.linear_replayssm_cache_len < 1: raise ValueError( "--linear-replayssm-cache-len must be >= 1, got " - f"{self.linear_replayssm_cache_len}." + f"{cfg.linear_replayssm_cache_len}." ) # ReplaySSM spec-verify (Part B of #28511): linear-chain target verify via @@ -6818,14 +6859,14 @@ def _handle_linear_attn_backend(self): # GDN sizes the window to the draft maximum; KDA (kda_backend) keeps a # --linear-replayssm-cache-len window and folds via its own fused # verify ring-write + commit_kda_replayssm_after_verify. - if self.enable_linear_replayssm_spec: - if self.speculative_eagle_topk not in (None, 1): + if cfg.enable_linear_replayssm_spec: + if cfg.speculative_eagle_topk not in (None, 1): raise ValueError( "--enable-linear-replayssm-spec requires a linear draft chain " "(--speculative-eagle-topk in {None, 1}); the chunked verify " "kernel uses a strictly-lower causal mask and is invalid for " "EAGLE tree verify. Got " - f"--speculative-eagle-topk={self.speculative_eagle_topk!r}." + f"--speculative-eagle-topk={cfg.speculative_eagle_topk!r}." ) if decode not in ("triton", "flashinfer"): raise ValueError( @@ -6846,8 +6887,8 @@ def _handle_linear_attn_backend(self): # not take the ragged layout and the flashinfer verify kernel # never writes the ring -> a stale ring would be folded; keep # refusing those combinations. - _algo = (self.speculative_algorithm or "").upper() - verify = self.linear_attn_verify_backend + _algo = (cfg.speculative_algorithm or "").upper() + verify = cfg.linear_attn_verify_backend if _algo not in ("DSPARK", "DFLASH") or verify not in ( "triton", "nv_cutedsl", @@ -6858,23 +6899,23 @@ def _handle_linear_attn_backend(self): "KDA fold-every-commit family (DSPARK/DFLASH) and a " "ring-writing verify kernel (--linear-attn-verify-backend " "triton or nv_cutedsl); got " - f"algorithm={self.speculative_algorithm!r}, " + f"algorithm={cfg.speculative_algorithm!r}, " f"verify={verify!r}. Use SGLANG_RAGGED_VERIFY_MODE=static." ) - if self.disaggregation_mode == "prefill": + if cfg.disaggregation_mode == "prefill": raise ValueError( "--enable-linear-replayssm-spec is not supported on a PD " "prefill server: the ring is spec-verify-only scratch and " "the prefill server never runs spec verify." ) - if self.enable_linear_replayssm: + if cfg.enable_linear_replayssm: raise ValueError( "--enable-linear-replayssm-spec and --enable-linear-replayssm are " "mutually exclusive: they share the ring storage but drive it " "with incompatible cursor protocols (per-decode-forward vs " "per-verify-commit advance)." ) - if self.mamba_ssm_dtype is None: + if cfg.mamba_ssm_dtype is None: logger.info( "--enable-linear-replayssm-spec: setting --mamba-ssm-dtype " "float32 (the closed-loop exact fold keeps the SSM checkpoint " @@ -6884,17 +6925,18 @@ def _handle_linear_attn_backend(self): "_handle_linear_attn_backend", mamba_ssm_dtype="float32", ) - elif self.mamba_ssm_dtype != "float32": + elif cfg.mamba_ssm_dtype != "float32": logger.warning( "--enable-linear-replayssm-spec with --mamba-ssm-dtype=%s: the " "closed-loop fold re-quantizes the committed state each " "commit/flush (fp32 keeps it bit-exact to the fp32 recurrent " "baseline), so it may drift over long sequences. Validate " "accuracy for your model.", - self.mamba_ssm_dtype, + cfg.mamba_ssm_dtype, ) def _handle_legacy_cp_arguments(self): + cfg = resolving_view(self) legacy_mode_to_strategy = { "in-seq-split": "zigzag", "round-robin-split": "interleave", @@ -6905,36 +6947,36 @@ def _handle_legacy_cp_arguments(self): } if ( - self.enable_prefill_context_parallel - or self.enable_dsa_prefill_context_parallel + cfg.enable_prefill_context_parallel + or cfg.enable_dsa_prefill_context_parallel ): self._declare( "_handle_legacy_cp_arguments", enable_prefill_cp=True, ) - if self.enable_prefill_context_parallel and self.cp_strategy is None: + if cfg.enable_prefill_context_parallel and cfg.cp_strategy is None: self._declare( "_handle_legacy_cp_arguments", - cp_strategy=legacy_mode_to_strategy[self.prefill_cp_mode], + cp_strategy=legacy_mode_to_strategy[cfg.prefill_cp_mode], ) - if self.enable_dsa_prefill_context_parallel and self.cp_strategy is None: + if cfg.enable_dsa_prefill_context_parallel and cfg.cp_strategy is None: self._declare( "_handle_legacy_cp_arguments", - cp_strategy=legacy_mode_to_strategy[self.dsa_prefill_cp_mode], + cp_strategy=legacy_mode_to_strategy[cfg.dsa_prefill_cp_mode], ) if ( - self.enable_prefill_context_parallel - and self.enable_dsa_prefill_context_parallel + cfg.enable_prefill_context_parallel + and cfg.enable_dsa_prefill_context_parallel ): return - if not self.enable_prefill_cp or self.cp_strategy is None: + if not cfg.enable_prefill_cp or cfg.cp_strategy is None: return - mode = strategy_to_legacy_mode[self.cp_strategy] - use_dsa_legacy_aliases = self.enable_dsa_prefill_context_parallel or getattr( + mode = strategy_to_legacy_mode[cfg.cp_strategy] + use_dsa_legacy_aliases = cfg.enable_dsa_prefill_context_parallel or getattr( self._resolved(), "attention_backend", None ) in ("dsa", "dsv4") if use_dsa_legacy_aliases: @@ -6961,7 +7003,8 @@ def _handle_legacy_cp_arguments(self): ) def _handle_context_parallelism(self): - if parse_connector_type(self.model_path) != ConnectorType.INSTANCE: + cfg = resolving_view(self) + if parse_connector_type(cfg.model_path) != ConnectorType.INSTANCE: from sglang.srt.configs.model_config import is_deepseek_dsa from sglang.srt.layers.cp.utils import CP_V2_DEFAULT_MODEL_CLASSES @@ -6972,38 +7015,38 @@ def _handle_context_parallelism(self): is_dsa_default_model = is_deepseek_dsa(hf_config) # DSA CP-v2 currently supports only the interleave strategy. enable_default_cp_v2 = not is_dsa_default_model or ( - self.enable_prefill_cp and self.cp_strategy == "interleave" + cfg.enable_prefill_cp and cfg.cp_strategy == "interleave" ) if enable_default_cp_v2 and not envs.SGLANG_ENABLE_CP_V2.is_set(): envs.SGLANG_ENABLE_CP_V2.set(True) if ( - self.enable_prefill_cp + cfg.enable_prefill_cp and model_arch in ("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM") and envs.SGLANG_ENABLE_CP_V2.get() ): - if self.cp_strategy != "zigzag": + if cfg.cp_strategy != "zigzag": raise ValueError( "MiMo V2 CP-v2 only supports --cp-strategy zigzag." ) if ( model_config.is_multimodal - and not self.language_only - and not self.language_model_only + and not cfg.language_only + and not cfg.language_model_only ): raise ValueError( "MiMo V2 CP-v2 only supports text inference; add " "--language-only." ) - if self.enable_prefill_cp and self.cp_strategy is None: + if cfg.enable_prefill_cp and cfg.cp_strategy is None: raise ValueError( "--cp-strategy must be set when --enable-prefill-cp is enabled." ) if ( - self.enable_prefill_context_parallel - and self.enable_dsa_prefill_context_parallel + cfg.enable_prefill_context_parallel + and cfg.enable_dsa_prefill_context_parallel ): raise ValueError( "--enable-prefill-context-parallel and " @@ -7018,38 +7061,38 @@ def _handle_context_parallelism(self): if view.attn_cp_size > 1: # The tp_size is the world size, not the real tensor parallel size assert ( - self.tp_size % view.attn_cp_size == 0 + cfg.tp_size % view.attn_cp_size == 0 ), "tp_size must be divisible by attn_cp_size" assert ( - self.tp_size % (self.dp_size * view.attn_cp_size) == 0 + cfg.tp_size % (cfg.dp_size * view.attn_cp_size) == 0 ), "tp_size must be divisible by dp_size * attn_cp_size" assert ( - not self.enable_aiter_allreduce_fusion + not cfg.enable_aiter_allreduce_fusion ), "Aiter allreduce fusion is not supported with context parallelism" - if self.moe_dp_size > 1: + if cfg.moe_dp_size > 1: # The tp_size is the world size, not the real tensor parallel size assert ( - self.tp_size % self.moe_dp_size == 0 + cfg.tp_size % cfg.moe_dp_size == 0 ), "tp_size must be divisible by moe_dp_size" assert ( - view.ep_size * self.moe_dp_size <= self.tp_size + view.ep_size * cfg.moe_dp_size <= cfg.tp_size ), "ep_size * moe_dp_size must be less than or equal to tp_size" - assert self.pp_size == 1, "PP is not supported with context parallelism" + assert cfg.pp_size == 1, "PP is not supported with context parallelism" if view.ep_size > 1: assert ( - view.ep_size * self.moe_dp_size == self.tp_size + view.ep_size * cfg.moe_dp_size == cfg.tp_size ), "ep_size * moe_dp_size must be equal to tp_size" assert ( - not self.enable_aiter_allreduce_fusion + not cfg.enable_aiter_allreduce_fusion ), "Aiter allreduce fusion is not supported with context parallelism" - if view.attn_cp_size != self.moe_dp_size: + if view.attn_cp_size != cfg.moe_dp_size: assert ( - self.moe_dp_size == 1 + cfg.moe_dp_size == 1 ), "attn_cp_size != moe_dp_size is only supported when moe_dp_size == 1" from sglang.srt.layers.cp.base import init_cp_strategy @@ -7057,31 +7100,32 @@ def _handle_context_parallelism(self): init_cp_strategy(self) def _handle_dwdp(self): - if self.dwdp_size <= 1: + cfg = resolving_view(self) + if cfg.dwdp_size <= 1: return assert ( - self.dwdp_size >= 2 - ), f"dwdp_size must be >= 2 when enabled, got {self.dwdp_size}" + cfg.dwdp_size >= 2 + ), f"dwdp_size must be >= 2 when enabled, got {cfg.dwdp_size}" assert ( - self.dwdp_size == self.tp_size - ), f"dwdp_size ({self.dwdp_size}) must equal tp_size ({self.tp_size})" - assert self.disaggregation_mode in ( + cfg.dwdp_size == cfg.tp_size + ), f"dwdp_size ({cfg.dwdp_size}) must equal tp_size ({cfg.tp_size})" + assert cfg.disaggregation_mode in ( "null", "prefill", ), "DWDP requires --disaggregation-mode null or prefill" assert ( - not self.enable_eplb + not cfg.enable_eplb ), "EPLB dynamic migration conflicts with static DWDP partitioning" assert ( - self.speculative_algorithm is None + cfg.speculative_algorithm is None ), "DWDP does not support speculative decoding (MTP/draft workers)" - assert self.pp_size == 1, "DWDP requires pp_size == 1" + assert cfg.pp_size == 1, "DWDP requires pp_size == 1" assert ( - not self.enable_two_batch_overlap + not cfg.enable_two_batch_overlap ), "DWDP's prefetch event protocol does not support two-batch overlap" - if self.disaggregation_mode == "null": + if cfg.disaggregation_mode == "null": logger.warning( "DWDP with --disaggregation-mode null: decode steps re-fetch all " "remote expert weights every step, which is slow. DWDP is " @@ -7090,7 +7134,7 @@ def _handle_dwdp(self): self._declare( "_handle_dwdp", - dp_size=self.dwdp_size, + dp_size=cfg.dwdp_size, ) self._declare( "_handle_dwdp", @@ -7107,9 +7151,9 @@ def _handle_dwdp(self): ) self._declare( "_handle_dwdp", - ep_size=self.dwdp_size, + ep_size=cfg.dwdp_size, ) - self.moe_ep_size = self.dwdp_size + self.moe_ep_size = cfg.dwdp_size self._declare( "_handle_dwdp", moe_dp_size=1, @@ -7127,8 +7171,8 @@ def _handle_dwdp(self): ) logger.info( - f"DWDP enabled: dwdp_size={self.dwdp_size}, " - f"auto-forced dp_size={self.dp_size}, moe_ep_size={self.moe_ep_size}, " + f"DWDP enabled: dwdp_size={cfg.dwdp_size}, " + f"auto-forced dp_size={cfg.dp_size}, moe_ep_size={self.moe_ep_size}, " f"moe_dense_tp_size=1, moe_a2a_backend=none, " f"dp_attention_local_control_broadcast=True, " f"enable_dp_lm_head=True, SCHEDULER_SKIP_ALL_GATHER=True, " @@ -7138,6 +7182,7 @@ def _handle_dwdp(self): def _handle_data_parallelism(self): # The dp_size==1 resets moved to the resolution pipeline # (arg_groups/overrides.py: _data_parallelism_defaults). + cfg = resolving_view(self) from sglang.srt.arg_groups.overrides import ( _data_parallelism_defaults, run_post_process_pass, @@ -7145,8 +7190,8 @@ def _handle_data_parallelism(self): run_post_process_pass(self, _data_parallelism_defaults) - if self.mm_enable_dp_encoder: - if self.tp_size == 1: + if cfg.mm_enable_dp_encoder: + if cfg.tp_size == 1: logger.warning( "--mm-enable-dp-encoder is enabled with TP=1, so the encoder " "has no data-parallel work to distribute. Disable it unless " @@ -7160,23 +7205,23 @@ def _handle_data_parallelism(self): "prefill is a material part of TTFT. Measure against the default " "for small-image workloads because replication and aggregation " "can increase memory use and overhead.", - self.tp_size, + cfg.tp_size, ) if self._resolved().enable_dp_attention: self._declare( "_handle_data_parallelism", - schedule_conservativeness=self.schedule_conservativeness * 0.3, + schedule_conservativeness=cfg.schedule_conservativeness * 0.3, ) - assert self.tp_size % self.dp_size == 0 - original_chunked_prefill_size = self.chunked_prefill_size + assert cfg.tp_size % cfg.dp_size == 0 + original_chunked_prefill_size = cfg.chunked_prefill_size self._declare( "_handle_data_parallelism", - chunked_prefill_size=self.chunked_prefill_size // self.dp_size, + chunked_prefill_size=cfg.chunked_prefill_size // cfg.dp_size, ) logger.warning( f"DP attention is enabled. chunked prefill size is adjusted " - f"from {original_chunked_prefill_size} to {self.chunked_prefill_size}." + f"from {original_chunked_prefill_size} to {cfg.chunked_prefill_size}." ) # The prefill CUDA graph max_bs was derived from the pre-DP-division @@ -7185,14 +7230,14 @@ def _handle_data_parallelism(self): # the per-DP-rank chunked_prefill_size so breakable CUDA graph # capture never exceeds the MoE all-to-all's max_num_tokens budget, # which is also sized from the DP-adjusted chunked_prefill_size. - prefill_cfg = self.cuda_graph_config.prefill + prefill_cfg = cfg.cuda_graph_config.prefill if ( prefill_cfg.backend != Backend.DISABLED and prefill_cfg.max_bs is not None - and prefill_cfg.max_bs > self.chunked_prefill_size + and prefill_cfg.max_bs > cfg.chunked_prefill_size and (Phase.PREFILL, "max_bs") not in self._cuda_graph_config_locked ): - prefill_cfg.max_bs = self.chunked_prefill_size + prefill_cfg.max_bs = cfg.chunked_prefill_size if (Phase.PREFILL, "bs") not in self._cuda_graph_config_locked: prefill_cfg.bs = self._generate_prefill_cuda_graph_batch_sizes( prefill_cfg.max_bs @@ -7212,6 +7257,7 @@ def _handle_moe_kernel_config(self): # The quantization-driven runner resolutions moved to the pipeline # (arg_groups/overrides.py: _moe_runner_backend_quant_constraints); # the compatibility asserts and fusion writes stay below. + cfg = resolving_view(self) from sglang.srt.arg_groups.overrides import ( _moe_runner_backend_quant_constraints, _moe_runner_fusion_disable, @@ -7230,7 +7276,7 @@ def _handle_moe_kernel_config(self): ], f"Invalid quantization '{view.quantization}'. \nFlashInfer Cutlass MOE supports only: 'modelopt_fp4', 'modelopt_fp8', 'modelopt_mixed', or bfloat16 (None)." assert view.ep_size in [ 1, - self.tp_size, + cfg.tp_size, ], "The expert parallel size must be 1 or the same as the tensor parallel size" if view.moe_runner_backend == "flashinfer_cutedsl": @@ -7241,7 +7287,7 @@ def _handle_moe_kernel_config(self): ), f"Invalid quantization '{view.quantization}'. \nFlashInfer CuteDSL MOE currently supports only: 'modelopt_fp4', 'modelopt_mixed' (with NVFP4 MoE layers), 'nvfp4_online', or hybrid NVFP4 models." assert view.ep_size in [ 1, - self.tp_size, + cfg.tp_size, ], "The expert parallel size must be 1 or the same as the tensor parallel size" assert view.moe_a2a_backend in [ "none", @@ -7305,12 +7351,13 @@ def cutedsl_moe_max_num_tokens(self) -> int: capture, and decode/verify bounds; num_tokens_per_req is speculative_num_draft_tokens under speculative decoding, else 1. """ - if self.speculative_algorithm: - num_tokens_per_req = self.speculative_num_draft_tokens or 1 + cfg = resolving_view(self) + if cfg.speculative_algorithm: + num_tokens_per_req = cfg.speculative_num_draft_tokens or 1 else: num_tokens_per_req = 1 - prefill_tokens = self.max_prefill_tokens - cg_config = self.cuda_graph_config + prefill_tokens = cfg.max_prefill_tokens + cg_config = cfg.cuda_graph_config if cg_config is not None and cg_config.prefill.backend == Backend.TC_PIECEWISE: prefill_tokens = max(prefill_tokens, cg_config.prefill.max_bs or 0) decode_max_bs = (cg_config.decode.max_bs if cg_config is not None else 0) or 0 @@ -7320,29 +7367,29 @@ def cutedsl_moe_max_num_tokens(self) -> int: def max_prefill_buffer_tokens(self) -> int: """Prefill-buffer ceiling: chunked_prefill_size, except PP dynamic chunking can grow chunks toward max_prefill_tokens and probe at 1.25x.""" + cfg = resolving_view(self) chunked = ( - self.chunked_prefill_size - if self.chunked_prefill_size and self.chunked_prefill_size > 0 + cfg.chunked_prefill_size + if cfg.chunked_prefill_size and cfg.chunked_prefill_size > 0 else 0 ) tokens = chunked - if self.enable_dynamic_chunking and self.pp_size > 1 and chunked: - tokens = max( - tokens, self.max_prefill_tokens or 0, math.ceil(chunked * 1.25) - ) + if cfg.enable_dynamic_chunking and cfg.pp_size > 1 and chunked: + tokens = max(tokens, cfg.max_prefill_tokens or 0, math.ceil(chunked * 1.25)) return tokens def _validate_cutedsl_a2a_token_budget(self): """Fail fast if the FlashInfer A2A dispatcher workspace cannot cover the largest CuteDSL MoE forward. Runs after speculative decoding is resolved so cutedsl_moe_max_num_tokens() sees the final num_tokens_per_req.""" + cfg = resolving_view(self) view = resolved_view(self) if not ( view.moe_a2a_backend == "flashinfer" and view.moe_runner_backend == "flashinfer_cutedsl" - and self.max_prefill_tokens > 0 - and self.disaggregation_mode != "decode" + and cfg.max_prefill_tokens > 0 + and cfg.disaggregation_mode != "decode" ): return required_tokens = self.cutedsl_moe_max_num_tokens() @@ -7374,6 +7421,7 @@ def _handle_a2a_moe(self): # the resolution pipeline (arg_groups/overrides.py: # _a2a_backend_overrides / _a2a_ep_size); the per-backend logs, # asserts, fusion/deepep_mode/env/cuda-graph writes stay below. + cfg = resolving_view(self) from sglang.srt.arg_groups.overrides import ( _a2a_backend_overrides, _a2a_ep_size, @@ -7390,13 +7438,13 @@ def _handle_a2a_moe(self): run_post_process_pass(self, _a2a_fusion_adjustments) a2a_backend = resolved_view(self).moe_a2a_backend - if self.enable_waterfill: + if cfg.enable_waterfill: self._declare("_handle_a2a_moe", enforce_shared_experts_fusion=True) logger.info(f"Waterfill is enabled with moe_a2a_backend='{a2a_backend}'.") if a2a_backend == "deepep": - if self.moe_runner_backend == "flashinfer_cutedsl": - if self.deepep_mode == "auto": + if cfg.moe_runner_backend == "flashinfer_cutedsl": + if cfg.deepep_mode == "auto": self._declare( "_handle_a2a_moe", deepep_mode="low_latency", @@ -7407,17 +7455,17 @@ def _handle_a2a_moe(self): "deepep auto mode would crash during prefill. " "low_latency covers both prefill and decode." ) - elif self.deepep_mode == "normal": + elif cfg.deepep_mode == "normal": raise ValueError( "flashinfer_cutedsl FP4 MoE only supports DeepEP " "low_latency dispatch (masked layout). DeepEP normal " "(prefill) dispatch has no CuteDSL FP4 handler. Pass " "--deepep-mode low_latency or auto." ) - if self.deepep_mode == "normal": + if cfg.deepep_mode == "normal": logger.warning("Cuda graph is disabled because deepep_mode=`normal`") - self.cuda_graph_config.decode.backend = Backend.DISABLED - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED # The resolving view, not the field: `_a2a_backend_overrides` may have # moved this already (waterfill forces `deepep`). @@ -7429,11 +7477,11 @@ def _handle_a2a_moe(self): moe_a2a_backend="none", ) - if self.moe_a2a_backend == "flashinfer": + if cfg.moe_a2a_backend == "flashinfer": assert ( - resolved_view(self).enable_dp_attention and self.dp_size == self.tp_size + resolved_view(self).enable_dp_attention and cfg.dp_size == cfg.tp_size ), "Flashinfer MoE A2A is only supported with dp_size == tp_size and --enable-dp-attention" - if self.deepep_mode != "auto": + if cfg.deepep_mode != "auto": logger.warning("--deepep-mode is ignored for Flashinfer MoE A2A") if not envs.SGLANG_MOE_NVFP4_DISPATCH.is_set() and ( resolved_view(self).quantization == "modelopt_fp4" @@ -7450,7 +7498,7 @@ def _handle_a2a_moe(self): ], "Flashinfer MoE A2A is only supported with flashinfer_cutlass, flashinfer_cutedsl or flashinfer_trtllm_routed moe runner backend" if a2a_backend == "mori": - if self.deepep_mode == "auto": + if cfg.deepep_mode == "auto": self._declare( "_handle_a2a_moe", deepep_mode="normal", @@ -7460,7 +7508,7 @@ def _handle_a2a_moe(self): # Check chunked prefill for mori # Skip validation if chunked prefill is disabled (i.e., size <= 0). # Skip validation if disaggregation mode is decode. - if self.chunked_prefill_size > 0 and self.disaggregation_mode != "decode": + if cfg.chunked_prefill_size > 0 and cfg.disaggregation_mode != "decode": assert ( self._required_mori_dispatch_tokens_per_rank() ) <= envs.SGLANG_MORI_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get(), ( @@ -7470,12 +7518,12 @@ def _handle_a2a_moe(self): ) if a2a_backend == "pplx": - if self.deepep_mode == "normal": + if cfg.deepep_mode == "normal": raise ValueError( "moe_a2a_backend='pplx' only supports low-latency mode; " "set --deepep-mode to 'low_latency' or 'auto'." ) - if self.deepep_mode == "auto": + if cfg.deepep_mode == "auto": self._declare( "_handle_a2a_moe", deepep_mode="low_latency", @@ -7484,7 +7532,7 @@ def _handle_a2a_moe(self): # pplx-kernels' AllToAll needs numDPGroups (== attention dp_size) > 1; # without DP attention numDPGroups == 1 and construction fails deep in # the kernel. This also implies ep_size >= 2. - assert resolved_view(self).enable_dp_attention and self.dp_size >= 2, ( + assert resolved_view(self).enable_dp_attention and cfg.dp_size >= 2, ( "moe_a2a_backend='pplx' requires --enable-dp-attention with at " "least 2 DP groups (--dp-size >= 2)." ) @@ -7496,7 +7544,7 @@ def _handle_a2a_moe(self): "moe_a2a_backend='pplx' is only supported with --moe-runner-backend " "deep_gemm (or auto)." ) - if self.moe_runner_backend == "auto": + if cfg.moe_runner_backend == "auto": self._declare( "_handle_a2a_moe", moe_runner_backend="deep_gemm", @@ -7506,7 +7554,7 @@ def _handle_a2a_moe(self): # Check per-rank dispatch tokens for pplx # Skip validation if chunked prefill is disabled (i.e., size <= 0) # Skip validation if disaggregation mode is decode - if self.chunked_prefill_size > 0 and self.disaggregation_mode != "decode": + if cfg.chunked_prefill_size > 0 and cfg.disaggregation_mode != "decode": assert ( self._required_pplx_dispatch_tokens_per_rank() ) <= envs.SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get(), ( @@ -7517,17 +7565,20 @@ def _handle_a2a_moe(self): def _required_mori_dispatch_tokens_per_rank(self) -> int: """Max tokens a single rank dispatches through MoRI in one forward.""" - return self.chunked_prefill_size + cfg = resolving_view(self) + return cfg.chunked_prefill_size def _required_pplx_dispatch_tokens_per_rank(self) -> int: """Max tokens a single rank dispatches through pplx in one forward.""" - required = self.chunked_prefill_size - if self.cuda_graph_max_bs_decode is not None: - required = max(required, self.cuda_graph_max_bs_decode) + cfg = resolving_view(self) + required = cfg.chunked_prefill_size + if cfg.cuda_graph_max_bs_decode is not None: + required = max(required, cfg.cuda_graph_max_bs_decode) return required def _handle_eplb_and_dispatch(self): - if self.enable_eplb and (self.expert_distribution_recorder_mode is None): + cfg = resolving_view(self) + if cfg.enable_eplb and (cfg.expert_distribution_recorder_mode is None): self._declare( "_handle_eplb_and_dispatch", expert_distribution_recorder_mode="stat", @@ -7540,8 +7591,8 @@ def _handle_eplb_and_dispatch(self): # sum their partial outputs, so the pick has to agree across ranks. needs_rank_invariant_dispatch = self._resolved().moe_a2a_backend == "none" - if (self.enable_eplb or (self.init_expert_location != "trivial")) and ( - self.ep_dispatch_algorithm is None + if (cfg.enable_eplb or (cfg.init_expert_location != "trivial")) and ( + cfg.ep_dispatch_algorithm is None ): self._declare( "_handle_eplb_and_dispatch", @@ -7552,23 +7603,24 @@ def _handle_eplb_and_dispatch(self): # `dynamic` / `fake` switch to the row-index pick; `static` reads a # per-rank table and `lp` samples inside its kernel. - if needs_rank_invariant_dispatch and self.ep_dispatch_algorithm in ( + if needs_rank_invariant_dispatch and cfg.ep_dispatch_algorithm in ( "static", "lp", ): raise ValueError( - f"--ep-dispatch-algorithm {self.ep_dispatch_algorithm} picks a " + f"--ep-dispatch-algorithm {cfg.ep_dispatch_algorithm} picks a " "different physical replica per rank, which only holds up when an " "a2a backend routes each token to a single rank. Use " "--ep-dispatch-algorithm dynamic with --moe-a2a-backend none." ) - if self.enable_eplb and self.ep_join_mode != "scale": + if cfg.enable_eplb and cfg.ep_join_mode != "scale": assert self._resolved().ep_size > 1 def _handle_elastic_ep(self): - if self.elastic_ep_rejoin: - if self.ep_join_mode is None: + cfg = resolving_view(self) + if cfg.elastic_ep_rejoin: + if cfg.ep_join_mode is None: logger.warning( "--elastic-ep-rejoin is deprecated, use --elastic-ep-join-mode recover instead." ) @@ -7577,65 +7629,65 @@ def _handle_elastic_ep(self): ep_join_mode="recover", ) else: - assert self.ep_join_mode == "recover", ( + assert cfg.ep_join_mode == "recover", ( "--elastic-ep-rejoin (deprecated) conflicts with " - f"--elastic-ep-join-mode {self.ep_join_mode}." + f"--elastic-ep-join-mode {cfg.ep_join_mode}." ) - if self.elastic_ep_backend is not None: - if self.enable_eplb: - if self.eplb_algorithm == "auto": + if cfg.elastic_ep_backend is not None: + if cfg.enable_eplb: + if cfg.eplb_algorithm == "auto": self._declare( "_handle_elastic_ep", eplb_algorithm="elasticity_aware", ) - assert self.eplb_algorithm in [ + assert cfg.eplb_algorithm in [ "elasticity_aware", "elasticity_aware_hierarchical", ], "Elastic EP requires eplb_algorithm to be set to 'auto' or 'elasticity_aware(_hierarchical)'." - assert self.pp_size == 1, "PP size should be set to 1 under elastic EP" + assert cfg.pp_size == 1, "PP size should be set to 1 under elastic EP" - if self.elastic_ep_backend == "mooncake": + if cfg.elastic_ep_backend == "mooncake": self._declare( "_handle_elastic_ep", mooncake_ib_device=self._validate_ib_devices( - self.mooncake_ib_device + cfg.mooncake_ib_device ), ) - if self.ep_join_mode is not None: + if cfg.ep_join_mode is not None: assert ( - self.elastic_ep_backend is not None + cfg.elastic_ep_backend is not None ), "--elastic-ep-join-mode requires --elastic-ep-backend to be set." - if self.ep_join_mode == "scale": - assert self.node_rank == 1, ( + if cfg.ep_join_mode == "scale": + assert cfg.node_rank == 1, ( "Elastic EP scale-up requires one joining TP group at " - f"--node-rank 1 (got {self.node_rank})." + f"--node-rank 1 (got {cfg.node_rank})." ) - assert self.ep_join_rank_offset > 0, ( + assert cfg.ep_join_rank_offset > 0, ( "Elastic EP scale joiners require " "--elastic-ep-join-rank-offset set to the current " "effective EP size." ) - if self.ep_join_rank_offset != 0: - assert self.ep_join_mode == "scale", ( + if cfg.ep_join_rank_offset != 0: + assert cfg.ep_join_mode == "scale", ( "--elastic-ep-join-rank-offset is only valid with " "--elastic-ep-join-mode scale." ) assert ( - self.ep_join_rank_offset >= 0 + cfg.ep_join_rank_offset >= 0 ), "elastic EP join rank offset must be >= 0." - if self.max_ep_size is not None: + if cfg.max_ep_size is not None: assert ( - self.elastic_ep_backend is not None + cfg.elastic_ep_backend is not None ), "--max-ep-size requires --elastic-ep-backend to be set." - assert self.max_ep_size > 0, "--max-ep-size must be a positive integer." + assert cfg.max_ep_size > 0, "--max-ep-size must be a positive integer." scaling_active = ( - self.elastic_ep_backend is not None - and self.max_ep_size is not None - and self.max_ep_size > self.tp_size + cfg.elastic_ep_backend is not None + and cfg.max_ep_size is not None + and cfg.max_ep_size > cfg.tp_size ) - if self.elastic_ep_initial_size is not None: + if cfg.elastic_ep_initial_size is not None: assert scaling_active, ( "--elastic-ep-initial-size is only valid for an Elastic EP " "deployment with --max-ep-size larger than its local TP size." @@ -7643,16 +7695,16 @@ def _handle_elastic_ep(self): if scaling_active: resolved = self._resolved() assert ( - self.elastic_ep_scale_timeout > 0 + cfg.elastic_ep_scale_timeout > 0 ), "--elastic-ep-scale-timeout must be greater than zero." - assert self.tokenizer_worker_num == 1, ( + assert cfg.tokenizer_worker_num == 1, ( "Elastic EP runtime scale-up currently requires " "--tokenizer-worker-num 1." ) assert ( - not self.use_ray + not cfg.use_ray ), "Elastic EP runtime scale-up does not support --use-ray." - assert not self.enable_elastic_expert_backup, ( + assert not cfg.enable_elastic_expert_backup, ( "Elastic EP runtime scale-up does not support " "--enable-elastic-expert-backup." ) @@ -7660,57 +7712,57 @@ def _handle_elastic_ep(self): "_handle_elastic_ep", enable_dp_attention_local_control_broadcast=True, ) - if self.ep_join_mode == "scale": - assert self.elastic_ep_initial_size is not None, ( + if cfg.ep_join_mode == "scale": + assert cfg.elastic_ep_initial_size is not None, ( "Elastic EP scale joiners require --elastic-ep-initial-size " "set to the primary deployment's launch-time EP size." ) - assert self.elastic_ep_initial_size <= self.ep_join_rank_offset, ( + assert cfg.elastic_ep_initial_size <= cfg.ep_join_rank_offset, ( "--elastic-ep-initial-size cannot exceed the current EP size " - f"(initial={self.elastic_ep_initial_size}, " - f"current={self.ep_join_rank_offset})." + f"(initial={cfg.elastic_ep_initial_size}, " + f"current={cfg.ep_join_rank_offset})." ) - join_target = self.ep_join_rank_offset + self.tp_size - assert join_target <= self.max_ep_size, ( + join_target = cfg.ep_join_rank_offset + cfg.tp_size + assert join_target <= cfg.max_ep_size, ( "Elastic EP joining group exceeds --max-ep-size " - f"(join_target={join_target}, max_ep_size={self.max_ep_size})." + f"(join_target={join_target}, max_ep_size={cfg.max_ep_size})." ) - if self.tp_size == 1: - assert self.moe_dense_tp_size == 1, ( + if cfg.tp_size == 1: + assert cfg.moe_dense_tp_size == 1, ( "A single-rank Elastic EP joining group requires " "--moe-dense-tp-size 1." ) else: - if self.elastic_ep_initial_size is None: + if cfg.elastic_ep_initial_size is None: self._declare( "_handle_elastic_ep", - elastic_ep_initial_size=self.tp_size, + elastic_ep_initial_size=cfg.tp_size, ) - assert self.elastic_ep_initial_size == self.tp_size, ( + assert cfg.elastic_ep_initial_size == cfg.tp_size, ( "The primary --elastic-ep-initial-size must equal its " - f"launch-time TP size ({self.tp_size})." + f"launch-time TP size ({cfg.tp_size})." ) - assert self.elastic_ep_initial_size > 0 - assert self.load_balance_method == "round_robin", ( + assert cfg.elastic_ep_initial_size > 0 + assert cfg.load_balance_method == "round_robin", ( "Elastic EP scale-up requires --load-balance-method round_robin; " "load-aware methods " "require global-rank load snapshots after scale " - f"(got {self.load_balance_method})." + f"(got {cfg.load_balance_method})." ) - assert self.elastic_ep_backend == "mooncake", ( + assert cfg.elastic_ep_backend == "mooncake", ( "Elastic EP runtime scale-up requires --elastic-ep-backend " - f"mooncake (got elastic_ep_backend={self.elastic_ep_backend})." + f"mooncake (got elastic_ep_backend={cfg.elastic_ep_backend})." ) - assert self.pp_size == 1, ( + assert cfg.pp_size == 1, ( "Elastic EP scale-up requires --pp-size 1 " - f"(got pp_size={self.pp_size}); WORLD must not span PP stages." + f"(got pp_size={cfg.pp_size}); WORLD must not span PP stages." ) decode_cuda_graph_disabled = ( - self.cuda_graph_config.decode.backend == Backend.DISABLED + cfg.cuda_graph_config.decode.backend == Backend.DISABLED ) prefill_cuda_graph_disabled = ( - self.cuda_graph_config.prefill.backend == Backend.DISABLED + cfg.cuda_graph_config.prefill.backend == Backend.DISABLED ) assert decode_cuda_graph_disabled and prefill_cuda_graph_disabled, ( "Elastic EP runtime scale-up requires decode and prefill CUDA " @@ -7729,18 +7781,18 @@ def _handle_elastic_ep(self): "Elastic EP scale-up requires --attn-cp-size 1 " f"(got attn_cp_size={resolved.attn_cp_size})." ) - assert self.moe_dp_size == 1, ( + assert cfg.moe_dp_size == 1, ( "Elastic EP scale-up requires --moe-dp-size 1 " - f"(got moe_dp_size={self.moe_dp_size})." + f"(got moe_dp_size={cfg.moe_dp_size})." ) - assert resolved.ep_size == self.tp_size, ( + assert resolved.ep_size == cfg.tp_size, ( "Elastic EP scale-up requires ep_size == tp_size " - f"(got ep_size={resolved.ep_size}, tp_size={self.tp_size}); EP, TP " + f"(got ep_size={resolved.ep_size}, tp_size={cfg.tp_size}); EP, TP " "and the attention DP group must all coincide with WORLD." ) - assert self.dp_size == self.tp_size, ( + assert cfg.dp_size == cfg.tp_size, ( "Elastic EP scale-up requires dp_size == tp_size " - f"(got dp_size={self.dp_size}, tp_size={self.tp_size})." + f"(got dp_size={cfg.dp_size}, tp_size={cfg.tp_size})." ) assert resolved.moe_a2a_backend == "nixl", ( "Elastic EP scale-up requires --moe-a2a-backend nixl " @@ -7761,6 +7813,7 @@ def _validate_experimental_sgl_marlin(self): # ===== END TO BE REFACTORED ==== def _handle_expert_distribution_metrics(self): + cfg = resolving_view(self) if "SGLANG_ENABLE_EPLB_BALANCEDNESS_METRIC" in os.environ: raise ValueError( "SGLANG_ENABLE_EPLB_BALANCEDNESS_METRIC is no longer supported. Use " @@ -7769,20 +7822,20 @@ def _handle_expert_distribution_metrics(self): ) if self.should_report_expert_balancedness() and ( - self.expert_distribution_recorder_mode is None + cfg.expert_distribution_recorder_mode is None ): self._declare( "_handle_expert_distribution_metrics", expert_distribution_recorder_mode="stat", ) - if self.expert_distribution_recorder_buffer_size is None: - if (x := self.eplb_rebalance_num_iterations) is not None: + if cfg.expert_distribution_recorder_buffer_size is None: + if (x := cfg.eplb_rebalance_num_iterations) is not None: self._declare( "_handle_expert_distribution_metrics", expert_distribution_recorder_buffer_size=x, ) - elif self.expert_distribution_recorder_mode is not None: + elif cfg.expert_distribution_recorder_mode is not None: self._declare( "_handle_expert_distribution_metrics", expert_distribution_recorder_buffer_size=1000, @@ -7804,26 +7857,27 @@ def _validate_prefill_only_disable_kv_cache_args(self): Backend resolution is checked separately by _handle_prefill_only_disable_kv_cache after backends settle. """ - if not self.prefill_only_disable_kv_cache: + cfg = resolving_view(self) + if not cfg.prefill_only_disable_kv_cache: return # This flag is intentionally scoped to embedding mode for now. Other # prefill-only paths (for example scoring and MIS) can benefit from # the same idea later, but some of them still stage K/V through the # paged cache today. - if not self.is_embedding: + if not cfg.is_embedding: raise ValueError( "--prefill-only-disable-kv-cache currently requires --is-embedding. " "Other prefill-only workloads may be supported in a future change once " "their attention paths stop reading or writing the paged KV cache." ) - if self.kv_cache_dtype in ("nvfp4", "fp4_mx_block16"): + if cfg.kv_cache_dtype in ("nvfp4", "fp4_mx_block16"): raise ValueError( "--prefill-only-disable-kv-cache does not currently support " "--kv-cache-dtype=nvfp4 or --kv-cache-dtype=fp4_mx_block16 because " "the FP4 pool uses a separate allocation path." ) - if self.kv_cache_dtype == "mxfp8": + if cfg.kv_cache_dtype == "mxfp8": raise ValueError( "--prefill-only-disable-kv-cache does not currently support " "--kv-cache-dtype=mxfp8 because the MXFP8 pool stores separate " @@ -7836,13 +7890,13 @@ def _validate_prefill_only_disable_kv_cache_args(self): # so K/V never has to be reused across prefill chunks. # - disable_radix_cache stops the prefix cache from indexing pool # slots that no longer hold real data. - if self.chunked_prefill_size != -1: + if cfg.chunked_prefill_size != -1: raise ValueError( "--prefill-only-disable-kv-cache requires --chunked-prefill-size=-1 so the FA " "backend takes the fa_skip_kv_cache path; otherwise the pool would be touched " "between prefill chunks." ) - if not self.disable_radix_cache: + if not cfg.disable_radix_cache: raise ValueError( "--prefill-only-disable-kv-cache requires --disable-radix-cache because the " "radix cache indexes KV pool slots that no longer hold real data." @@ -7857,7 +7911,7 @@ def _validate_prefill_only_disable_kv_cache_args(self): "the context-parallel attention path writes K/V to the pool via set_kv_buffer, " "which the no-op pool intentionally rejects." ) - if self.enable_prefill_cp: + if cfg.enable_prefill_cp: raise ValueError( "--prefill-only-disable-kv-cache is incompatible with " "--enable-prefill-cp: the prefill-CP path stages K/V through " @@ -7866,7 +7920,7 @@ def _validate_prefill_only_disable_kv_cache_args(self): # HiSparse selects a different pool class (HiSparseDSATokenToKVPool / # HiSparseTokenToKVPoolAllocator) that is not the no-op pool. - if self.enable_hisparse: + if cfg.enable_hisparse: raise ValueError( "--prefill-only-disable-kv-cache is incompatible with --enable-hisparse: " "HiSparse uses a dedicated pool family that is not the no-op MHA pool." @@ -7882,8 +7936,9 @@ def _handle_prefill_only_disable_kv_cache(self): still None, backends haven't settled yet and the resolved (prefill, decode) pair would be a stale (None, None). """ + cfg = resolving_view(self) - if not self.prefill_only_disable_kv_cache: + if not cfg.prefill_only_disable_kv_cache: return assert resolved_view(self).attention_backend is not None, ( @@ -7910,11 +7965,12 @@ def _handle_hicache_ratio_default(self): A decode server keeps the ratio unset here: kv_cache_builder resolves it against the retraction-backup backend (1.0 for host_pool, else 2.0). """ - if self.hicache_ratio is None and self.disaggregation_mode != "decode": + cfg = resolving_view(self) + if cfg.hicache_ratio is None and cfg.disaggregation_mode != "decode": self._declare( "_handle_hicache_ratio_default", hicache_ratio=( - 1.2 if self.hicache_host_memory_mode == "buffer_only" else 2.0 + 1.2 if cfg.hicache_host_memory_mode == "buffer_only" else 2.0 ), ) @@ -7925,13 +7981,14 @@ def _handle_hicache(self): 1) Layout <-> I/O compatibility for direct conflicts. 2) Storage <-> layout compatibility (may rewrite layout). """ + cfg = resolving_view(self) # Skip all normalization when neither hicache nor decode-offload path is active. if not ( - self.enable_hierarchical_cache - or self.disaggregation_decode_enable_offload_kvcache + cfg.enable_hierarchical_cache + or cfg.disaggregation_decode_enable_offload_kvcache or ( - self.disaggregation_mode == "decode" - and self.disaggregation_decode_retraction_backup in (None, "host_pool") + cfg.disaggregation_mode == "decode" + and cfg.disaggregation_decode_retraction_backup in (None, "host_pool") ) ): return @@ -7948,42 +8005,43 @@ def _handle_hicache(self): self._resolve_hicache_dcp_compatibility() def _validate_hicache_host_memory_mode(self): - if self.hicache_host_memory_mode not in ("cache", "buffer_only"): + cfg = resolving_view(self) + if cfg.hicache_host_memory_mode not in ("cache", "buffer_only"): raise ValueError( "hicache_host_memory_mode must be 'cache' or 'buffer_only', " - f"got {self.hicache_host_memory_mode!r}" + f"got {cfg.hicache_host_memory_mode!r}" ) # Both modes are defaulted upstream (a decode server resolves the # ratio later, in kv_cache_builder), so this fires only if that # defaulting regresses -- never build an unsized host pool. if ( - self.hicache_size <= 0 - and self.hicache_ratio is None - and self.disaggregation_mode != "decode" + cfg.hicache_size <= 0 + and cfg.hicache_ratio is None + and cfg.disaggregation_mode != "decode" ): raise ValueError( - f"--hicache-host-memory-mode {self.hicache_host_memory_mode} " + f"--hicache-host-memory-mode {cfg.hicache_host_memory_mode} " "requires a host pool size: pass --hicache-size or " "--hicache-ratio." ) - if self.hicache_host_memory_mode == "cache": + if cfg.hicache_host_memory_mode == "cache": return - if self.hicache_storage_backend is None: + if cfg.hicache_storage_backend is None: raise ValueError( "--hicache-host-memory-mode buffer_only requires a storage backend " "(--hicache-storage-backend): host memory is only a staging buffer " "and all cached data lives in storage." ) - if self.hicache_write_policy == "write_back": + if cfg.hicache_write_policy == "write_back": raise ValueError( "--hicache-host-memory-mode buffer_only does not support " "--hicache-write-policy write_back; use write_through or " "write_through_selective." ) - if self.disaggregation_mode == "decode": + if cfg.disaggregation_mode == "decode": raise ValueError( "--hicache-host-memory-mode buffer_only is not supported on " "decode instances: the decode-side prefetch and offload paths " @@ -7993,9 +8051,10 @@ def _validate_hicache_host_memory_mode(self): ) def _resolve_hicache_dcp_compatibility(self): - if self.dcp_size <= 1 or not self.enable_hierarchical_cache: + cfg = resolving_view(self) + if cfg.dcp_size <= 1 or not cfg.enable_hierarchical_cache: return - if self.hicache_storage_backend is not None: + if cfg.hicache_storage_backend is not None: raise NotImplementedError( "--hicache-storage-backend (L3) with --dcp-size > 1 is not " "supported yet: under DCP each rank holds a distinct " @@ -8003,18 +8062,18 @@ def _resolve_hicache_dcp_compatibility(self): "backup and the storage keys must become dcp_rank-aware " "first. Run HiCache+DCP with L1/L2 only." ) - if self.speculative_algorithm not in (None, "DSPARK"): + if cfg.speculative_algorithm not in (None, "DSPARK"): raise NotImplementedError( "HiCache with --dcp-size > 1 only supports DSPARK speculative " "decoding; other draft-model host pools have no DCP index " "translation." ) - if self.enable_lmcache: + if cfg.enable_lmcache: raise NotImplementedError( "--enable-lmcache with --dcp-size > 1 is not supported: " "LMCache has no DCP-aware index translation." ) - if self.enable_hisparse: + if cfg.enable_hisparse: raise NotImplementedError( "--enable-hisparse with --dcp-size > 1 is not supported: the " "HiSparse host pool is constructed without DCP translation." @@ -8029,13 +8088,14 @@ def _resolve_hicache_dcp_compatibility(self): "HiCache + DCP enabled (L1/L2 only): host pool uses widened " "logical slot accounting with per-rank physical translation at " "the transfer boundary (dcp_size=%d).", - self.dcp_size, + cfg.dcp_size, ) def _resolve_layout_io_compatibility(self): + cfg = resolving_view(self) if ( - self.hicache_mem_layout == "page_first_direct" - and self.hicache_io_backend == "kernel" + cfg.hicache_mem_layout == "page_first_direct" + and cfg.hicache_io_backend == "kernel" ): self._declare( "_resolve_layout_io_compatibility", @@ -8046,8 +8106,8 @@ def _resolve_layout_io_compatibility(self): ) if ( - self.hicache_mem_layout == "page_first" - and self.hicache_io_backend == "direct" + cfg.hicache_mem_layout == "page_first" + and cfg.hicache_io_backend == "direct" ): self._declare( "_resolve_layout_io_compatibility", @@ -8058,19 +8118,20 @@ def _resolve_layout_io_compatibility(self): ) def _resolve_storage_layout_compatibility(self): + cfg = resolving_view(self) if ( - self.hicache_storage_backend != "mooncake" - or self.hicache_mem_layout != "layer_first" + cfg.hicache_storage_backend != "mooncake" + or cfg.hicache_mem_layout != "layer_first" ): return - if self.hicache_io_backend == "direct": + if cfg.hicache_io_backend == "direct": new_layout = "page_first_direct" - elif self.hicache_io_backend == "kernel": + elif cfg.hicache_io_backend == "kernel": new_layout = "page_first" else: # Keep current behavior for unknown backends (e.g., kernel_ascend). - new_layout = self.hicache_mem_layout + new_layout = cfg.hicache_mem_layout self._declare( "_resolve_storage_layout_compatibility", @@ -8078,17 +8139,18 @@ def _resolve_storage_layout_compatibility(self): ) logger.warning( f"Mooncake storage backend does not support layer_first layout, " - f"switching to {new_layout} layout for {self.hicache_io_backend} io backend" + f"switching to {new_layout} layout for {cfg.hicache_io_backend} io backend" ) def _resolve_hf_gguf_model_path(self): """Turn a Hub reference to a .gguf into a local file path.""" + cfg = resolving_view(self) from sglang.srt.utils.hf_transformers_utils import resolve_hf_gguf_reference - resolved = resolve_hf_gguf_reference(self.model_path, revision=self.revision) + resolved = resolve_hf_gguf_reference(cfg.model_path, revision=cfg.revision) if resolved is not None: - logger.info("Resolved GGUF %s -> %s", self.model_path, resolved) - if self.tokenizer_path == self.model_path: + logger.info("Resolved GGUF %s -> %s", cfg.model_path, resolved) + if cfg.tokenizer_path == cfg.model_path: self._declare( "_resolve_hf_gguf_model_path", tokenizer_path=resolved, @@ -8100,15 +8162,15 @@ def _resolve_hf_gguf_model_path(self): # A speculative draft can be a .gguf too, and it is loaded by path, so it # needs the same Hub-reference resolution as the target. - if self.speculative_draft_model_path: + if cfg.speculative_draft_model_path: resolved_draft = resolve_hf_gguf_reference( - self.speculative_draft_model_path, - revision=self.speculative_draft_model_revision, + cfg.speculative_draft_model_path, + revision=cfg.speculative_draft_model_revision, ) if resolved_draft is not None: logger.info( "Resolved draft GGUF %s -> %s", - self.speculative_draft_model_path, + cfg.speculative_draft_model_path, resolved_draft, ) self._declare( @@ -8125,21 +8187,22 @@ def _handle_load_format(self): # The quantization side of the gguf coupling moved to the pipeline # (arg_groups/overrides.py: _gguf_quantization); load_format itself is # genuine config (runtime user updates write it) and stays imperative. + cfg = resolving_view(self) from sglang.srt.arg_groups.overrides import ( _gguf_quantization, run_post_process_pass, ) run_post_process_pass(self, _gguf_quantization) - if ( - self.load_format == "auto" or self.load_format == "gguf" - ) and check_gguf_file(self.model_path): + if (cfg.load_format == "auto" or cfg.load_format == "gguf") and check_gguf_file( + cfg.model_path + ): self._declare( "_handle_load_format", load_format="gguf", ) - if self.load_format == "auto" and self._is_mistral_native_format(): + if cfg.load_format == "auto" and self._is_mistral_native_format(): self._declare( "_handle_load_format", load_format="mistral", @@ -8148,34 +8211,34 @@ def _handle_load_format(self): "Detected Mistral native format checkpoint, setting load_format='mistral'" ) - if is_runai_obj_uri(self.model_path): + if is_runai_obj_uri(cfg.model_path): self._declare( "_handle_load_format", load_format="runai_streamer", ) - elif is_remote_url(self.model_path): + elif is_remote_url(cfg.model_path): self._declare( "_handle_load_format", load_format="remote", ) if ( - self.speculative_draft_model_path is not None - and is_runai_obj_uri(self.speculative_draft_model_path) - and self.speculative_draft_load_format is None + cfg.speculative_draft_model_path is not None + and is_runai_obj_uri(cfg.speculative_draft_model_path) + and cfg.speculative_draft_load_format is None ): self._declare( "_handle_load_format", speculative_draft_load_format="runai_streamer", ) - if self.custom_weight_loader is None: + if cfg.custom_weight_loader is None: self._declare("_handle_load_format", custom_weight_loader=[]) - if self.load_format == "remote_instance": - if self.remote_instance_weight_loader_backend != "modelexpress" and ( - self.remote_instance_weight_loader_seed_instance_ip is None - or self.remote_instance_weight_loader_seed_instance_service_port is None + if cfg.load_format == "remote_instance": + if cfg.remote_instance_weight_loader_backend != "modelexpress" and ( + cfg.remote_instance_weight_loader_seed_instance_ip is None + or cfg.remote_instance_weight_loader_seed_instance_service_port is None ): logger.warning( "Fallback load_format to 'auto' due to incomplete remote instance weight loader settings." @@ -8185,8 +8248,8 @@ def _handle_load_format(self): load_format="auto", ) elif ( - self.remote_instance_weight_loader_send_weights_group_ports is None - and self.remote_instance_weight_loader_backend == "nccl" + cfg.remote_instance_weight_loader_send_weights_group_ports is None + and cfg.remote_instance_weight_loader_backend == "nccl" ): logger.warning( "Fallback load_format to 'auto' due to incomplete remote instance weight loader NCCL group ports settings." @@ -8196,7 +8259,7 @@ def _handle_load_format(self): load_format="auto", ) elif ( - self.remote_instance_weight_loader_backend == "transfer_engine" + cfg.remote_instance_weight_loader_backend == "transfer_engine" and not self.validate_transfer_engine() ): logger.warning( @@ -8208,7 +8271,7 @@ def _handle_load_format(self): ) # Check whether TransferEngine can be used when users want to start seed service that supports TransferEngine backend. - if self.remote_instance_weight_loader_start_seed_via_transfer_engine: + if cfg.remote_instance_weight_loader_start_seed_via_transfer_engine: self._declare( "_handle_load_format", remote_instance_weight_loader_start_seed_via_transfer_engine=self.validate_transfer_engine(), @@ -8220,7 +8283,7 @@ def _handle_load_format(self): # launched, and fallback_load_format inherits a nonsensical format), so # reject it and point at the knob (defense-in-depth; the CLI already # rejects it via LOAD_FORMAT_CHOICES). - if self.load_format == "ipc_cache": + if cfg.load_format == "ipc_cache": raise ValueError( "load_format='ipc_cache' is an internal-only format and must not " "be set directly. Enable the weight cache via --weight-cache-mode " @@ -8231,7 +8294,7 @@ def _handle_load_format(self): # Speculative decoding loads an extra draft model whose weights the # daemon does not export, so refuse the combination up front instead of # failing deep inside draft-worker load (draft-model daemon TBD). - if self.weight_cache_mode != "off" and self.speculative_algorithm is not None: + if cfg.weight_cache_mode != "off" and cfg.speculative_algorithm is not None: raise ValueError( "--weight-cache-mode is not supported together with speculative " "decoding (--speculative-algorithm): the weight cache daemon does " @@ -8258,13 +8321,14 @@ def _is_mistral_native_format(self) -> bool: is present -- those families need Mistral weight loading regardless of which weight files happen to be present. """ + cfg = resolving_view(self) _MISTRAL_NATIVE_PATTERNS = ( "mistral-large-3", "mistral-small-4", "leanstral", ) name_matches = any( - p in str(self.model_path).lower() for p in _MISTRAL_NATIVE_PATTERNS + p in str(cfg.model_path).lower() for p in _MISTRAL_NATIVE_PATTERNS ) def _check_format(has_params, has_consolidated, has_hf_weights) -> bool: @@ -8272,23 +8336,21 @@ def _check_format(has_params, has_consolidated, has_hf_weights) -> bool: return True return has_consolidated and not has_hf_weights - if os.path.isdir(self.model_path): + if os.path.isdir(cfg.model_path): return _check_format( - has_params=os.path.exists(os.path.join(self.model_path, "params.json")), + has_params=os.path.exists(os.path.join(cfg.model_path, "params.json")), has_consolidated=bool( - glob.glob( - os.path.join(self.model_path, "consolidated*.safetensors") - ) + glob.glob(os.path.join(cfg.model_path, "consolidated*.safetensors")) ), has_hf_weights=bool( - glob.glob(os.path.join(self.model_path, "model*.safetensors")) + glob.glob(os.path.join(cfg.model_path, "model*.safetensors")) ), ) try: from huggingface_hub import HfApi - files = {s.rfilename for s in HfApi().model_info(self.model_path).siblings} + files = {s.rfilename for s in HfApi().model_info(cfg.model_path).siblings} return _check_format( has_params="params.json" in files, has_consolidated=any( @@ -8308,23 +8370,24 @@ def _check_format(has_params, has_consolidated, has_hf_weights) -> bool: LANGUAGE_MODEL_ONLY_ARCHITECTURES = ("MuseGlimmerForConditionalGeneration",) def _handle_language_model_only(self): - if not self.language_model_only: + cfg = resolving_view(self) + if not cfg.language_model_only: return for flag, name in ( - (self.encoder_only, "--encoder-only"), - (self.language_only, "--language-only"), - (self.enable_prefix_mm_cache, "--enable-prefix-mm-cache"), + (cfg.encoder_only, "--encoder-only"), + (cfg.language_only, "--language-only"), + (cfg.enable_prefix_mm_cache, "--enable-prefix-mm-cache"), ( - self.enable_broadcast_mm_inputs_process, + cfg.enable_broadcast_mm_inputs_process, "--enable-broadcast-mm-inputs-process", ), - (self.mm_enable_dp_encoder, "--mm-enable-dp-encoder"), + (cfg.mm_enable_dp_encoder, "--mm-enable-dp-encoder"), ): if flag: raise ValueError( f"--language-model-only cannot be combined with {name}" ) - if self.disaggregation_mode != "null": + if cfg.disaggregation_mode != "null": raise ValueError( "--language-model-only is incompatible with --disaggregation-mode " "prefill/decode" @@ -8337,19 +8400,20 @@ def _handle_language_model_only(self): ) def _handle_encoder_disaggregation(self): + cfg = resolving_view(self) self._handle_language_model_only() - if self.enable_prefix_mm_cache and not self.encoder_only: + if cfg.enable_prefix_mm_cache and not cfg.encoder_only: raise ValueError( "--enable-prefix-mm-cache requires --encoder-only to be enabled" ) - if self.encoder_only and self.language_only: + if cfg.encoder_only and cfg.language_only: raise ValueError("Cannot set --encoder-only and --language-only together") - if self.encoder_only and not self.disaggregation_mode == "null": + if cfg.encoder_only and not cfg.disaggregation_mode == "null": raise ValueError( "Cannot set --encoder-only and --disaggregation-mode prefill/decode together" ) - if self.language_only and len(self.encoder_urls) == 0: + if cfg.language_only and len(cfg.encoder_urls) == 0: logger.info( "--language-only is set without --encoder-urls. Encoders are " "expected to register dynamically via the " @@ -8358,34 +8422,34 @@ def _handle_encoder_disaggregation(self): # Validate IB devices when mooncake backend is used if ( - self.disaggregation_transfer_backend == "mooncake" - and self.disaggregation_mode in ("prefill", "decode") - ) or self.encoder_transfer_backend == "mooncake": + cfg.disaggregation_transfer_backend == "mooncake" + and cfg.disaggregation_mode in ("prefill", "decode") + ) or cfg.encoder_transfer_backend == "mooncake": self._declare( "_handle_encoder_disaggregation", disaggregation_ib_device=self._validate_ib_devices( - self.disaggregation_ib_device + cfg.disaggregation_ib_device ), ) # Validate model type for encoder disaggregation hf_config = self.get_model_config().hf_config model_arch = hf_config.architectures[0] - if self.encoder_transfer_backend == "auto": + if cfg.encoder_transfer_backend == "auto": self._declare( "_handle_encoder_disaggregation", encoder_transfer_backend=resolve_encoder_transfer_backend( - self.encoder_transfer_backend, model_arch, self.tp_size + cfg.encoder_transfer_backend, model_arch, cfg.tp_size ), ) - if self.encoder_only or self.language_only: + if cfg.encoder_only or cfg.language_only: logger.info( "Encoder transfer backend auto-resolved to %s for %s at TP%d.", - self.encoder_transfer_backend, + cfg.encoder_transfer_backend, model_arch, - self.tp_size, + cfg.tp_size, ) - if (self.encoder_only or self.language_only) and model_arch not in [ + if (cfg.encoder_only or cfg.language_only) and model_arch not in [ "Qwen2VLForConditionalGeneration", "Qwen3VLForConditionalGeneration", "Qwen2_5_VLForConditionalGeneration", @@ -8484,23 +8548,24 @@ def _normalize_device_group(raw_value: str, context: str) -> str: return json.dumps(normalized_mapping, separators=(",", ":")) def _handle_tokenizer_batching(self): - if self.enable_tokenizer_batch_encode and self.enable_dynamic_batch_tokenizer: + cfg = resolving_view(self) + if cfg.enable_tokenizer_batch_encode and cfg.enable_dynamic_batch_tokenizer: raise ValueError( "Cannot enable both --enable-tokenizer-batch-encode and --enable-dynamic-batch-tokenizer. " "Please choose one tokenizer batching approach." ) - if self.skip_tokenizer_init and not envs.SGLANG_RUST_SERVER.get(): + if cfg.skip_tokenizer_init and not envs.SGLANG_RUST_SERVER.get(): # Tokenizer workers still serve HTTP / state / output work, so # their fanout is preserved; detokenizer workers only decode. - if self.detokenizer_worker_num != 1: + if cfg.detokenizer_worker_num != 1: logger.warning( "skip_tokenizer_init=True leaves no decode work for detokenizer workers; " - f"forcing detokenizer_worker_num=1 (requested {self.detokenizer_worker_num})." + f"forcing detokenizer_worker_num=1 (requested {cfg.detokenizer_worker_num})." ) self._declare("_handle_tokenizer_batching", detokenizer_worker_num=1) - if self.enable_tokenizer_batch_encode: + if cfg.enable_tokenizer_batch_encode: logger.warning( "skip_tokenizer_init=True ignores --enable-tokenizer-batch-encode; disabling it." ) @@ -8509,7 +8574,7 @@ def _handle_tokenizer_batching(self): enable_tokenizer_batch_encode=False, ) - if self.enable_dynamic_batch_tokenizer: + if cfg.enable_dynamic_batch_tokenizer: logger.warning( "skip_tokenizer_init=True ignores --enable-dynamic-batch-tokenizer; disabling it." ) @@ -8531,11 +8596,12 @@ def _handle_multimodal_feature_transport(self): may still auto-select CUDA VMM. The legacy CUDA IPC flag and environment variable remain supported so existing deployments map to this policy. """ - requested_transport = self.mm_feature_transport + cfg = resolving_view(self) + requested_transport = cfg.mm_feature_transport legacy_ipc_is_set = envs.SGLANG_USE_CUDA_IPC_TRANSPORT.is_set() legacy_ipc_enabled = envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get() - if self.keep_mm_feature_on_device: + if cfg.keep_mm_feature_on_device: if requested_transport not in (None, "cuda_ipc"): raise ValueError( "--keep-mm-feature-on-device conflicts with " @@ -8556,7 +8622,7 @@ def _handle_multimodal_feature_transport(self): "--mm-feature-transport=%s instead.", requested_transport, ) - elif self.encoder_only: + elif cfg.encoder_only: requested_transport = "cpu" logger.info( "Multimodal feature transport auto-resolved to cpu for " @@ -8566,14 +8632,14 @@ def _handle_multimodal_feature_transport(self): elif ( self.get_model_config().is_multimodal and is_cuda() - and self.disaggregation_mode == "null" + and cfg.disaggregation_mode == "null" ): # A full GPU pool always degrades to CPU transport per tensor. # Keep CUDA IPC opt-in because even an idle pool consumes HBM # that would otherwise back the KV cache. Multi-node # auto-selection is limited to GB200/GB300 systems where the # runtime already enables the MNNVL/IMEX communication stack. - if self.nnodes == 1: + if cfg.nnodes == 1: requested_transport = "cpu" elif is_mnnvl_fabric_device() and os.path.exists( "/dev/nvidia-caps-imex-channels/channel0" @@ -8616,7 +8682,7 @@ def _handle_multimodal_feature_transport(self): int(legacy_ipc_enabled), ) - if self.encoder_only and requested_transport in ("cuda_ipc", "cuda_vmm"): + if cfg.encoder_only and requested_transport in ("cuda_ipc", "cuda_vmm"): logger.warning( "--mm-feature-transport=%s does not control encoder-only " "output transfer; using cpu for this inactive transport. Select " @@ -8630,7 +8696,7 @@ def _handle_multimodal_feature_transport(self): raise ValueError( "--mm-feature-transport=cuda_vmm requires NVIDIA CUDA." ) - if self.pp_size != 1: + if cfg.pp_size != 1: raise ValueError( "--mm-feature-transport=cuda_vmm does not support pipeline " "parallelism." @@ -8641,7 +8707,7 @@ def _handle_multimodal_feature_transport(self): "SGLANG_RUST_SERVER." ) pool_budget_mb = envs.SGLANG_MM_FEATURE_CACHE_MB.get() - handle_kind = "CUDA FABRIC" if self.nnodes > 1 else "POSIX FD" + handle_kind = "CUDA FABRIC" if cfg.nnodes > 1 else "POSIX FD" logger.info( "Using CUDA VMM for multimodal features with %s sharing: " "reserving up to %d MiB on base GPU %d across %d tokenizer " @@ -8649,8 +8715,8 @@ def _handle_multimodal_feature_transport(self): "back to inline CPU transport.", handle_kind, pool_budget_mb, - self.base_gpu_id, - self.tokenizer_worker_num, + cfg.base_gpu_id, + cfg.tokenizer_worker_num, ) if requested_transport == "cuda_ipc": @@ -8658,7 +8724,7 @@ def _handle_multimodal_feature_transport(self): raise ValueError( "--mm-feature-transport=cuda_ipc requires NVIDIA CUDA." ) - if self.nnodes != 1: + if cfg.nnodes != 1: raise ValueError( "--mm-feature-transport=cuda_ipc only supports a single node." ) @@ -8669,8 +8735,8 @@ def _handle_multimodal_feature_transport(self): "on base GPU %d across %d tokenizer worker(s). This reduces KV " "cache headroom; a full pool falls back to CPU transport.", pool_budget_mb, - self.base_gpu_id, - self.tokenizer_worker_num, + cfg.base_gpu_id, + cfg.tokenizer_worker_num, ) logger.info( "CUDA IPC pool-handle caching is %s. It reuses mappings to the " @@ -8698,19 +8764,20 @@ def _handle_multimodal_feature_transport(self): ) def _handle_environment_variables(self): + cfg = resolving_view(self) self._handle_multimodal_feature_transport() - envs.SGLANG_ENABLE_TORCH_COMPILE.set("1" if self.enable_torch_compile else "0") - if self.mamba_ssm_dtype is not None: - envs.SGLANG_MAMBA_SSM_DTYPE.set(self.mamba_ssm_dtype) + envs.SGLANG_ENABLE_TORCH_COMPILE.set("1" if cfg.enable_torch_compile else "0") + if cfg.mamba_ssm_dtype is not None: + envs.SGLANG_MAMBA_SSM_DTYPE.set(cfg.mamba_ssm_dtype) envs.SGLANG_DISABLE_OUTLINES_DISK_CACHE.set( - "1" if self.disable_outlines_disk_cache else "0" + "1" if cfg.disable_outlines_disk_cache else "0" ) envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.set( - "1" if self.enable_deterministic_inference else "0" + "1" if cfg.enable_deterministic_inference else "0" ) - if self.enable_deterministic_inference: + if cfg.enable_deterministic_inference: envs.SGLANG_FLASHINFER_MOE_FUSED_FINALIZE.set("0") - if self.debug_cuda_graph: + if cfg.debug_cuda_graph: if not (is_cuda() or is_hip()): logger.warning( "--debug-cuda-graph is not supported on non CUDA/HIP devices. " @@ -8723,7 +8790,7 @@ def _handle_environment_variables(self): "Debug mode for CUDA graph is enabled via breakable CUDA graph. " "All operations will run eagerly through the graph capture/replay path." ) - if self.enable_deepseek_v4_fp4_indexer and not ( + if cfg.enable_deepseek_v4_fp4_indexer and not ( is_sm100_supported() or is_sm120_supported() ): raise ValueError( @@ -8755,48 +8822,49 @@ def _handle_environment_variables(self): envs.SGLANG_OPT_FP8_WO_A_GEMM.set(False) def _handle_cache_compatibility(self): + cfg = resolving_view(self) if ( - self.disaggregation_decode_retraction_backup == "host_pool" - and self.disaggregation_mode != "decode" + cfg.disaggregation_decode_retraction_backup == "host_pool" + and cfg.disaggregation_mode != "decode" ): raise ValueError( "--disaggregation-decode-retraction-backup=host_pool is only " "supported on a PD decode server." ) if ( - self.disaggregation_decode_retraction_backup == "host_pool" - and self.dcp_size > 1 + cfg.disaggregation_decode_retraction_backup == "host_pool" + and cfg.dcp_size > 1 ): raise ValueError( "--disaggregation-decode-retraction-backup=host_pool does not " "support --dcp-size > 1." ) if ( - self.disaggregation_decode_retraction_backup == "host_pool" - and self.enable_priority_scheduling - and not self.disable_priority_preemption + cfg.disaggregation_decode_retraction_backup == "host_pool" + and cfg.enable_priority_scheduling + and not cfg.disable_priority_preemption ): raise ValueError( "--disaggregation-decode-retraction-backup=host_pool requires " "--disable-priority-preemption when priority scheduling is enabled." ) - if self.enable_hierarchical_cache and self.disable_radix_cache: + if cfg.enable_hierarchical_cache and cfg.disable_radix_cache: raise ValueError( "The arguments enable-hierarchical-cache and disable-radix-cache are mutually exclusive " "and cannot be used at the same time. Please use only one of them." ) - if self.disaggregation_decode_enable_offload_kvcache: - if self.disaggregation_mode != "decode": + if cfg.disaggregation_decode_enable_offload_kvcache: + if cfg.disaggregation_mode != "decode": raise ValueError( "The argument disaggregation-decode-enable-offload-kvcache is only supported for decode side." ) - if self.hicache_storage_backend is None: + if cfg.hicache_storage_backend is None: raise ValueError( "The argument disaggregation-decode-enable-offload-kvcache is only supported when hicache-storage-backend is provided." ) - if self.disaggregation_decode_retraction_backup == "host_pool": + if cfg.disaggregation_decode_retraction_backup == "host_pool": raise ValueError( "The arguments disaggregation-decode-enable-offload-kvcache and " "disaggregation-decode-retraction-backup=host_pool are mutually exclusive: " @@ -8810,7 +8878,8 @@ def _handle_cache_compatibility(self): raise ValueError("--swa-full-tokens-ratio should be in range (0, 1.0].") def _handle_deterministic_inference(self): - if self.rl_on_policy_target is not None: + cfg = resolving_view(self) + if cfg.rl_on_policy_target is not None: logger.warning( "Enable deterministic inference because of rl_on_policy_target." ) @@ -8824,8 +8893,8 @@ def _handle_deterministic_inference(self): # TODO remove this environment variable as a whole envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.set(True) - if self.enable_deterministic_inference: - if self.enable_aiter_allreduce_fusion: + if cfg.enable_deterministic_inference: + if cfg.enable_aiter_allreduce_fusion: logger.warning( "Disable --enable-aiter-allreduce-fusion because deterministic inference is enabled." ) @@ -8855,7 +8924,7 @@ def _handle_deterministic_inference(self): run_post_process_pass(self, _deterministic_sampling_backend) is_deepseek_model = False - if parse_connector_type(self.model_path) != ConnectorType.INSTANCE: + if parse_connector_type(cfg.model_path) != ConnectorType.INSTANCE: try: hf_config = self.get_model_config().hf_config model_arch = hf_config.architectures[0] @@ -8902,7 +8971,7 @@ def _handle_deterministic_inference(self): ) # Check TP size - if self.tp_size > 1: + if cfg.tp_size > 1: if is_hip(): # AMD: use 1-stage all-reduce kernel which is inherently deterministic # (each GPU reads all data from all GPUs, reduces locally in fixed order) @@ -8937,17 +9006,18 @@ def _handle_deterministic_inference(self): ) def _handle_unified_memory_pool(self): - if not self.enable_unified_memory: + cfg = resolving_view(self) + if not cfg.enable_unified_memory: return - if self.disaggregation_mode != "null": + if cfg.disaggregation_mode != "null": # Constraints of the whole-envelope transfer; see # UnifiedMLATokenToKVPool.get_contiguous_buf_infos. - assert self.disaggregation_transfer_backend == "mooncake", ( + assert cfg.disaggregation_transfer_backend == "mooncake", ( "--enable-unified-memory with PD disaggregation supports only " "the mooncake transfer backend; got " - f"{self.disaggregation_transfer_backend!r}." + f"{cfg.disaggregation_transfer_backend!r}." ) - assert self.pp_size == 1, ( + assert cfg.pp_size == 1, ( "--enable-unified-memory with PD disaggregation does not support " "pipeline parallelism (whole-envelope transfer has no per-layer " "entries to subset)." @@ -8956,24 +9026,24 @@ def _handle_unified_memory_pool(self): "--enable-unified-memory with PD disaggregation requires lazy " "compaction; unset SGLANG_DISABLE_LAZY_COMPACTION." ) - assert not self.enable_hisparse, ( + assert not cfg.enable_hisparse, ( "--enable-unified-memory with PD disaggregation is not compatible " "with --enable-hisparse: the decode-side HiSparse prealloc path " "ships host/C4 rows straight from the allocator, bypassing the " "virtual->physical translation the unified pool needs." ) - assert self.speculative_algorithm in (None, "DSPARK"), ( + assert cfg.speculative_algorithm in (None, "DSPARK"), ( "--enable-unified-memory only supports --speculative-algorithm " "DSPARK (chain draft); other speculative algorithms are not yet " "audited for the unified pool's virtual/dense loc translation. Got " - f"--speculative-algorithm={self.speculative_algorithm!r}." + f"--speculative-algorithm={cfg.speculative_algorithm!r}." ) - if self.speculative_algorithm == "DSPARK": - assert self.speculative_eagle_topk in (None, 1), ( + if cfg.speculative_algorithm == "DSPARK": + assert cfg.speculative_eagle_topk in (None, 1), ( "--enable-unified-memory + DSPARK supports a linear draft " "chain only (--speculative-eagle-topk in {None, 1}); tree " "verify is not audited for the unified pool. Got " - f"--speculative-eagle-topk={self.speculative_eagle_topk!r}." + f"--speculative-eagle-topk={cfg.speculative_eagle_topk!r}." ) # Both roles: verify routes to either backend depending on # --speculative-attention-mode. @@ -8987,14 +9057,14 @@ def _handle_unified_memory_pool(self): "not translate speculative verify indices to the unified " "pool's dense space yet." ) - assert not (self.enable_hierarchical_cache or self.enable_lmcache), ( + assert not (cfg.enable_hierarchical_cache or cfg.enable_lmcache), ( "--enable-unified-memory is not yet compatible with hierarchical / " "host-tiered KV cache (--enable-hierarchical-cache / --enable-lmcache): " "the unified-memory-pool init wires up no host pools, and its device mamba / " "full-attention slots are VIRTUAL — the host-offload path does not " "translate them to physical." ) - assert self.dcp_size == 1, ( + assert cfg.dcp_size == 1, ( "--enable-unified-memory is not yet compatible with decode context " "parallelism (--dcp-size > 1): the pool has no DCP-aware masked write " "path (UnifiedMHATokenToKVPool.set_kv_buffer asserts dcp_kv_mask is None), " @@ -9002,7 +9072,7 @@ def _handle_unified_memory_pool(self): ) # Only monolithic decode cuda-graph capture is wired; piecewise prefill # capture is not. Guard when the user opts into it. - _cg_cfg = self.cuda_graph_config + _cg_cfg = cfg.cuda_graph_config if _cg_cfg is not None and _cg_cfg.prefill.backend == Backend.TC_PIECEWISE: raise ValueError( "--enable-unified-memory supports monolithic (decode) " @@ -9018,12 +9088,13 @@ def _handle_page_major_kv_layout(self): # The unified pool stores state in the page-major envelope-strided layout, so # enabling it implies --enable-page-major-kv-layout — routing it through the # single page-major path + stride-aware Triton asserts (set before the guard). - if self.enable_unified_memory: + cfg = resolving_view(self) + if cfg.enable_unified_memory: self._declare( "_handle_page_major_kv_layout", enable_page_major_kv_layout=True, ) - if not self.enable_page_major_kv_layout: + if not cfg.enable_page_major_kv_layout: return # Only the Triton attention kernels read the strided 4-D envelope K/V # views; FA3 / FlashInfer do not. EXCEPTION: the unified-memory MLA pool @@ -9037,7 +9108,7 @@ def _handle_page_major_kv_layout(self): # page_table (in-kernel for captured decode, one funnel for eager). # flashmla / cutlass_mla share the create_flashmla block-table path and # can be added the same way once exercised. - if self.enable_unified_memory and self.use_mla_backend(): + if cfg.enable_unified_memory and self.use_mla_backend(): allowed_full = { "triton", "fa3", @@ -9074,10 +9145,10 @@ def _handle_page_major_kv_layout(self): decode_allowed.update({"cutedsl", "helion"}) prefill_allowed.update({"cutedsl", "helion"}) resolved_linear_decode = ( - self.linear_attn_decode_backend or self.linear_attn_backend + cfg.linear_attn_decode_backend or cfg.linear_attn_backend ) resolved_linear_prefill = ( - self.linear_attn_prefill_backend or self.linear_attn_backend + cfg.linear_attn_prefill_backend or cfg.linear_attn_backend ) assert resolved_linear_decode in decode_allowed | {None}, ( "--enable-page-major-kv-layout: linear-attention DECODE backend must " @@ -9089,28 +9160,29 @@ def _handle_page_major_kv_layout(self): f"be one of {sorted(prefill_allowed)} for the strided conv/SSM state; " f"got {resolved_linear_prefill!r}." ) - assert self.mamba_backend in (None, "triton"), ( + assert cfg.mamba_backend in (None, "triton"), ( "--enable-page-major-kv-layout requires the Triton Mamba kernels for " - f"the strided conv/SSM state; got {self.mamba_backend!r}. Pass " + f"the strided conv/SSM state; got {cfg.mamba_backend!r}. Pass " "--mamba-backend triton." ) def _handle_dllm_inference(self): - if self.dllm_algorithm is None: + cfg = resolving_view(self) + if cfg.dllm_algorithm is None: return # On AMD/HIP, disable cuda graph for DLLM (the attention_backend # resolution moved to the pipeline: arg_groups/overrides.py # _dllm_attention_backend, invoked below at its legacy slot). if is_hip(): if ( - self.cuda_graph_config.decode.backend != Backend.DISABLED - or self.cuda_graph_config.prefill.backend != Backend.DISABLED + cfg.cuda_graph_config.decode.backend != Backend.DISABLED + or cfg.cuda_graph_config.prefill.backend != Backend.DISABLED ): logger.warning( "Cuda graph is disabled for diffusion LLM inference on AMD GPUs" ) - self.cuda_graph_config.decode.backend = Backend.DISABLED - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED from sglang.srt.arg_groups.overrides import ( _dllm_attention_backend, @@ -9130,8 +9202,8 @@ def _handle_dllm_inference(self): run_post_process_pass(self, _dllm_page_size) - if not self.disable_radix_cache: - if self.enable_hierarchical_cache: + if not cfg.disable_radix_cache: + if cfg.enable_hierarchical_cache: logger.warning( "Hierarchical cache is disabled because of using diffusion LLM inference" ) @@ -9139,18 +9211,18 @@ def _handle_dllm_inference(self): "_handle_dllm_inference", enable_hierarchical_cache=False, ) - if self.enable_lmcache: + if cfg.enable_lmcache: logger.warning( "LMCache is disabled because of using diffusion LLM inference" ) self._declare("_handle_dllm_inference", enable_lmcache=False) - if self.enable_flexkv: + if cfg.enable_flexkv: logger.warning( "FlexKV is disabled because of using diffusion LLM inference" ) self._declare("_handle_dllm_inference", enable_flexkv=False) - if self.pp_size > 1: + if cfg.pp_size > 1: logger.warning( "Pipeline parallelism is disabled because of using diffusion LLM inference" ) @@ -9159,13 +9231,13 @@ def _handle_dllm_inference(self): pp_size=1, ) - if self.enable_lora: + if cfg.enable_lora: logger.warning( "Currently LoRA is not supported by diffusion LLM inference." ) self._declare("_handle_dllm_inference", enable_lora=False) - if self.disaggregation_mode != "null": + if cfg.disaggregation_mode != "null": logger.warning( "Currently disaggregation is not supported by diffusion LLM inference." ) @@ -9174,7 +9246,7 @@ def _handle_dllm_inference(self): disaggregation_mode="null", ) - if self.enable_mixed_chunk: + if cfg.enable_mixed_chunk: logger.warning( "Mixed chunked prefill is disabled because of using diffusion LLM inference." ) @@ -9185,43 +9257,43 @@ def _handle_dllm_inference(self): def _handle_asr_validation(self): """Validate transcription/ASR-specific server args.""" - if self.asr_max_buffer_seconds <= 0: + cfg = resolving_view(self) + if cfg.asr_max_buffer_seconds <= 0: raise ValueError( f"--asr-max-buffer-seconds must be positive " - f"(got {self.asr_max_buffer_seconds})." + f"(got {cfg.asr_max_buffer_seconds})." ) - if self.asr_max_concurrent_sessions <= 0: + if cfg.asr_max_concurrent_sessions <= 0: raise ValueError( f"--asr-max-concurrent-sessions must be positive " - f"(got {self.asr_max_concurrent_sessions})." + f"(got {cfg.asr_max_concurrent_sessions})." ) def _validate_prefill_decode_interval(self): - if self.prefill_decode_interval < 0: + cfg = resolving_view(self) + if cfg.prefill_decode_interval < 0: raise ValueError("--prefill-decode-interval must be non-negative.") def _handle_other_validations(self): - if self.default_chat_template_kwargs is not None and not isinstance( - self.default_chat_template_kwargs, dict + cfg = resolving_view(self) + if cfg.default_chat_template_kwargs is not None and not isinstance( + cfg.default_chat_template_kwargs, dict ): raise ValueError( "--default-chat-template-kwargs must decode to a JSON object" ) # Handle optimistic prefill validation - if ( - self.optimistic_prefill_attempts > 0 - and self.disaggregation_mode == "prefill" - ): - if self.pp_size > 1: + if cfg.optimistic_prefill_attempts > 0 and cfg.disaggregation_mode == "prefill": + if cfg.pp_size > 1: logger.warning("Optimistic prefill does not support pp_size > 1") self._declare( "_handle_other_validations", optimistic_prefill_attempts=0, ) - elif self.enable_hierarchical_cache and ( - self.hicache_storage_backend is not None - or self.hicache_write_policy != "write_back" + elif cfg.enable_hierarchical_cache and ( + cfg.hicache_storage_backend is not None + or cfg.hicache_write_policy != "write_back" ): logger.warning( "Optimistic prefill only supports L2 hierarchical cache " @@ -9242,37 +9314,35 @@ def _handle_other_validations(self): ) # Handle model inference tensor dump. - if self.debug_tensor_dump_output_folder is not None: + if cfg.debug_tensor_dump_output_folder is not None: logger.warning( "Cuda graph and server warmup are disabled because of using tensor dump mode" ) - self.cuda_graph_config.decode.backend = Backend.DISABLED - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED self._declare("_handle_other_validations", skip_server_warmup=True) - if self.msprobe_dump_config is not None: + if cfg.msprobe_dump_config is not None: logger.warning( "When msProbe is enabled, " "cuda graph is disabled because msProbe only supports dump in eager mode, " "warmup is disabled(skip_server_warmup=True) because there is no need to dump data for this stage." ) - self.cuda_graph_config.decode.backend = Backend.DISABLED - self.cuda_graph_config.prefill.backend = Backend.DISABLED + cfg.cuda_graph_config.decode.backend = Backend.DISABLED + cfg.cuda_graph_config.prefill.backend = Backend.DISABLED self._declare("_handle_other_validations", skip_server_warmup=True) # Validate limit_mm_per_prompt modalities - if self.limit_mm_data_per_request: - if isinstance(self.limit_mm_data_per_request, str): + if cfg.limit_mm_data_per_request: + if isinstance(cfg.limit_mm_data_per_request, str): self._declare( "_handle_other_validations", - limit_mm_data_per_request=json.loads( - self.limit_mm_data_per_request - ), + limit_mm_data_per_request=json.loads(cfg.limit_mm_data_per_request), ) - if isinstance(self.limit_mm_data_per_request, dict): + if isinstance(cfg.limit_mm_data_per_request, dict): allowed_modalities = {"image", "video", "audio"} - for modality in self.limit_mm_data_per_request.keys(): + for modality in cfg.limit_mm_data_per_request.keys(): if modality not in allowed_modalities: raise ValueError( f"Invalid modality '{modality}' in --limit-mm-data-per-request." @@ -9280,25 +9350,24 @@ def _handle_other_validations(self): ) # Validate preferred_sampling_params - if self.preferred_sampling_params: - if isinstance(self.preferred_sampling_params, str): + if cfg.preferred_sampling_params: + if isinstance(cfg.preferred_sampling_params, str): self._declare( "_handle_other_validations", - preferred_sampling_params=json.loads( - self.preferred_sampling_params - ), + preferred_sampling_params=json.loads(cfg.preferred_sampling_params), ) # Validate preferred_sampling_params doesn't use tokenizer-dependent features - if self.skip_tokenizer_init: + if cfg.skip_tokenizer_init: from sglang.srt.sampling.sampling_params import SamplingParams - test_params = SamplingParams(**self.preferred_sampling_params) + test_params = SamplingParams(**cfg.preferred_sampling_params) # raises if tokenizer-dependent features used test_params.normalize(None) def _handle_crash_dump_env(self): - if not self.crash_dump_folder: + cfg = resolving_view(self) + if not cfg.crash_dump_folder: return _CUDA_COREDUMP_DEFAULTS = { "CUDA_ENABLE_COREDUMP_ON_EXCEPTION": "1", @@ -9308,7 +9377,7 @@ def _handle_crash_dump_env(self): "skip_nonrelocated_elf_images,skip_global_memory," "skip_shared_memory,skip_local_memory,skip_constbank_memory" ), - "CUDA_COREDUMP_FILE": f"{self.crash_dump_folder}/%h/core.cuda.%t.%p", + "CUDA_COREDUMP_FILE": f"{cfg.crash_dump_folder}/%h/core.cuda.%t.%p", "CUDA_COREDUMP_PIPE": "/tmp/corepipe.cuda.%h.%p", } for key, value in _CUDA_COREDUMP_DEFAULTS.items(): @@ -9338,7 +9407,8 @@ def _handle_crash_dump_env(self): ) def _handle_debug_utils(self): - if is_in_ci() and self.soft_watchdog_timeout is None: + cfg = resolving_view(self) + if is_in_ci() and cfg.soft_watchdog_timeout is None: logger.info("Set soft_watchdog_timeout since in CI") self._declare("_handle_debug_utils", soft_watchdog_timeout=300) @@ -9676,6 +9746,7 @@ def ssl_verify(self): def get_model_config(self): # Lazy init to avoid circular import + cfg = resolving_view(self) from sglang.srt.configs.model_config import ModelConfig memo = getattr(self, "model_config", None) @@ -9688,12 +9759,12 @@ def get_model_config(self): # object-store URI, so its field is not the key. A configuration a # fixture supplied carries no key and is handed back as it is. built_from = getattr(self, "_model_config_built_from", None) - if built_from is None or built_from == self.model_path: + if built_from is None or built_from == cfg.model_path: return memo model_config = ModelConfig.from_server_args(self) self.model_config = model_config - self._model_config_built_from = self.model_path + self._model_config_built_from = cfg.model_path if model_config.is_hybrid_swa: logger.info( "Hybrid SWA model detected. architectures=%s", @@ -10301,12 +10372,13 @@ def validate_buckets_rule(self, arg_name: str, buckets_rule: List[str]): ), f"{arg_name} custom rule bucket values should be non-negative" def adjust_mem_fraction_for_vlm(self, model_config): + cfg = resolving_view(self) vision_config = getattr(model_config.hf_config, "vision_config", None) if vision_config is None: return # roughly reduce the mem_fraction_static base on params of Vit - original_server_arg_mem_fraction = self.mem_fraction_static + original_server_arg_mem_fraction = cfg.mem_fraction_static # a base mem_fraction_static factor for regular Vit base_mem_fraction_reduction_ratio = 0.95 @@ -10340,6 +10412,7 @@ def adjust_mem_fraction_for_vlm(self, model_config): ) def validate_transfer_engine(self): + cfg = resolving_view(self) try: mooncake_available = importlib.util.find_spec("mooncake.engine") is not None except (ModuleNotFoundError, ValueError): @@ -10349,7 +10422,7 @@ def validate_transfer_engine(self): "Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend." ) return False - elif self.enable_memory_saver: + elif cfg.enable_memory_saver: logger.warning( "Memory saver is enabled, which is not compatible with TransferEngine. Does not support using TransferEngine as remote instance weight loader backend." ) @@ -10483,7 +10556,8 @@ def describe_kv_events_publisher(self) -> Optional[dict]: } def should_report_expert_balancedness(self) -> bool: - return self.expert_balancedness_report_mode != "off" + cfg = resolving_view(self) + return cfg.expert_balancedness_report_mode != "off" def should_log_expert_balancedness_to_server_log(self) -> bool: return self.expert_balancedness_report_mode in ("server_log", "both") From 1c2b3ee822581fe1ab3bc0cf0bba45af8e0d103d Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 11:46:40 +0000 Subject: [PATCH 09/26] config: the readers resolution calls with the record follow too The last of the mid-resolution readers outside the pipeline's own modules: `ModelConfig.from_server_args` (23 reads), the spec-algo hook, the adaptive-spec support check, the CP strategy binder and the BCG predicate. They are called *by* resolution with the record in hand, so they have the same problem as the handlers -- a declaration-only resolver leaves the field holding the raw input. `resolving_view` is imported inside the function at these sites: they sit under `configs/`, `layers/` and `speculative/`, and a module-level import of `arg_groups.overrides` there would be a new import edge into the resolution pipeline. `test_model_config_reads_resolved_input` learns the spelling: a local bound to `resolving_view(sa)` / `resolved_view(sa)` / `sa._resolved()` is the record for scanning purposes, so its two pins keep describing the reads they were written for. 20 launch shapes, zero differences in the resolution result. --- python/sglang/srt/configs/model_config.py | 52 +++++++++---------- python/sglang/srt/layers/cp/base.py | 7 ++- python/sglang/srt/layers/cp/bcg.py | 9 ++-- .../srt/speculative/adaptive_spec_params.py | 18 +++---- python/sglang/srt/speculative/spec_info.py | 5 +- .../sglang/srt/speculative/spec_registry.py | 5 +- .../test_model_config_reads_resolved_input.py | 40 ++++++++++++-- 7 files changed, 91 insertions(+), 45 deletions(-) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 32806be688e0..5ec73e5410c6 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -589,46 +589,46 @@ def from_server_args( context_length: Optional[int] = None, **kwargs, ): + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) quantization = ( - server_args.speculative_draft_model_quantization + cfg.speculative_draft_model_quantization if is_draft_model - else server_args.quantization + else cfg.quantization ) override_config_file = ( - server_args.decrypted_draft_config_file + cfg.decrypted_draft_config_file if is_draft_model - else server_args.decrypted_config_file + else cfg.decrypted_config_file ) return ModelConfig( - model_path=model_path or server_args.model_path, - trust_remote_code=server_args.trust_remote_code, - revision=model_revision or server_args.revision, + model_path=model_path or cfg.model_path, + trust_remote_code=cfg.trust_remote_code, + revision=model_revision or cfg.revision, context_length=( - context_length - if context_length is not None - else server_args.context_length + context_length if context_length is not None else cfg.context_length ), - model_override_args=server_args.json_model_override_args, - is_embedding=server_args.is_embedding, - enable_multimodal=server_args.enable_multimodal, - dtype=server_args.dtype, + model_override_args=cfg.json_model_override_args, + is_embedding=cfg.is_embedding, + enable_multimodal=cfg.enable_multimodal, + dtype=cfg.dtype, quantization=quantization, - model_impl=server_args.model_impl, - sampling_defaults=server_args.sampling_defaults, - quantize_and_serve=server_args.quantize_and_serve, + model_impl=cfg.model_impl, + sampling_defaults=cfg.sampling_defaults, + quantize_and_serve=cfg.quantize_and_serve, override_config_file=override_config_file, - is_multi_layer_eagle=server_args.enable_multi_layer_eagle, - language_only=server_args.language_only, - language_model_only=server_args.language_model_only, - encoder_only=server_args.encoder_only, + is_multi_layer_eagle=cfg.enable_multi_layer_eagle, + language_only=cfg.language_only, + language_model_only=cfg.language_model_only, + encoder_only=cfg.encoder_only, is_draft_model=is_draft_model, is_draft_quantization_explicit=( - is_draft_model - and server_args._speculative_draft_quantization_explicitly_set + is_draft_model and cfg._speculative_draft_quantization_explicitly_set ), - disable_hybrid_swa_memory=server_args.disable_hybrid_swa_memory, - model_config_parser=server_args.model_config_parser, - speculative_algorithm=server_args.speculative_algorithm, + disable_hybrid_swa_memory=cfg.disable_hybrid_swa_memory, + model_config_parser=cfg.model_config_parser, + speculative_algorithm=cfg.speculative_algorithm, **kwargs, ) diff --git a/python/sglang/srt/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index 4191089ed6a8..b58cc68bf620 100644 --- a/python/sglang/srt/layers/cp/base.py +++ b/python/sglang/srt/layers/cp/base.py @@ -239,6 +239,9 @@ def _is_dsa_active() -> bool: def init_cp_strategy(server_args: ServerArgs) -> None: """Bind the configured CP strategy for this process.""" + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) global _STRATEGY if not getattr(server_args, "enable_prefill_cp", False): @@ -250,7 +253,7 @@ def init_cp_strategy(server_args: ServerArgs) -> None: _STRATEGY = None return - kind = ContextParallelStrategyKind.from_string(server_args.cp_strategy) + kind = ContextParallelStrategyKind.from_string(cfg.cp_strategy) if kind == ContextParallelStrategyKind.ZIGZAG: from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy @@ -262,7 +265,7 @@ def init_cp_strategy(server_args: ServerArgs) -> None: else: raise ValueError( f"Unsupported cp_strategy kind {kind} for " - f"cp_strategy={server_args.cp_strategy!r}" + f"cp_strategy={cfg.cp_strategy!r}" ) diff --git a/python/sglang/srt/layers/cp/bcg.py b/python/sglang/srt/layers/cp/bcg.py index 6abbddcf03a0..6137ef091604 100644 --- a/python/sglang/srt/layers/cp/bcg.py +++ b/python/sglang/srt/layers/cp/bcg.py @@ -42,12 +42,15 @@ def supports_prefill_cp_bcg(server_args: ServerArgs) -> bool: """Return whether the selected prefill-CP configuration supports BCG.""" + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) resolved = server_args._resolved() prefill_attention_backend, _ = server_args._resolved_attention_backends() return ( - server_args.enable_prefill_cp - and resolved.attn_cp_size == server_args.tp_size - and server_args.cp_strategy == "zigzag" + cfg.enable_prefill_cp + and resolved.attn_cp_size == cfg.tp_size + and cfg.cp_strategy == "zigzag" and prefill_attention_backend == "trtllm_mha" ) diff --git a/python/sglang/srt/speculative/adaptive_spec_params.py b/python/sglang/srt/speculative/adaptive_spec_params.py index d8b24dce0100..c3c252f82310 100644 --- a/python/sglang/srt/speculative/adaptive_spec_params.py +++ b/python/sglang/srt/speculative/adaptive_spec_params.py @@ -49,19 +49,19 @@ def adaptive_unsupported_reason(server_args: ServerArgs) -> str | None: """Return why adaptive spec cannot run under the given server args, or None if supported.""" + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) from sglang.srt.arg_groups.overrides import resolved_view - if server_args.speculative_algorithm not in ("EAGLE", "EAGLE3"): + if cfg.speculative_algorithm not in ("EAGLE", "EAGLE3"): return ( - f"speculative_algorithm={server_args.speculative_algorithm} " + f"speculative_algorithm={cfg.speculative_algorithm} " "(only EAGLE/EAGLE3 are supported)" ) - if ( - server_args.speculative_eagle_topk is not None - and server_args.speculative_eagle_topk != 1 - ): + if cfg.speculative_eagle_topk is not None and cfg.speculative_eagle_topk != 1: return ( - f"speculative_eagle_topk={server_args.speculative_eagle_topk} " + f"speculative_eagle_topk={cfg.speculative_eagle_topk} " "(only topk=1 is supported)" ) if resolved_view(server_args).enable_dp_attention: @@ -74,12 +74,12 @@ def adaptive_unsupported_reason(server_args: ServerArgs) -> str | None: "enable_multi_layer_eagle=True is not supported " "(MultiLayerEagleWorkerV2 does not implement adaptive)" ) - if server_args.enable_two_batch_overlap: + if cfg.enable_two_batch_overlap: return ( "enable_two_batch_overlap=True is not supported " "(adaptive state swap would discard the TboAttnBackend wrapper)" ) - if server_args.enable_pdmux: + if cfg.enable_pdmux: return ( "enable_pdmux=True is not supported " "(adaptive state swap does not update decode_attn_backend_group)" diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index 03967668e83c..97ca92d31f0f 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -254,6 +254,9 @@ def get_num_tokens_per_bs_for_target_verify( def create_worker( self, server_args: ServerArgs ) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]: + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) assert ( not self.is_none() ), "Cannot create worker for NONE speculative algorithm." @@ -283,7 +286,7 @@ def create_worker( # EAGLE / EAGLE3 / STANDALONE / MULTI_LAYER always use the V2 worker, # even with overlap disabled (scheduler drives it synchronously). - if self.is_eagle() and server_args.enable_multi_layer_eagle: + if self.is_eagle() and cfg.enable_multi_layer_eagle: from sglang.srt.speculative.multi_layer_eagle_worker_v2 import ( MultiLayerEagleWorkerV2, ) diff --git a/python/sglang/srt/speculative/spec_registry.py b/python/sglang/srt/speculative/spec_registry.py index 1e3b84853c60..933d4e05d7b0 100644 --- a/python/sglang/srt/speculative/spec_registry.py +++ b/python/sglang/srt/speculative/spec_registry.py @@ -108,7 +108,10 @@ def handle_server_args(self, server_args: ServerArgs) -> None: pass def create_worker(self, server_args: ServerArgs) -> Type: - if not server_args.disable_overlap_schedule and not self.supports_overlap: + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) + if not cfg.disable_overlap_schedule and not self.supports_overlap: raise ValueError( f"Speculative algorithm {self.name} does not support overlap scheduling." ) diff --git a/test/registered/unit/server_args/test_model_config_reads_resolved_input.py b/test/registered/unit/server_args/test_model_config_reads_resolved_input.py index e21d6a462fe7..8dabf06aa34c 100644 --- a/test/registered/unit/server_args/test_model_config_reads_resolved_input.py +++ b/test/registered/unit/server_args/test_model_config_reads_resolved_input.py @@ -125,6 +125,13 @@ def _registry_collection_is_after_the_build(): def _server_args_names(tree, path): + """Every local that names the record, including the read views over it. + + A resolution-time reader reads through `resolving_view(server_args)` (the + declaration stash over the fields): declaration-only resolvers write no + field, so a field read there answers with the raw input. `cfg.dtype` after `cfg = resolving_view(sa)` is + the same read this scan is looking for, so the local it binds counts. + """ names = {"self"} if path.name == "server_args.py" else {"server_args"} for node in ast.walk(tree): if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): @@ -142,6 +149,32 @@ def _server_args_names(tree, path): continue if text == "ServerArgs": names.add(arg.arg) + # `cfg = resolving_view(server_args)` / `resolved_view(server_args)` + for _ in range(2): # a view over a view-holding local is still one + for node in ast.walk(tree): + if not isinstance(node, ast.Assign): + continue + value = node.value + bare = ( + isinstance(value, ast.Call) + and isinstance(value.func, ast.Name) + and value.func.id in ("resolving_view", "resolved_view") + and value.args + and isinstance(value.args[0], ast.Name) + and value.args[0].id in names + ) + # `resolved = self._resolved()` is the same view, spelled as the + # record's own member. + member = ( + isinstance(value, ast.Call) + and isinstance(value.func, ast.Attribute) + and value.func.attr == "_resolved" + and isinstance(value.func.value, ast.Name) + and value.func.value.id in names + ) + if not (bare or member): + continue + names |= {t.id for t in node.targets if isinstance(t, ast.Name)} return names @@ -191,11 +224,12 @@ def _late_resolution_fields(): for name in ( "server_args.py", "arg_groups/overrides.py", - "utils/template_detection.py", + "parser/template_detection.py", ): path = _SRT / name - if not path.exists(): - continue + # A named file that moved away has to be loud; skipping it silently + # leaves the scan believing it read a module it never opened. + assert path.exists(), f"{name} is not where this scan looks for it" tree = _parsed(path) for node in ast.walk(tree): if not isinstance(node, ast.Call): From d2ce9d53966de8e8810ae29513556c2179de61dc Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 12:01:55 +0000 Subject: [PATCH 10/26] config: the record's own members answer from the declarations `ServerArgs`' public members read their fields, which is the same problem the handlers had: a declaration-only resolver leaves the field holding the raw input, so `max_speculative_num_draft_tokens`, `is_ep_joiner`, `is_startup_weight_load_overlap`, the expert-balancedness predicates and `describe_kv_events_publisher` could answer for what was typed instead of what resolution decided. They read through the view now, and the file settles on one spelling for it: `cfg = resolving_view(self)`, replacing the `resolved = resolved_view(self)` / `resolved = self._resolved()` mix this file had accumulated. One exception keeps its own name -- `describe_kv_events_publisher` already binds `cfg` to a `KVEventsConfig`, and the view has to not shadow it. --- python/sglang/srt/server_args.py | 219 +++++++++++++++++-------------- 1 file changed, 118 insertions(+), 101 deletions(-) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index de4a5cc44535..59623be30b6e 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -9708,15 +9708,21 @@ def engine_info_bootstrap_url(self): @property def is_ep_joiner(self) -> bool: """True for processes launched as elastic-EP joiners.""" - return self.ep_join_mode in ("scale", "recover") + cfg = resolving_view(self) + + return cfg.ep_join_mode in ("scale", "recover") @property def is_ep_scale_joiner(self) -> bool: - return self.ep_join_mode == "scale" + cfg = resolving_view(self) + + return cfg.ep_join_mode == "scale" @property def is_startup_weight_load_overlap(self) -> bool: - return self.startup_weight_load_mode == "overlap" + cfg = resolving_view(self) + + return cfg.startup_weight_load_mode == "overlap" def ssl_verify(self): """Return the value for the requests library's verify= parameter. @@ -9849,20 +9855,22 @@ def max_speculative_num_draft_tokens(self) -> Optional[int]: sizing fills `speculative_num_draft_tokens` in), and a cache filled that early would keep answering with it. """ + cfg = resolving_view(self) + memo = self.__dict__.get("_max_speculative_num_draft_tokens") if memo is not None: return memo - if self.speculative_num_draft_tokens is None: + if cfg.speculative_num_draft_tokens is None: result = None - elif not self.speculative_adaptive: - result = self.speculative_num_draft_tokens + elif not cfg.speculative_adaptive: + result = cfg.speculative_num_draft_tokens else: from sglang.srt.speculative.adaptive_spec_params import ( resolve_candidate_steps_from_config, ) candidate_steps = resolve_candidate_steps_from_config( - cfg_path=self.speculative_adaptive_config, + cfg_path=cfg.speculative_adaptive_config, ) # TODO: adaptive spec currently requires topk=1, so each runtime # state needs steps + 1 draft-token slots. Revisit this if topk>1 @@ -9907,15 +9915,17 @@ def _check_two_batch_overlap(self): # DP TP-MoE path (overlapping the DP all_gatherv / reduce_scatterv with # the other ubatch's compute), which requires DP attention. Enabling it # there needs no extra opt-in env flag. + cfg = resolving_view(self) + cp_tbo = ( is_hip() - and self.enable_dsa_prefill_context_parallel - and self.dsa_prefill_cp_mode == "round-robin-split" + and cfg.enable_dsa_prefill_context_parallel + and cfg.dsa_prefill_cp_mode == "round-robin-split" ) if ( - self.enable_two_batch_overlap - and self.moe_a2a_backend == "none" - and not self.enable_dp_attention + cfg.enable_two_batch_overlap + and cfg.moe_a2a_backend == "none" + and not cfg.enable_dp_attention and not cp_tbo ): raise ValueError( @@ -9925,20 +9935,22 @@ def _check_two_batch_overlap(self): ) def check_server_args(self): + cfg = resolving_view(self) + # Check parallel size constraints - if self.ep_join_mode != "scale": + if cfg.ep_join_mode != "scale": assert ( - self.tp_size * self.pp_size - ) % self.nnodes == 0, "tp_size must be divisible by number of nodes" + cfg.tp_size * cfg.pp_size + ) % cfg.nnodes == 0, "tp_size must be divisible by number of nodes" assert ( - self.pp_max_micro_batch_size is None or self.pp_max_micro_batch_size >= 1 + cfg.pp_max_micro_batch_size is None or cfg.pp_max_micro_batch_size >= 1 ), ( "pp_max_micro_batch_size must be a positive integer or None (for auto-compute). " - f"Got: {self.pp_max_micro_batch_size}" + f"Got: {cfg.pp_max_micro_batch_size}" ) - assert not (self.disable_cuda_graph_padding and self.enable_torch_compile), ( + assert not (cfg.disable_cuda_graph_padding and cfg.enable_torch_compile), ( "--disable-cuda-graph-padding is incompatible with --enable-torch-compile. " "With padding disabled, every distinct batch size gets its own torch.compile + " "Triton autotune cycle (O(max_batch_size) compilations) instead of the small fixed " @@ -9946,67 +9958,67 @@ def check_server_args(self): "Remove --disable-cuda-graph-padding or --enable-torch-compile." ) - if self.pp_size > 1: + if cfg.pp_size > 1: assert ( - self.disable_overlap_schedule and self.speculative_algorithm is None + cfg.disable_overlap_schedule and cfg.speculative_algorithm is None ), "Pipeline parallelism is not compatible with overlap schedule, speculative decoding" - assert self.min_free_slots_delay is None, ( + assert cfg.min_free_slots_delay is None, ( "--min-free-slots-delay is not supported with pipeline " "parallelism: allocatable slots per microbatch are bounded by " "pp-max-micro-batch-size, so the threshold may never be reached" ) assert not ( - self.dp_size > 1 and self.nnodes != 1 and not self.enable_dp_attention + cfg.dp_size > 1 and cfg.nnodes != 1 and not cfg.enable_dp_attention ), "multi-node data parallel is not supported unless dp attention!" - assert self.base_gpu_id >= 0, "base_gpu_id must be non-negative" - assert self.gpu_id_step >= 1, "gpu_id_step must be positive" + assert cfg.base_gpu_id >= 0, "base_gpu_id must be non-negative" + assert cfg.gpu_id_step >= 1, "gpu_id_step must be positive" - assert self.moe_dense_tp_size in ( + assert cfg.moe_dense_tp_size in ( None, 1, - self.tp_size, + cfg.tp_size, ), "moe_dense_tp_size only supports None, 1, or tp_size currently" # Check served model name to not have colon as it is reserved for LoRA adapter syntax - if not is_runai_obj_uri(self.served_model_name): - assert ":" not in self.served_model_name, ( + if not is_runai_obj_uri(cfg.served_model_name): + assert ":" not in cfg.served_model_name, ( "served_model_name cannot contain a colon (':') character. " "The colon is reserved for the 'model:adapter' syntax used in LoRA adapter specification. " - f"Invalid value: '{self.served_model_name}'" + f"Invalid value: '{cfg.served_model_name}'" ) # Check LoRA self.check_lora_server_args() # Check speculative decoding - if self.speculative_algorithm is not None: + if cfg.speculative_algorithm is not None: assert ( - not self.enable_mixed_chunk + not cfg.enable_mixed_chunk ), "enable_mixed_chunk is required for speculative decoding" # Check chunked prefill # Skip validation if chunked prefill is disabled (i.e., size <= 0). # Skip validation if disaggregation mode is decode. - if self.chunked_prefill_size > 0 and self.disaggregation_mode != "decode": + if cfg.chunked_prefill_size > 0 and cfg.disaggregation_mode != "decode": assert ( - self.chunked_prefill_size % self.page_size == 0 + cfg.chunked_prefill_size % cfg.page_size == 0 ), "chunked_prefill_size must be divisible by page_size" # Check pdmux - if self.enable_pdmux: + if cfg.enable_pdmux: assert ( - self.pp_size == 1 + cfg.pp_size == 1 ), "PD-Multiplexing is only supported with pipeline parallelism disabled (pp_size=1)." assert ( - self.chunked_prefill_size == -1 + cfg.chunked_prefill_size == -1 ), "PD-Multiplexing is not compatible with chunked prefill." assert ( - self.disaggregation_mode == "null" + cfg.disaggregation_mode == "null" ), "PD-Multiplexing is not compatible with disaggregation mode." assert ( - self.disable_overlap_schedule + cfg.disable_overlap_schedule ), "PD-Multiplexing is not compatible with overlap schedule." # NOTE: CUDA Green Context may encounter potential issues with CudaGraph on torch 2.7.x – 2.8.x, leading to performance degradation. @@ -10019,41 +10031,39 @@ def check_server_args(self): " Please manually install torch 2.6.x." ) - assert self.tokenizer_worker_num > 0, "Tokenizer worker num must >= 1" - assert self.detokenizer_worker_num > 0, "Detokenizer worker num must >= 1" + assert cfg.tokenizer_worker_num > 0, "Tokenizer worker num must >= 1" + assert cfg.detokenizer_worker_num > 0, "Detokenizer worker num must >= 1" assert ( - self.mm_processor_worker_num >= 0 + cfg.mm_processor_worker_num >= 0 ), "Multimodal processor worker num must >= 0" - assert self.mm_io_worker_num >= 0, "Multimodal I/O worker num must >= 0" - self.validate_buckets_rule( - "--prompt-tokens-buckets", self.prompt_tokens_buckets - ) + assert cfg.mm_io_worker_num >= 0, "Multimodal I/O worker num must >= 0" + self.validate_buckets_rule("--prompt-tokens-buckets", cfg.prompt_tokens_buckets) self.validate_buckets_rule( - "--generation-tokens-buckets", self.generation_tokens_buckets + "--generation-tokens-buckets", cfg.generation_tokens_buckets ) # Check scheduling policy - if self.enable_priority_scheduling: - assert self.schedule_policy in [ + if cfg.enable_priority_scheduling: + assert cfg.schedule_policy in [ "fcfs", "lof", - ], f"To use priority scheduling, schedule_policy must be 'fcfs' or 'lof'. '{self.schedule_policy}' is not supported." - if self.default_priority_value is None: + ], f"To use priority scheduling, schedule_policy must be 'fcfs' or 'lof'. '{cfg.schedule_policy}' is not supported." + if cfg.default_priority_value is None: logger.warning( "--default-priority-value is not set while --enable-priority-scheduling is enabled. " "Requests without explicit priority will have priority=None, " "resulting in priority='None' string labels in Prometheus metrics." ) else: - if self.disable_priority_preemption: + if cfg.disable_priority_preemption: logger.warning( "--disable-priority-preemption has no effect without --enable-priority-scheduling" ) - if self.default_priority_value is not None: + if cfg.default_priority_value is not None: logger.warning( "--default-priority-value has no effect without --enable-priority-scheduling" ) - if self.retraction_policy == "priority" and not self.enable_priority_scheduling: + if cfg.retraction_policy == "priority" and not cfg.enable_priority_scheduling: raise ValueError( "--retraction-policy priority requires --enable-priority-scheduling" ) @@ -10069,23 +10079,23 @@ def check_server_args(self): run_post_process_pass(self, _hisparse_validation) assert ( - self.schedule_conservativeness >= 0 + cfg.schedule_conservativeness >= 0 ), "schedule_conservativeness must be non-negative" - if self.model_impl == "mindspore": + if cfg.model_impl == "mindspore": assert is_npu(), "MindSpore model impl is only supported on Ascend npu." # Check metrics labels if ( - not self.tokenizer_metrics_custom_labels_header - and self.tokenizer_metrics_allowed_custom_labels + not cfg.tokenizer_metrics_custom_labels_header + and cfg.tokenizer_metrics_allowed_custom_labels ): raise ValueError( "Please set --tokenizer-metrics-custom-labels-header when setting --tokenizer-metrics-allowed-custom-labels." ) # Check metrics exporters - if self.export_metrics_to_file and self.export_metrics_to_file_dir is None: + if cfg.export_metrics_to_file and cfg.export_metrics_to_file_dir is None: raise ValueError( "--export-metrics-to-file-dir is required when --export-metrics-to-file is enabled" ) @@ -10094,65 +10104,67 @@ def check_server_args(self): self._check_two_batch_overlap() # Check communications compression - if self.enable_quant_communications and self.tp_size == 1: + if cfg.enable_quant_communications and cfg.tp_size == 1: raise ValueError( "Communications quantization is only used with tp_size != 1" ) - if self.enable_quant_communications and self.device != "npu": + if cfg.enable_quant_communications and cfg.device != "npu": raise ValueError( "Communications quantization is only supported for NPU device" ) # grpc_port is None for HTTP-only launches, so the == comparison is # already False there; no explicit None check needed. - if not (self.smg_grpc_mode or self.grpc_mode) and self.grpc_port == self.port: + if not (cfg.smg_grpc_mode or cfg.grpc_mode) and cfg.grpc_port == cfg.port: raise ValueError( - f"--grpc-port ({self.grpc_port}) must differ from --port ({self.port})" + f"--grpc-port ({cfg.grpc_port}) must differ from --port ({cfg.port})" ) # TODO: Also validate grpc_port != metrics_http_port and grpc_port != nccl_port # to avoid opaque bind errors at runtime. Deferred because metrics_http_port # and nccl_port have dynamic defaults that may not be resolved yet here. - if self.gc_threshold: - if not (1 <= len(self.gc_threshold) <= 3): + if cfg.gc_threshold: + if not (1 <= len(cfg.gc_threshold) <= 3): raise ValueError( "When setting gc_threshold, it must contain 1 to 3 integers." ) - if self.kv_canary_sweep_interval > 0 and self.kv_canary == "none": + if cfg.kv_canary_sweep_interval > 0 and cfg.kv_canary == "none": raise ValueError( "--kv-canary-sweep-interval requires --kv-canary in {log, raise}" ) def check_lora_server_args(self): - assert self.max_loras_per_batch > 0, "max_loras_per_batch must be positive" + cfg = resolving_view(self) + + assert cfg.max_loras_per_batch > 0, "max_loras_per_batch must be positive" # Enable LoRA if any LoRA paths are provided for backward compatibility. - if self.lora_paths: - if self.enable_lora is None: + if cfg.lora_paths: + if cfg.enable_lora is None: self._late_resolution("check_lora_server_args", enable_lora=True) logger.warning( "--enable-lora is set to True because --lora-paths is provided." ) - elif self.enable_lora is False: + elif cfg.enable_lora is False: logger.warning( "--enable-lora is set to False, any provided lora_paths will be ignored." ) - if self.enable_lora: - if self.enable_lora_overlap_loading is None: + if cfg.enable_lora: + if cfg.enable_lora_overlap_loading is None: self._late_resolution( "check_lora_server_args", enable_lora_overlap_loading=False ) - if self.enable_lora_overlap_loading: + if cfg.enable_lora_overlap_loading: # TODO (glenliu21): use some sort of buffer with eviction instead of enforcing a limit - max_loaded_loras_limit = self.max_loras_per_batch * 2 + max_loaded_loras_limit = cfg.max_loras_per_batch * 2 assert ( - self.max_loaded_loras is not None - and self.max_loaded_loras <= max_loaded_loras_limit + cfg.max_loaded_loras is not None + and cfg.max_loaded_loras <= max_loaded_loras_limit ), ( "Enabling LoRA overlap loading requires pinning LoRA adapter weights in CPU memory, " f"so --max-loaded-loras must be less than or equal to double --max-loras-per-batch: {max_loaded_loras_limit}" @@ -10162,9 +10174,9 @@ def check_lora_server_args(self): self._check_lora_speculative_compatibility() # Parse lora_paths - if isinstance(self.lora_paths, list): + if isinstance(cfg.lora_paths, list): parsed_lora_paths = [] - for lora_path in self.lora_paths: + for lora_path in cfg.lora_paths: if isinstance(lora_path, str): if "=" in lora_path: name, path = lora_path.split("=", 1) @@ -10202,7 +10214,7 @@ def check_lora_server_args(self): self._late_resolution( "check_lora_server_args", lora_paths=parsed_lora_paths ) - elif isinstance(self.lora_paths, dict): + elif isinstance(cfg.lora_paths, dict): self._late_resolution( "check_lora_server_args", lora_paths=[ @@ -10212,56 +10224,56 @@ def check_lora_server_args(self): lora_path=v, pinned=False, ) - for k, v in self.lora_paths.items() + for k, v in cfg.lora_paths.items() ], ) - elif self.lora_paths is None: + elif cfg.lora_paths is None: self._late_resolution("check_lora_server_args", lora_paths=[]) else: raise ValueError( - f"Invalid type for --lora-paths: {type(self.lora_paths)}. " + f"Invalid type for --lora-paths: {type(cfg.lora_paths)}. " "Expected a list or a dictionary." ) # Normalize target modules to a set; keep {"all"} as a sentinel # that gets resolved model-awarely in lora_manager.init_lora_shapes(). - if self.lora_target_modules: + if cfg.lora_target_modules: self._late_resolution( "check_lora_server_args", - lora_target_modules=set(self.lora_target_modules), + lora_target_modules=set(cfg.lora_target_modules), ) - if "all" in self.lora_target_modules: + if "all" in cfg.lora_target_modules: assert ( - len(self.lora_target_modules) == 1 + len(cfg.lora_target_modules) == 1 ), "If 'all' is specified in --lora-target-modules, it should be the only module specified." # Ensure sufficient information is provided for LoRA initialization. - assert self.lora_paths or ( - self.max_lora_rank and self.lora_target_modules + assert cfg.lora_paths or ( + cfg.max_lora_rank and cfg.lora_target_modules ), "When no initial --lora-paths is provided, you need to specify both --max-lora-rank and --lora-target-modules for LoRA initialization." # Validate max_loaded_loras - if self.max_loaded_loras is not None: - assert self.max_loaded_loras >= self.max_loras_per_batch, ( + if cfg.max_loaded_loras is not None: + assert cfg.max_loaded_loras >= cfg.max_loras_per_batch, ( "max_loaded_loras should be greater than or equal to max_loras_per_batch. " - f"max_loaded_loras={self.max_loaded_loras}, max_loras_per_batch={self.max_loras_per_batch}" + f"max_loaded_loras={cfg.max_loaded_loras}, max_loras_per_batch={cfg.max_loras_per_batch}" ) - assert len(self.lora_paths) <= self.max_loaded_loras, ( + assert len(cfg.lora_paths) <= cfg.max_loaded_loras, ( "The number of LoRA paths should not exceed max_loaded_loras. " - f"max_loaded_loras={self.max_loaded_loras}, lora_paths={len(self.lora_paths)}" + f"max_loaded_loras={cfg.max_loaded_loras}, lora_paths={len(cfg.lora_paths)}" ) - if self.max_lora_chunk_size is not None: + if cfg.max_lora_chunk_size is not None: assert ( - 16 <= self.max_lora_chunk_size <= 128 - and (self.max_lora_chunk_size & (self.max_lora_chunk_size - 1)) == 0 + 16 <= cfg.max_lora_chunk_size <= 128 + and (cfg.max_lora_chunk_size & (cfg.max_lora_chunk_size - 1)) == 0 ), "--max-lora-chunk-size must be a power of 2 between 16 and 128." - if self.lora_use_virtual_experts: + if cfg.lora_use_virtual_experts: logger.info("Virtual expert computation enabled.") assert ( - self.lora_drain_wait_threshold >= 0.0 + cfg.lora_drain_wait_threshold >= 0.0 ), "--lora-drain-wait-threshold must be non-negative." def _check_lora_speculative_compatibility(self): @@ -10519,8 +10531,9 @@ def describe_kv_events_publisher(self) -> Optional[dict]: # disaggregation / msgspec / zmq at module top level. from sglang.srt.disaggregation.kv_events import KVEventsConfig - raw = self.kv_events_config - page_size = self.page_size + resolved = resolving_view(self) + raw = resolved.kv_events_config + page_size = resolved.page_size if not raw or page_size is None or page_size <= 0: return None try: @@ -10560,10 +10573,14 @@ def should_report_expert_balancedness(self) -> bool: return cfg.expert_balancedness_report_mode != "off" def should_log_expert_balancedness_to_server_log(self) -> bool: - return self.expert_balancedness_report_mode in ("server_log", "both") + cfg = resolving_view(self) + + return cfg.expert_balancedness_report_mode in ("server_log", "both") def should_export_expert_balancedness_to_prometheus(self) -> bool: - return self.expert_balancedness_report_mode in ("prometheus", "both") + cfg = resolving_view(self) + + return cfg.expert_balancedness_report_mode in ("prometheus", "both") def compute_world_size(server_args: ServerArgs) -> int: From 3c38035dbb6cf3b0a6ce86f336cc49c61882e4d3 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 12:05:22 +0000 Subject: [PATCH 11/26] config: the last resolution-reachable readers move to the view A census over the modules the pipeline actually reaches -- the registered passes and providers plus the import map, the same derivation `test_resolution_reads_no_bag` uses -- left 26 field reads outside `arg_groups` and `ServerArgs`: the dLLM config builder, the experimental Marlin LoRA validator, and the NPU platform defaults. The NPU one is the reason to bother. `set_default_server_args` asks "did anyone decide `page_size` yet?" before declaring its own default, and that question has to be asked of the declarations: reading the field would answer "no" for a size an earlier pass had already declared, and the hook would overwrite it. No A/B on a CUDA host can catch that, which is why the census is the check here. `configure_logger`'s single read stays: `log_level` is raw input, and that function is called with stand-ins. --- python/sglang/srt/dllm/config.py | 21 ++++++++++--------- .../sglang/srt/hardware_backend/npu/utils.py | 17 ++++++++------- .../srt/lora/marlin_lora_temp/policy.py | 21 +++++++++++-------- 3 files changed, 33 insertions(+), 26 deletions(-) diff --git a/python/sglang/srt/dllm/config.py b/python/sglang/srt/dllm/config.py index f0f2b9d11087..a2206feb92a9 100644 --- a/python/sglang/srt/dllm/config.py +++ b/python/sglang/srt/dllm/config.py @@ -25,13 +25,16 @@ def __init__( def from_server_args( server_args: ServerArgs, ): - if server_args.dllm_algorithm is None: + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) + if cfg.dllm_algorithm is None: return None model_config = ModelConfig.from_server_args( server_args, - model_path=server_args.model_path, - model_revision=server_args.revision, + model_path=cfg.model_path, + model_revision=cfg.revision, ) DLLM_PARAMS = { "LLaDA2MoeModelLM": {"block_size": 32, "mask_id": 156895}, @@ -48,13 +51,11 @@ def from_server_args( raise RuntimeError(f"Unknown diffusion LLM: {arch}") max_running_requests = ( - 1 - if server_args.max_running_requests is None - else server_args.max_running_requests + 1 if cfg.max_running_requests is None else cfg.max_running_requests ) algorithm_config = {} - if server_args.dllm_algorithm_config is not None: + if cfg.dllm_algorithm_config is not None: try: import yaml except ImportError: @@ -62,17 +63,17 @@ def from_server_args( "Please install PyYAML to use YAML config files. " "`pip install pyyaml`" ) - with open(server_args.dllm_algorithm_config, "r") as f: + with open(cfg.dllm_algorithm_config, "r") as f: algorithm_config = yaml.safe_load(f) # Parse common algorithm configurations block_size = algorithm_config.get("block_size", block_size) return DllmConfig( - algorithm=server_args.dllm_algorithm, + algorithm=cfg.dllm_algorithm, algorithm_config=algorithm_config, block_size=block_size, mask_id=mask_id, max_running_requests=max_running_requests, - first_done_first_out_mode=server_args.dllm_fdfo, + first_done_first_out_mode=cfg.dllm_fdfo, ) diff --git a/python/sglang/srt/hardware_backend/npu/utils.py b/python/sglang/srt/hardware_backend/npu/utils.py index 3a285f728ab4..857ae86341de 100644 --- a/python/sglang/srt/hardware_backend/npu/utils.py +++ b/python/sglang/srt/hardware_backend/npu/utils.py @@ -43,6 +43,9 @@ def set_default_server_args(args: "ServerArgs"): """ Set default server arguments for NPU backend. """ + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(args) # NPU only works with "ascend" attention backend for now declare_resolution( @@ -60,7 +63,7 @@ def set_default_server_args(args: "ServerArgs"): "set_default_server_args", decode_attention_backend="ascend", ) - if args.page_size is None: + if cfg.page_size is None: declare_resolution( args, "set_default_server_args", @@ -68,33 +71,33 @@ def set_default_server_args(args: "ServerArgs"): ) # NPU memory settings - decode = args.cuda_graph_config.decode + decode = cfg.cuda_graph_config.decode npu_mem = get_npu_memory_capacity() if npu_mem <= 32 * 1024: # Ascend 910B4,910B4_1 # (chunked_prefill_size 4k, max_bs 16 if tp < 4 else 64) - if args.chunked_prefill_size is None: + if cfg.chunked_prefill_size is None: declare_resolution( args, "set_default_server_args", chunked_prefill_size=4 * 1024, ) if decode.max_bs is None: - if args.tp_size < 4: + if cfg.tp_size < 4: decode.max_bs = 16 else: decode.max_bs = 64 elif npu_mem <= 64 * 1024: # Ascend 910B1,910B2,910B2C,910B3,910_9391,910_9392,910_9381,910_9382,910_9372,910_9362 # (chunked_prefill_size 8k, max_bs 64 if tp < 4 else 256) - if args.chunked_prefill_size is None: + if cfg.chunked_prefill_size is None: declare_resolution( args, "set_default_server_args", chunked_prefill_size=8 * 1024, ) if decode.max_bs is None: - if args.tp_size < 4: + if cfg.tp_size < 4: decode.max_bs = 64 else: decode.max_bs = 256 @@ -107,7 +110,7 @@ def set_default_server_args(args: "ServerArgs"): ) # handles hierarchical cache configs - if args.enable_hierarchical_cache: + if cfg.enable_hierarchical_cache: declare_resolution( args, "set_default_server_args", diff --git a/python/sglang/srt/lora/marlin_lora_temp/policy.py b/python/sglang/srt/lora/marlin_lora_temp/policy.py index 76259ead0010..c05838fa1365 100644 --- a/python/sglang/srt/lora/marlin_lora_temp/policy.py +++ b/python/sglang/srt/lora/marlin_lora_temp/policy.py @@ -14,6 +14,9 @@ def validate_experimental_sgl_marlin_server_args( server_args: Any, resolved_args: Any ) -> None: """Validate startup options before the experimental runner is constructed.""" + from sglang.srt.arg_groups.overrides import resolving_view + + cfg = resolving_view(server_args) if resolved_args.ep_size > 1 and resolved_args.moe_a2a_backend != "none": raise ValueError("experimental_sgl_marlin EP requires --moe-a2a-backend none") @@ -21,16 +24,16 @@ def validate_experimental_sgl_marlin_server_args( # A provided adapter path implicitly enables LoRA later unless it was # explicitly disabled. No-LoRA delegates to the stock Marlin fused path. lora_enabled = bool(resolved_args.enable_lora) or ( - resolved_args.enable_lora is None and bool(server_args.lora_paths) + resolved_args.enable_lora is None and bool(cfg.lora_paths) ) if not lora_enabled: return - if not server_args.lora_use_virtual_experts: + if not cfg.lora_use_virtual_experts: raise ValueError( "experimental_sgl_marlin LoRA requires --lora-use-virtual-experts" ) - if server_args.lora_backend != "triton": + if cfg.lora_backend != "triton": # The temporary dense/sink kernels consume Triton SGEMM batch metadata # directly; other global backends are not adapted in this tree. raise ValueError("experimental_sgl_marlin LoRA requires --lora-backend triton") @@ -38,12 +41,12 @@ def validate_experimental_sgl_marlin_server_args( return if ( - server_args.init_expert_location != "trivial" - or server_args.ep_num_redundant_experts != 0 - or server_args.enable_eplb - or server_args.elastic_ep_backend is not None - or server_args.enable_elastic_expert_backup - or server_args.elastic_ep_rejoin + cfg.init_expert_location != "trivial" + or cfg.ep_num_redundant_experts != 0 + or cfg.enable_eplb + or cfg.elastic_ep_backend is not None + or cfg.enable_elastic_expert_backup + or cfg.elastic_ep_rejoin ): raise ValueError( "experimental_sgl_marlin EP requires trivial expert placement " From 931dc630a2e1edffef11f2d1014199df3edb9a98 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 12:26:06 +0000 Subject: [PATCH 12/26] config: the dispatcher and seven test files read the resolution result Two reads in the dispatcher itself, one of which matters: `get_device_memory_capacity(self.device)` runs right after the platform defaults declare `device`, so a field read there would size memory for `auto`. The dummy-model boundary check moves with it for uniformity. The tests follow the same rule as the golden model-override ones: what they assert is what resolution decided, so they read `resolution_result` rather than the field -- the CPU-EAGLE overlap constraint, the dSpark draft-path default, the media-domain normalization, the multimodal piecewise-graph gates, the encoder transfer backend, and the spec-registry algorithm name. The multimodal processor fixture seeds the worker counts through `override_server_args` instead of a MagicMock, because the processor reads them from `get_mm()` now. All no-ops today (the declaration still writes the field as it declares); they are what the flip needs in place first. --- python/sglang/srt/server_args.py | 9 ++++-- .../dspark/test_dspark_draft_path_default.py | 12 +++++-- .../test_multimodal_piecewise_cuda_graph.py | 32 ++++++++++++++----- .../test_kimi_k3_encoder_mode.py | 3 +- .../unit/managers/test_mm_process_config.py | 12 ++++--- .../spec/test_spec_cpu_overlap_constraint.py | 11 ++++--- .../unit/spec/test_spec_registry.py | 5 ++- .../unit/test_server_args_migration.py | 6 +++- 8 files changed, 64 insertions(+), 26 deletions(-) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 59623be30b6e..d949367e9c5b 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -3840,6 +3840,10 @@ def _run_resolution_pipeline(self): # _handle_model_specific_adjustments never runs. self._resolved_overrides = [] + # Read through the declarations from here on: the handlers below declare + # rather than assign, so a field read would answer with the raw input. + cfg = resolving_view(self) + from sglang.srt.arg_groups.mega_moe_hook import handle_mega_moe handle_mega_moe(self) @@ -3851,7 +3855,7 @@ def _run_resolution_pipeline(self): # Reject an explicitly enabled but incompatible hardware runtime before # model path resolution, downloads, or the dummy-model short circuit. self._handle_hardware_runtime_validation() - if self.model_path.lower() in ["none", "dummy"]: + if cfg.model_path.lower() in ["none", "dummy"]: return self._handle_model_source_paths() @@ -3912,8 +3916,7 @@ def _run_resolution_pipeline(self): current_platform.apply_server_args_defaults, ) - # Get GPU memory capacity, which is a common dependency for several configuration steps. - gpu_mem = get_device_memory_capacity(self.device) + gpu_mem = get_device_memory_capacity(cfg.device) # Handle memory-related, chunked prefill, and CUDA graph batch size configurations. self._handle_gpu_memory_settings(gpu_mem) diff --git a/test/registered/spec/dspark/test_dspark_draft_path_default.py b/test/registered/spec/dspark/test_dspark_draft_path_default.py index 53fbc1051bbc..12b4c3621536 100644 --- a/test/registered/spec/dspark/test_dspark_draft_path_default.py +++ b/test/registered/spec/dspark/test_dspark_draft_path_default.py @@ -1,6 +1,7 @@ import unittest from types import SimpleNamespace +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.arg_groups.speculative_hook import ( _handle_dspark, _target_checkpoint_bundles_dspark_draft, @@ -63,8 +64,13 @@ def test_bundled_checkpoint_defaults_draft_path_to_model_path(self): model_path=_BUNDLED_MODEL_PATH, hf_config=_bundled_hf_config() ) _handle_dspark(server_args) - self.assertEqual(server_args.speculative_draft_model_path, _BUNDLED_MODEL_PATH) - self.assertEqual(server_args.speculative_num_draft_tokens, 6) + self.assertEqual( + resolution_result(server_args, "speculative_draft_model_path"), + _BUNDLED_MODEL_PATH, + ) + self.assertEqual( + resolution_result(server_args, "speculative_num_draft_tokens"), 6 + ) def test_plain_target_without_draft_path_raises(self): server_args = _make_dspark_server_args( @@ -80,7 +86,7 @@ def test_explicit_draft_path_is_not_overwritten(self): server_args.speculative_draft_model_path = "deepseek-ai/some-other-dspark-draft" _handle_dspark(server_args) self.assertEqual( - server_args.speculative_draft_model_path, + resolution_result(server_args, "speculative_draft_model_path"), "deepseek-ai/some-other-dspark-draft", ) diff --git a/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py b/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py index b81df8b7b1f0..b16d22a0be03 100644 --- a/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py +++ b/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py @@ -4,6 +4,7 @@ from types import SimpleNamespace from unittest.mock import patch +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.configs.embedding_model_spec import resolve_embedding_model_spec from sglang.srt.configs.model_config import ( is_multimodal_piecewise_cuda_graph_supported, @@ -90,7 +91,10 @@ def test_supported_multimodal_model_upgrades_default_to_tc_piecewise(self): ): args._apply_cuda_graph_compatibility() - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.TC_PIECEWISE, + ) disable_if_incompatible.assert_called_once() def test_trtllm_mla_stays_on_breakable(self): @@ -118,7 +122,10 @@ def test_trtllm_mla_stays_on_breakable(self): ): args._apply_cuda_graph_compatibility() - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.BREAKABLE, + ) def test_explicit_tc_piecewise_overrides_trtllm_mla_default(self): args = ServerArgs(model_path="dummy") @@ -134,7 +141,10 @@ def test_explicit_tc_piecewise_overrides_trtllm_mla_default(self): ): args._apply_cuda_graph_compatibility() - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.TC_PIECEWISE, + ) def test_multimodal_inputs_keep_tc_piecewise_prefill_enabled(self): runner = self._make_prefill_runner(Backend.TC_PIECEWISE) @@ -178,10 +188,16 @@ def test_embedding_gemma_forces_breakable_prefill(self): ): args._handle_model_capability_adjustments() - self.assertTrue(args.disable_radix_cache) - self.assertEqual(args.chunked_prefill_size, -1) - self.assertEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED) - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE) + self.assertTrue(resolution_result(args, "disable_radix_cache")) + self.assertEqual(resolution_result(args, "chunked_prefill_size"), -1) + self.assertEqual( + resolution_result(args, "cuda_graph_config").decode.backend, + Backend.DISABLED, + ) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.BREAKABLE, + ) def test_encoder_embedding_model_enables_embedding_mode_without_flag(self): args = ServerArgs(model_path="dummy") @@ -199,7 +215,7 @@ def test_encoder_embedding_model_enables_embedding_mode_without_flag(self): with patch.object(args, "get_model_config", return_value=args.model_config): args._handle_model_capability_adjustments() - self.assertTrue(args.is_embedding) + self.assertTrue(resolution_result(args, "is_embedding")) if __name__ == "__main__": 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 3e9c3dc79b9e..dd1a19ebb733 100644 --- a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py +++ b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py @@ -15,6 +15,7 @@ from fastapi import HTTPException from PIL import Image +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.disaggregation.encoder.preprocessor import ( EncoderPreprocessor, EncoderPreprocessResult, @@ -195,7 +196,7 @@ def env_field_flags(): finally: shutil.rmtree(config_dir, ignore_errors=True) - assert resolved.encoder_transfer_backend == "zmq_to_tokenizer" + assert resolution_result(resolved, "encoder_transfer_backend") == "zmq_to_tokenizer" # 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") diff --git a/test/registered/unit/managers/test_mm_process_config.py b/test/registered/unit/managers/test_mm_process_config.py index a5e8adc30970..ad1791d767b3 100644 --- a/test/registered/unit/managers/test_mm_process_config.py +++ b/test/registered/unit/managers/test_mm_process_config.py @@ -6,6 +6,7 @@ import torch +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.environ import envs from sglang.srt.runtime_context import get_context from sglang.srt.server_args import ServerArgs @@ -25,17 +26,20 @@ def _validate_config(self, mm_process_config): def test_valid_config_accepted(self): args = self._validate_config({"image": {"max_pixels": 5000000}}) - self.assertEqual(args.mm_process_config, {"image": {"max_pixels": 5000000}}) + self.assertEqual( + resolution_result(args, "mm_process_config"), + {"image": {"max_pixels": 5000000}}, + ) def test_empty_config_accepted(self): args = self._validate_config({}) - self.assertEqual(args.mm_process_config, {}) + self.assertEqual(resolution_result(args, "mm_process_config"), {}) def test_none_config_defaults_to_empty_dict(self): args = self._validate_config(None) # None is kept as-is for dummy models (default happens after early return) # but for real models it would be set to {} - self.assertIsNone(args.mm_process_config) + self.assertIsNone(resolution_result(args, "mm_process_config")) def test_top_level_non_dict_rejected(self): with self.assertRaises(TypeError) as ctx: @@ -64,7 +68,7 @@ def test_multi_modality_config_accepted(self): "audio": {"sample_rate": 16000}, } args = self._validate_config(config) - self.assertEqual(args.mm_process_config, config) + self.assertEqual(resolution_result(args, "mm_process_config"), config) class TestBaseProcessorConfigExtraction(CustomTestCase): diff --git a/test/registered/unit/spec/test_spec_cpu_overlap_constraint.py b/test/registered/unit/spec/test_spec_cpu_overlap_constraint.py index 6e21a69bee3e..68ef5eac9c9a 100644 --- a/test/registered/unit/spec/test_spec_cpu_overlap_constraint.py +++ b/test/registered/unit/spec/test_spec_cpu_overlap_constraint.py @@ -1,6 +1,7 @@ import unittest from types import SimpleNamespace +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci @@ -33,18 +34,18 @@ def _make_spec_args(device: str, algorithm: str = "EAGLE", **overrides) -> Serve class TestSpecCPUOverlapConstraint(CustomTestCase): def test_cpu_eagle_forces_disable_overlap_schedule(self): args = _make_spec_args(device="cpu") - self.assertFalse(args.disable_overlap_schedule) + self.assertFalse(resolution_result(args, "disable_overlap_schedule")) handle_speculative_decoding(args) - self.assertTrue(args.disable_overlap_schedule) + self.assertTrue(resolution_result(args, "disable_overlap_schedule")) def test_cpu_eagle3_forces_disable_overlap_schedule(self): args = _make_spec_args(device="cpu", algorithm="EAGLE3") handle_speculative_decoding(args) - self.assertTrue(args.disable_overlap_schedule) + self.assertTrue(resolution_result(args, "disable_overlap_schedule")) def test_cpu_explicit_disable_overlap_is_preserved(self): args = _make_spec_args(device="cpu", disable_overlap_schedule=True) @@ -56,7 +57,7 @@ def test_cpu_explicit_disable_overlap_is_preserved(self): ) as logs: handle_speculative_decoding(args) - self.assertTrue(args.disable_overlap_schedule) + self.assertTrue(resolution_result(args, "disable_overlap_schedule")) self.assertFalse( any("Overlap schedule" in message for message in logs.output), f"hook warned about overriding an already-disabled overlap: {logs.output}", @@ -68,7 +69,7 @@ def test_cuda_eagle_keeps_overlap_schedule(self): handle_speculative_decoding(args) - self.assertFalse(args.disable_overlap_schedule) + self.assertFalse(resolution_result(args, "disable_overlap_schedule")) if __name__ == "__main__": diff --git a/test/registered/unit/spec/test_spec_registry.py b/test/registered/unit/spec/test_spec_registry.py index 141b854c4017..047a8b7174b0 100644 --- a/test/registered/unit/spec/test_spec_registry.py +++ b/test/registered/unit/spec/test_spec_registry.py @@ -4,6 +4,7 @@ from types import SimpleNamespace from unittest.mock import MagicMock +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_registry import ( @@ -237,7 +238,9 @@ def _factory(server_args): handle_speculative_decoding(server_args) - self.assertEqual(server_args.speculative_algorithm, "MY_HANDLE_ARGS") + self.assertEqual( + resolution_result(server_args, "speculative_algorithm"), "MY_HANDLE_ARGS" + ) self.assertEqual(server_args.custom_spec_handle_seen, "MY_HANDLE_ARGS") self.assertEqual(server_args.speculative_num_draft_tokens, 7) diff --git a/test/registered/unit/test_server_args_migration.py b/test/registered/unit/test_server_args_migration.py index cac5d7cc6340..1ea89e8c31b4 100644 --- a/test/registered/unit/test_server_args_migration.py +++ b/test/registered/unit/test_server_args_migration.py @@ -7,6 +7,7 @@ import argparse import unittest +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.server_args import ServerArgs from sglang.srt.utils.common import configure_media_url_security from sglang.test.ci.ci_register import register_cpu_ci @@ -95,8 +96,11 @@ def test_media_url_security_args(self): "32", ] ) + # The normalization is a declaration, so it lands in the + # resolution result rather than on the field. self.assertEqual( - sa.allowed_media_domains, ["127.0.0.1", "media.example.com"] + resolution_result(sa, "allowed_media_domains"), + ["127.0.0.1", "media.example.com"], ) self.assertEqual(sa.media_url_max_file_size_mb, 32) finally: From 936fa53445664a0f00e9345ab7cb0c18dda96218 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 13:02:21 +0000 Subject: [PATCH 13/26] config: the assertions read the resolution result, not the record Everything that will change hands when the declarations stop writing the field, moved ahead of the flip so it can be checked while both still agree. `test_server_args` is the bulk of it (101 reads): what those cases assert is what resolution decided, and `resolution_result` answers that whether or not the declaration was written back. Four assertions go the *other* way and now read the field on purpose -- the FA4 page-size and waterfill cases exist to show the field staying pristine while the declaration wins, so they keep reading it and say so. `_comparable` in the reproducibility suite reads the projection too: comparing fields would have stopped covering the decisions a resolution leak would shift. The multimodal and Kimi processor fixtures seed their worker counts and cache budget through `override_server_args` instead of a stand-in, because the processor reads them from `get_mm()` / `get_serving()` now -- that also fixes `test_kimi_processor_workers_clone_the_gpu_wrapper`, which the processor conversion broke (the cache came out enabled and the fingerprint path ran into a `SimpleNamespace` hf_config). The supplied-instance exposure pin drops 19 entries: the model-config, dLLM, CP, Marlin-LoRA, adaptive-spec and spec-registry reads all go through the view now. --- python/sglang/srt/arg_groups/overrides.py | 27 +- python/sglang/srt/layers/cp/base.py | 4 +- .../sglang/srt/parser/template_detection.py | 20 +- python/sglang/srt/server_args.py | 60 ++-- .../test_resolution_declarations.py | 2 +- .../test_resolution_is_reproducible.py | 19 +- .../unit/server_args/test_server_args.py | 332 ++++++++++++------ test/registered/unit/test_model_overrides.py | 10 +- .../unit/test_runtime_context_override.py | 10 +- .../unit/test_server_args_migration.py | 3 +- ...test_supplied_instance_exposure_ratchet.py | 34 +- 11 files changed, 318 insertions(+), 203 deletions(-) diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index a52b66995ccb..556479e54850 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -648,7 +648,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: # lacks that DCP path (TypeError: unexpected kwarg 'causal_seqs'). overrides["speculative_attention_mode"] = "decode" - prefill_backend, decode_backend = attention_backends_of(server_args) + prefill_backend, decode_backend = attention_backends_of(cfg) if decode_backend == "cutedsl_mla" or decode_backend is None: _require_kimi_k3_cutedsl_dcp_support() logger.info( @@ -725,7 +725,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: # Explicit backend knobs keep priority, but the mode is a separate knob # that still needs declaring -- else verify stays on the prefill backend, # whose host-side plan (flashinfer by default) forces a per-step D2H. - _, backend = attention_backends_of(server_args) + _, backend = attention_backends_of(cfg) if _dspark_verify_on_decode_backend(backend, q_len, cfg.kv_cache_dtype): overrides["speculative_attention_mode"] = "decode" logger.info( @@ -1287,7 +1287,8 @@ def _moss_vl_overrides(server_args: Any, hf_config: Any) -> dict: @_register_for("MiniCPMForCausalLM", "MiniCPMSALAForCausalLM") def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict: - if server_args.enable_dp_attention: + cfg = resolving_view(server_args) + if cfg.enable_dp_attention: raise ValueError("MiniCPM does not support DP attention") has_sparse_attention = getattr(hf_config, "has_minicpm_sparse_attention", False) has_hybrid_attention = has_sparse_attention or getattr( @@ -1295,7 +1296,7 @@ def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict: ) overrides: Dict[str, Any] = {} if has_hybrid_attention: - if server_args.enable_hierarchical_cache: + if cfg.enable_hierarchical_cache: raise ValueError("MiniCPM SALA does not support hierarchical cache") overrides["disable_radix_cache"] = True if envs.SGLANG_MINICPM_FORCE_DENSE.get(): @@ -1305,29 +1306,29 @@ def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict: } # Literal keys keep the written-field set statically derivable; a loop # variable hides it from the census in test_chain_read_ratchet.py. - dense_attention = dense_backends.get(server_args.attention_backend) + dense_attention = dense_backends.get(cfg.attention_backend) if dense_attention is not None: overrides["attention_backend"] = dense_attention - dense_prefill = dense_backends.get(server_args.prefill_attention_backend) + dense_prefill = dense_backends.get(cfg.prefill_attention_backend) if dense_prefill is not None: overrides["prefill_attention_backend"] = dense_prefill - dense_decode = dense_backends.get(server_args.decode_attention_backend) + dense_decode = dense_backends.get(cfg.decode_attention_backend) if dense_decode is not None: overrides["decode_attention_backend"] = dense_decode elif has_sparse_attention: - uses_sparse_backend = server_args.is_attention_backend_not_set() or any( + uses_sparse_backend = cfg.is_attention_backend_not_set() or any( backend in ("minicpm_flashattn", "minicpm_flashinfer") for backend in ( - server_args.attention_backend, - server_args.prefill_attention_backend, - server_args.decode_attention_backend, + cfg.attention_backend, + cfg.prefill_attention_backend, + cfg.decode_attention_backend, ) ) - if uses_sparse_backend and server_args.disaggregation_mode != "null": + if uses_sparse_backend and cfg.disaggregation_mode != "null": raise ValueError( "MiniCPM sparse attention does not support PD disaggregation" ) - if server_args.is_attention_backend_not_set(): + if cfg.is_attention_backend_not_set(): overrides["attention_backend"] = ( "minicpm_flashinfer" if is_blackwell_supported() diff --git a/python/sglang/srt/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index b58cc68bf620..c57e33e97c8d 100644 --- a/python/sglang/srt/layers/cp/base.py +++ b/python/sglang/srt/layers/cp/base.py @@ -244,11 +244,11 @@ def init_cp_strategy(server_args: ServerArgs) -> None: cfg = resolving_view(server_args) global _STRATEGY - if not getattr(server_args, "enable_prefill_cp", False): + if not cfg.enable_prefill_cp: _STRATEGY = None return - cp_size = getattr(server_args, "attn_cp_size", 1) + cp_size = cfg.attn_cp_size if cp_size <= 1: _STRATEGY = None return diff --git a/python/sglang/srt/parser/template_detection.py b/python/sglang/srt/parser/template_detection.py index d353aa1c60ca..13d5465a7eb2 100644 --- a/python/sglang/srt/parser/template_detection.py +++ b/python/sglang/srt/parser/template_detection.py @@ -29,7 +29,7 @@ import jinja2.nodes import jinja2.sandbox -from sglang.srt.arg_groups.overrides import declare_late_resolution +from sglang.srt.arg_groups.overrides import declare_late_resolution, resolving_view logger = logging.getLogger(__name__) @@ -666,11 +666,12 @@ def _architecture_auto_parsers(server_args, needs: Tuple[str, ...]) -> Dict[str, """The parsers the model architecture implies, for the fields still on auto.""" from sglang.srt.utils.hf_transformers_utils import get_config + cfg = resolving_view(server_args) config = get_config( - server_args.model_path, - trust_remote_code=server_args.trust_remote_code, - revision=getattr(server_args, "revision", None), - model_config_parser=getattr(server_args, "model_config_parser", "auto"), + cfg.model_path, + trust_remote_code=cfg.trust_remote_code, + revision=getattr(cfg, "revision", None), + model_config_parser=getattr(cfg, "model_config_parser", "auto"), ) architectures = getattr(config, "architectures", None) or [] arch = architectures[0] if architectures else "" @@ -708,17 +709,18 @@ def resolve_auto_parsers(server_args) -> None: the schedulers it forks, the HTTP server, and the tokenizer workers it is serialized for. """ + cfg = resolving_view(server_args) needs = tuple( attr for attr in ("reasoning_parser", "tool_call_parser") - if getattr(server_args, attr) == "auto" + if getattr(cfg, attr) == "auto" ) if not needs: return from sglang.srt.utils.hf_transformers_utils import get_tokenizer - chat_template_arg = getattr(server_args, "chat_template", None) + chat_template_arg = getattr(cfg, "chat_template", None) try: explicit_jinja_template = _load_explicit_jinja_template(chat_template_arg) except Exception as e: @@ -731,8 +733,8 @@ def resolve_auto_parsers(server_args) -> None: tokenizer = None try: tokenizer = get_tokenizer( - server_args.model_path, - trust_remote_code=server_args.trust_remote_code, + cfg.model_path, + trust_remote_code=cfg.trust_remote_code, ) except Exception as e: logger.warning(f"Failed to load tokenizer for auto-detection: {e}") diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index d949367e9c5b..17a4afa4942f 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -3840,8 +3840,6 @@ def _run_resolution_pipeline(self): # _handle_model_specific_adjustments never runs. self._resolved_overrides = [] - # Read through the declarations from here on: the handlers below declare - # rather than assign, so a field read would answer with the raw input. cfg = resolving_view(self) from sglang.srt.arg_groups.mega_moe_hook import handle_mega_moe @@ -5488,19 +5486,20 @@ def post_capture_kv_sizing_planned(self) -> bool: def pre_capture_activation_reserve_mb(self, gpu_mem: Optional[float]) -> float: # Runtime activation working-set reserve for eager decode above the captured # max_bs and transient prefill/logits; also covers fixed state caches. - if self.disaggregation_mode == "decode": + cfg = resolving_view(self) + if cfg.disaggregation_mode == "decode": running_requests = ( - self.max_running_requests or self.cuda_graph_config.decode.max_bs or 1 + cfg.max_running_requests or cfg.cuda_graph_config.decode.max_bs or 1 ) activation_tokens = max( - running_requests * (self.speculative_num_draft_tokens or 1), 2048 + running_requests * (cfg.speculative_num_draft_tokens or 1), 2048 ) - elif self.chunked_prefill_size > 0: - activation_tokens = max(self.chunked_prefill_size, 2048) + elif cfg.chunked_prefill_size > 0: + activation_tokens = max(cfg.chunked_prefill_size, 2048) else: - activation_tokens = max(self.max_prefill_tokens, 2048) + activation_tokens = max(cfg.max_prefill_tokens, 2048) reserved_mem = ( - 512 + activation_tokens * 1.5 + self.tp_size * self.pp_size / 8 * 1024 + 512 + activation_tokens * 1.5 + cfg.tp_size * cfg.pp_size / 8 * 1024 ) if gpu_mem is not None and gpu_mem > 60 * 1024: reserved_mem = max(reserved_mem, 10 * 1024) @@ -6303,6 +6302,7 @@ def _get_default_attn_backend(self, use_mla_backend: bool, model_config): 2.2 We will use Flashinfer backend on blackwell. 2.3 Otherwise, we will use triton backend. """ + cfg = resolving_view(self) # OOT platforms provide their own default attention backend. if current_platform.is_out_of_tree(): return current_platform.get_default_attention_backend() @@ -6327,8 +6327,8 @@ def _get_default_attn_backend(self, use_mla_backend: bool, model_config): is_sm100_supported() and is_no_spec_infer_or_topk_one(resolved_view(self)) and ( - self.speculative_algorithm is None - or self.speculative_eagle_topk is not None + cfg.speculative_algorithm is None + or cfg.speculative_eagle_topk is not None ) ): # trtllm_mha requires equal K/V row widths; fa4 carries @@ -8305,7 +8305,7 @@ def _handle_load_format(self): "(--weight-cache-mode off) for this configuration." ) - if self.weight_cache_mode != "off" and self.enable_eplb: + if cfg.weight_cache_mode != "off" and cfg.enable_eplb: raise ValueError( "--weight-cache-mode is not supported together with --enable-eplb." ) @@ -9837,17 +9837,18 @@ def use_mla_backend(self): return model_config.attention_arch == AttentionArch.MLA def is_attention_backend_not_set(self): + cfg = resolving_view(self) return ( - self.attention_backend is None - and self.prefill_attention_backend is None - and self.decode_attention_backend is None + cfg.attention_backend is None + and cfg.prefill_attention_backend is None + and cfg.decode_attention_backend is None ) def enable_mamba_extra_buffer(self) -> bool: - return mamba_extra_buffer_of(self) + return mamba_extra_buffer_of(resolving_view(self)) def enable_mamba_extra_buffer_lazy(self) -> bool: - return mamba_extra_buffer_lazy_of(self) + return mamba_extra_buffer_lazy_of(resolving_view(self)) @property def max_speculative_num_draft_tokens(self) -> Optional[int]: @@ -10285,20 +10286,21 @@ def _check_lora_speculative_compatibility(self): Adapters apply to the target only; a shared draft runs unadapted. Matches resolved algorithm names (NEXTN has collapsed to EAGLE). """ - if self.speculative_algorithm in ["NGRAM", None]: + cfg = resolving_view(self) + if cfg.speculative_algorithm in ["NGRAM", None]: return - if self.speculative_algorithm not in _LORA_SPEC_ALGORITHMS: + if cfg.speculative_algorithm not in _LORA_SPEC_ALGORITHMS: promoted = ( " (NEXTN/EAGLE with a Gemma4 assistant draft is automatically " "promoted to FROZEN_KV_MTP, which does not support LoRA)" - if self.speculative_algorithm == "FROZEN_KV_MTP" + if cfg.speculative_algorithm == "FROZEN_KV_MTP" else "" ) raise ValueError( "LoRA is only compatible with NGRAM, EAGLE, NEXTN, EAGLE3, " "DFLASH, or DSPARK speculative decoding, not " - f"{self.speculative_algorithm}{promoted}." + f"{cfg.speculative_algorithm}{promoted}." ) ragged_mode = envs.SGLANG_RAGGED_VERIFY_MODE.get() @@ -10307,20 +10309,20 @@ def _check_lora_speculative_compatibility(self): # prefix so the message names the combination, not just the flag. unsupported = [ ( - self.speculative_algorithm == "DSPARK" and ragged_mode != "static", + cfg.speculative_algorithm == "DSPARK" and ragged_mode != "static", f"does not support SGLANG_RAGGED_VERIFY_MODE={ragged_mode!r}: " "the per-request verify lengths it schedules break the " "uniform-width LoRA segment layout", ), ( - self.speculative_adaptive, + cfg.speculative_adaptive, "does not support --speculative-adaptive: the draft is built " "from a static ServerArgs snapshot, and the runtime-state " "swap does not rebuild LoRA cuda-graph metadata", ), ( "experimental_sgl_trtllm" - in (self.moe_runner_backend, self.speculative_moe_runner_backend), + in (cfg.moe_runner_backend, cfg.speculative_moe_runner_backend), "does not support the experimental_sgl_trtllm MoE runner: its " "TopK reads the LoRA config per forward, which the draft " "resolves against the target's after its own publish ended", @@ -10471,14 +10473,15 @@ def modelexpress_transport(self) -> str: def remote_instance_weight_loader_use_transfer_engine(self, load_format=None): """``load_format`` overrides the seed's: a draft runner loading under ``--speculative-draft-load-format`` needs its own transfer engine.""" - return remote_instance_transfer_engine_of(self, load_format) + return remote_instance_transfer_engine_of(resolving_view(self), load_format) @property def kv_event_block_size(self) -> int: """Width KV events are emitted at: under DCP the radix tree pages at ``page_size * dcp_size`` (``mem_cache/kv_cache_builder.py``). """ - return self.page_size * self.dcp_size + cfg = resolving_view(self) + return cfg.page_size * self.dcp_size def describe_kv_events_publisher(self) -> Optional[dict]: """Return a structured description of this server's KV-event @@ -10751,6 +10754,7 @@ def init_new( dp_rank: Optional[int] = None, worker_ports: Optional[List[int]] = None, ) -> PortArgs: + cfg = resolving_view(server_args) if server_args.nccl_port is None: nccl_port = get_free_port() else: @@ -10783,7 +10787,7 @@ def init_new( rank=int(server_args.decoupled_spec_rank), ) - if not server_args.enable_dp_attention: + if not cfg.enable_dp_attention: # Normal case, use IPC within a single node return PortArgs( tokenizer_ipc_name=f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}", @@ -10814,7 +10818,7 @@ def init_new( # every init_new call agrees, decrementing below dist_init_port on # overflow. is_rust_server = envs.SGLANG_RUST_SERVER.get() - NUM_DERIVED_PORTS = 6 if not is_rust_server else 6 + server_args.dp_size + NUM_DERIVED_PORTS = 6 if not is_rust_server else 6 + cfg.dp_size if server_args.is_ep_scale_joiner: port_base = server_args.port + ZMQ_TCP_PORT_DELTA if port_base + NUM_DERIVED_PORTS > 65535: diff --git a/test/registered/unit/server_args/test_resolution_declarations.py b/test/registered/unit/server_args/test_resolution_declarations.py index 331e17779fc8..ee37b06cd4ec 100644 --- a/test/registered/unit/server_args/test_resolution_declarations.py +++ b/test/registered/unit/server_args/test_resolution_declarations.py @@ -764,7 +764,7 @@ def test_a_nested_resolution_decision_reaches_the_bags(self): # Snapshot before publishing: the bag serves the very object the record # holds, so comparing them after the fact compares an object with # itself and passes however the projection behaves. - expected = copy.deepcopy(server_args.cuda_graph_config) + expected = copy.deepcopy(resolution_result(server_args, "cuda_graph_config")) publish(server_args, role="scheduler") published = get_exec().graph.cuda_graph_config resolved = expected diff --git a/test/registered/unit/server_args/test_resolution_is_reproducible.py b/test/registered/unit/server_args/test_resolution_is_reproducible.py index 0badd493b7a4..dfeccd4b2333 100644 --- a/test/registered/unit/server_args/test_resolution_is_reproducible.py +++ b/test/registered/unit/server_args/test_resolution_is_reproducible.py @@ -36,6 +36,7 @@ import torch import sglang +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.environ import EnvField, envs from sglang.srt.server_args import ServerArgs from sglang.srt.utils import is_cuda @@ -270,7 +271,10 @@ def _comparable(self, server_args: ServerArgs) -> dict: for field in dataclasses.fields(server_args): if field.name in _NOT_COMPARABLE: continue - value = getattr(server_args, field.name) + # The resolution result, not the field: a declaration-only resolver + # never writes the field, so comparing fields would miss exactly + # the decisions a leak would shift. + value = resolution_result(server_args, field.name) # Nested dataclasses (cuda_graph_config) compare structurally, and # everything else is deep-copied: a snapshot that stored the live # list/dict would follow an in-place mutation, which is exactly the @@ -382,10 +386,13 @@ def test_a_resolution_does_not_leak_into_the_next(self): # differs from the cpu that `default_before` resolved to. expected = ( "cuda_ipc" - if intermediate.mm_feature_transport == "cuda_ipc" + if resolution_result(intermediate, "mm_feature_transport") + == "cuda_ipc" else "cpu" ) - self.assertEqual(after.mm_feature_transport, expected) + self.assertEqual( + resolution_result(after, "mm_feature_transport"), expected + ) def test_resolving_a_sibling_leaves_the_first_alone(self): for label, config, kwargs in _SHAPES: @@ -847,9 +854,9 @@ def test_the_change_reaches_the_bags(self): self.assertEqual(get_parallel().config.dist_init_addr, "1.2.3.4:5000") self.assertEqual( get_schedule().chunked_prefill_size, - parent.chunked_prefill_size, - "publishing the copy re-ran resolution; the bag disagrees with the " - "record the parent resolved", + resolution_result(parent, "chunked_prefill_size"), + "publishing the copy re-ran resolution; the bag disagrees with what " + "the parent's resolution decided", ) def test_no_bare_replace_of_a_record_outside_the_helper(self): diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 5c3e085b5951..234d268f52eb 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -10,6 +10,7 @@ import sglang.srt.server_args as server_args_module from sglang.srt.arg_groups import pd_disaggregation_hook +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding from sglang.srt.entrypoints.sidecar import ( SGLANG_GRPC_ENDPOINT_ENV, @@ -61,7 +62,7 @@ def test_enable_w4a4_mxfp4_megamoe_sets_deepgemm_env(self): args.resolve_once() - self.assertTrue(args.enable_w4a4_mxfp4_megamoe) + self.assertTrue(resolution_result(args, "enable_w4a4_mxfp4_megamoe")) self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "1") self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "1") @@ -76,14 +77,14 @@ def test_w4a4_mxfp4_megamoe_disabled_preserves_deepgemm_env(self): # nothing to be untouched by. args.resolve_once() - self.assertFalse(args.enable_w4a4_mxfp4_megamoe) + self.assertFalse(resolution_result(args, "enable_w4a4_mxfp4_megamoe")) self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "0") self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "0") def test_prefill_decode_interval(self): args = ServerArgs(model_path="dummy", prefill_decode_interval=16) args.resolve_once() - self.assertEqual(args.prefill_decode_interval, 16) + self.assertEqual(resolution_result(args, "prefill_decode_interval"), 16) with self.assertRaisesRegex( ValueError, "--prefill-decode-interval must be non-negative" @@ -114,22 +115,24 @@ def _resolved(**kwargs): return server_args disabled = _resolved(model_path="dummy") - self.assertFalse(disabled.enable_return_hidden_states) - self.assertIsNone(disabled.return_hidden_states_mode) + self.assertFalse(resolution_result(disabled, "enable_return_hidden_states")) + self.assertIsNone(resolution_result(disabled, "return_hidden_states_mode")) last = _resolved( model_path="dummy", return_hidden_states_mode="last", ) - self.assertTrue(last.enable_return_hidden_states) - self.assertEqual(last.return_hidden_states_mode, "last") + self.assertTrue(resolution_result(last, "enable_return_hidden_states")) + self.assertEqual(resolution_result(last, "return_hidden_states_mode"), "last") legacy_full = _resolved( model_path="dummy", enable_return_hidden_states=True, ) - self.assertTrue(legacy_full.enable_return_hidden_states) - self.assertEqual(legacy_full.return_hidden_states_mode, "full") + self.assertTrue(resolution_result(legacy_full, "enable_return_hidden_states")) + self.assertEqual( + resolution_result(legacy_full, "return_hidden_states_mode"), "full" + ) parsed_last = prepare_server_args( [ @@ -140,8 +143,10 @@ def _resolved(**kwargs): ] ) parsed_last.resolve_once() - self.assertTrue(parsed_last.enable_return_hidden_states) - self.assertEqual(parsed_last.return_hidden_states_mode, "last") + self.assertTrue(resolution_result(parsed_last, "enable_return_hidden_states")) + self.assertEqual( + resolution_result(parsed_last, "return_hidden_states_mode"), "last" + ) # The rejection is resolution's, not the constructor's. with self.assertRaisesRegex( @@ -156,13 +161,24 @@ def _resolved(**kwargs): def test_draft_quantization_explicitness_survives_asdict_round_trip(self): inherited = ServerArgs(model_path="dummy", quantization="modelopt_fp4") inherited._handle_missing_default_values() - self.assertEqual(inherited.speculative_draft_model_quantization, "modelopt_fp4") - self.assertFalse(inherited._speculative_draft_quantization_explicitly_set) + self.assertEqual( + resolution_result(inherited, "speculative_draft_model_quantization"), + "modelopt_fp4", + ) + self.assertFalse( + resolution_result( + inherited, "_speculative_draft_quantization_explicitly_set" + ) + ) reconstructed = ServerArgs(**dataclasses.asdict(inherited)) reconstructed._handle_missing_default_values() - self.assertFalse(reconstructed._speculative_draft_quantization_explicitly_set) + self.assertFalse( + resolution_result( + reconstructed, "_speculative_draft_quantization_explicitly_set" + ) + ) def test_config_nested_dict_args_are_json(self): with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f: @@ -219,8 +235,10 @@ def test_new_backend_does_not_set_legacy_flag(self): server_args._handle_deprecated_args() - self.assertEqual(server_args.image_processor_backend, "pil") - self.assertFalse(server_args.disable_fast_image_processor) + self.assertEqual( + resolution_result(server_args, "image_processor_backend"), "pil" + ) + self.assertFalse(resolution_result(server_args, "disable_fast_image_processor")) def test_legacy_flag_maps_to_pil_with_one_warning(self): server_args = ServerArgs(model_path="dummy", disable_fast_image_processor=True) @@ -228,8 +246,10 @@ def test_legacy_flag_maps_to_pil_with_one_warning(self): with self.assertLogs(server_args_module.logger, level="WARNING") as logs: server_args._handle_deprecated_args() - self.assertEqual(server_args.image_processor_backend, "pil") - self.assertTrue(server_args.disable_fast_image_processor) + self.assertEqual( + resolution_result(server_args, "image_processor_backend"), "pil" + ) + self.assertTrue(resolution_result(server_args, "disable_fast_image_processor")) self.assertEqual( sum( "--disable-fast-image-processor is deprecated" in x for x in logs.output @@ -266,7 +286,9 @@ def test_cuda_ipc_is_explicit_and_bounded(self, _mock_is_cuda): with self.assertLogs(server_args_module.logger, level="INFO") as logs: server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cuda_ipc") + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cuda_ipc" + ) self.assertTrue(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) output = "\n".join(logs.output) @@ -281,8 +303,12 @@ def test_legacy_keep_flag_maps_to_cuda_ipc(self, _mock_is_cuda): with self.assertLogs(server_args_module.logger, level="WARNING") as logs: server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cuda_ipc") - self.assertFalse(server_args.keep_mm_feature_on_device) + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cuda_ipc" + ) + self.assertFalse( + resolution_result(server_args, "keep_mm_feature_on_device") + ) self.assertTrue(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) self.assertIn("deprecated", logs.output[0]) @@ -305,7 +331,9 @@ def test_explicit_cpu_overrides_legacy_environment(self, _mock_is_cuda): with self.assertLogs(server_args_module.logger, level="WARNING") as logs: server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cpu") + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cpu" + ) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) self.assertIn("overrides", logs.output[0]) @@ -316,7 +344,9 @@ def test_default_transport_is_cpu(self): with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "0"}): server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cpu") + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cpu" + ) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) @patch("sglang.srt.server_args.is_cuda", return_value=True) @@ -329,7 +359,9 @@ def test_default_transport_is_cpu_for_text_only_model(self, _mock_is_cuda): with self.assertNoLogs(server_args_module.logger, level="INFO"): server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cpu") + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cpu" + ) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) @patch("sglang.srt.server_args.is_cuda", return_value=True) @@ -342,7 +374,9 @@ def test_default_transport_is_cpu_for_multimodal_model(self, _mock_is_cuda): with self.assertNoLogs(server_args_module.logger, level="INFO"): server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cpu") + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cpu" + ) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) @patch("sglang.srt.server_args.os.path.exists", return_value=True) @@ -367,7 +401,9 @@ def test_default_transport_is_cuda_vmm_for_supported_multinode_mnnvl( with self.assertLogs(server_args_module.logger, level="INFO") as logs: server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cuda_vmm") + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cuda_vmm" + ) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) output = "\n".join(logs.output) @@ -394,7 +430,7 @@ def test_default_transport_is_cpu_for_unsupported_multinode_model( with self.assertLogs(server_args_module.logger, level="INFO") as logs: server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cpu") + self.assertEqual(resolution_result(server_args, "mm_feature_transport"), "cpu") self.assertIn("has not opted into CUDA VMM", "\n".join(logs.output)) @patch("sglang.srt.server_args.os.path.exists", return_value=False) @@ -411,7 +447,9 @@ def test_default_transport_is_cpu_without_imex_channel( with self.assertLogs(server_args_module.logger, level="INFO") as logs: server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cpu") + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cpu" + ) self.assertIn("no IMEX channel", "\n".join(logs.output)) @@ -427,7 +465,9 @@ def test_default_transport_is_cpu_for_multinode_non_mnnvl( envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear() server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cpu") + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cpu" + ) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) @patch("sglang.srt.server_args.is_cuda", return_value=True) @@ -439,7 +479,9 @@ def test_default_transport_is_cpu_for_language_only_model(self, _mock_is_cuda): envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear() server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cpu") + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cpu" + ) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) @patch("sglang.srt.server_args.is_cuda", return_value=False) @@ -474,7 +516,9 @@ def test_cuda_vmm_is_explicit_and_uses_shared_budget(self, _mock_is_cuda): with self.assertLogs(server_args_module.logger, level="INFO") as logs: server_args._handle_multimodal_feature_transport() - self.assertEqual(server_args.mm_feature_transport, "cuda_vmm") + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cuda_vmm" + ) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) output = "\n".join(logs.output) @@ -555,15 +599,22 @@ def _load_balance_args(self, **kwargs): def test_non_pd_defaults_to_round_robin(self): server_args = self._load_balance_args(disaggregation_mode="null") - self.assertEqual(server_args.load_balance_method, "round_robin") + self.assertEqual( + resolution_result(server_args, "load_balance_method"), "round_robin" + ) def test_pd_prefill_defaults_to_follow_bootstrap_room(self): server_args = self._load_balance_args(disaggregation_mode="prefill") - self.assertEqual(server_args.load_balance_method, "follow_bootstrap_room") + self.assertEqual( + resolution_result(server_args, "load_balance_method"), + "follow_bootstrap_room", + ) def test_pd_decode_defaults_to_round_robin(self): server_args = self._load_balance_args(disaggregation_mode="decode") - self.assertEqual(server_args.load_balance_method, "round_robin") + self.assertEqual( + resolution_result(server_args, "load_balance_method"), "round_robin" + ) def test_pd_prefill_dcp_warns_about_performance(self): server_args = ServerArgs( @@ -581,7 +632,7 @@ def test_pd_decode_dcp_forces_chunk_cache(self): disaggregation_transfer_backend="mooncake", dcp_size=4, ) - self.assertTrue(server_args.disable_radix_cache) + self.assertTrue(resolution_result(server_args, "disable_radix_cache")) def test_pd_decode_dcp_rejects_unsupported_transfer_backend(self): server_args = ServerArgs( @@ -601,7 +652,7 @@ def test_pd_decode_dcp_allows_fake_transfer_backend(self): disaggregation_transfer_backend="fake", dcp_size=4, ) - self.assertTrue(server_args.disable_radix_cache) + self.assertTrue(resolution_result(server_args, "disable_radix_cache")) def test_pd_decode_dcp_rejects_radix_cache(self): server_args = ServerArgs( @@ -665,8 +716,11 @@ def test_pd_decode_radix_cache_allows_mooncake_tcp(self): disaggregation_transfer_backend="mooncake_tcp", ) - self.assertFalse(server_args.disable_radix_cache) - self.assertEqual(server_args.disaggregation_transfer_backend, "mooncake") + self.assertFalse(resolution_result(server_args, "disable_radix_cache")) + self.assertEqual( + resolution_result(server_args, "disaggregation_transfer_backend"), + "mooncake", + ) class TestSkipTokenizerInit(unittest.TestCase): @@ -681,8 +735,8 @@ def test_skip_tokenizer_worker_counts(self): server_args._handle_tokenizer_batching() # Tokenizer fanout preserved; detokenizer coerced to 1 (no decode work). - self.assertEqual(server_args.tokenizer_worker_num, 4) - self.assertEqual(server_args.detokenizer_worker_num, 1) + self.assertEqual(resolution_result(server_args, "tokenizer_worker_num"), 4) + self.assertEqual(resolution_result(server_args, "detokenizer_worker_num"), 1) class TestHiSparseDsaBackendPolicy(unittest.TestCase): @@ -870,7 +924,7 @@ def test_combined_attention_backend_fa4_forces_page_size_128( from sglang.srt.arg_groups.overrides import resolved_view - self.assertEqual(args.page_size, 1) # dual-apply retired: pristine + self.assertEqual(args.page_size, 1) # the field stays pristine self.assertEqual(resolved_view(args).page_size, 128) @patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True) @@ -883,7 +937,7 @@ def test_explicit_prefill_fa4_forces_page_size_128(self, _mock_mla, _mock_sm100) from sglang.srt.arg_groups.overrides import resolved_view - self.assertEqual(args.page_size, 1) # dual-apply retired: pristine + self.assertEqual(args.page_size, 1) # the field stays pristine self.assertEqual(resolved_view(args).page_size, 128) @@ -918,12 +972,12 @@ def _new_cp_args(self, **overrides): def test_canonical_prefill_cp_requires_strategy(self): args = self.parser.parse_args(["--model", "dummy", "--enable-prefill-cp"]) - self.assertTrue(args.enable_prefill_cp) - self.assertIsNone(args.cp_strategy) + self.assertTrue(resolution_result(args, "enable_prefill_cp")) + self.assertIsNone(resolution_result(args, "cp_strategy")) server_args = self._new_cp_args( - enable_prefill_cp=args.enable_prefill_cp, - cp_strategy=args.cp_strategy, + enable_prefill_cp=resolution_result(args, "enable_prefill_cp"), + cp_strategy=resolution_result(args, "cp_strategy"), ) with self.assertRaisesRegex(ValueError, "--cp-strategy"): server_args._handle_context_parallelism() @@ -940,16 +994,18 @@ def test_deprecated_dsa_cp_mode_maps_to_unified_strategy(self): ) server_args = self._new_cp_args( enable_dsa_prefill_context_parallel=( - args.enable_dsa_prefill_context_parallel + resolution_result(args, "enable_dsa_prefill_context_parallel") ), - dsa_prefill_cp_mode=args.dsa_prefill_cp_mode, + dsa_prefill_cp_mode=resolution_result(args, "dsa_prefill_cp_mode"), ) server_args._handle_legacy_cp_arguments() - self.assertTrue(server_args.enable_prefill_cp) - self.assertEqual(server_args.cp_strategy, "interleave") - self.assertEqual(server_args.dsa_prefill_cp_mode, "round-robin-split") + self.assertTrue(resolution_result(server_args, "enable_prefill_cp")) + self.assertEqual(resolution_result(server_args, "cp_strategy"), "interleave") + self.assertEqual( + resolution_result(server_args, "dsa_prefill_cp_mode"), "round-robin-split" + ) def test_canonical_interleave_cp_mirrors_to_dsa_runtime_aliases(self): server_args = self._new_cp_args( @@ -961,10 +1017,18 @@ def test_canonical_interleave_cp_mirrors_to_dsa_runtime_aliases(self): server_args._handle_legacy_cp_arguments() server_args._handle_context_parallelism() - self.assertTrue(server_args.enable_dsa_prefill_context_parallel) - self.assertFalse(server_args.enable_prefill_context_parallel) - self.assertEqual(server_args.dsa_prefill_cp_mode, "round-robin-split") - self.assertEqual(server_args.prefill_cp_mode, "round-robin-split") + self.assertTrue( + resolution_result(server_args, "enable_dsa_prefill_context_parallel") + ) + self.assertFalse( + resolution_result(server_args, "enable_prefill_context_parallel") + ) + self.assertEqual( + resolution_result(server_args, "dsa_prefill_cp_mode"), "round-robin-split" + ) + self.assertEqual( + resolution_result(server_args, "prefill_cp_mode"), "round-robin-split" + ) def test_context_parallel_handler_initializes_cp_strategy(self): server_args = self._new_cp_args( @@ -1049,15 +1113,25 @@ def test_registered_cp_legacy_args_map_to_unified_strategy(self): server_args._handle_legacy_cp_arguments() server_args._handle_context_parallelism() - self.assertTrue(server_args.enable_prefill_cp) - self.assertEqual(server_args.cp_strategy, strategy) - self.assertEqual(server_args.dsa_prefill_cp_mode, mode) - self.assertEqual(server_args.prefill_cp_mode, mode) + self.assertTrue(resolution_result(server_args, "enable_prefill_cp")) + self.assertEqual( + resolution_result(server_args, "cp_strategy"), strategy + ) self.assertEqual( - server_args.enable_dsa_prefill_context_parallel, expect_dsa + resolution_result(server_args, "dsa_prefill_cp_mode"), mode ) self.assertEqual( - server_args.enable_prefill_context_parallel, expect_generic + resolution_result(server_args, "prefill_cp_mode"), mode + ) + self.assertEqual( + resolution_result( + server_args, "enable_dsa_prefill_context_parallel" + ), + expect_dsa, + ) + self.assertEqual( + resolution_result(server_args, "enable_prefill_context_parallel"), + expect_generic, ) @@ -1308,7 +1382,7 @@ def test_enable_ssl_refresh_with_ssl_accepted(self, _mock_isfile): ssl_certfile="cert.pem", enable_ssl_refresh=True, ) - self.assertTrue(server_args.enable_ssl_refresh) + self.assertTrue(resolution_result(server_args, "enable_ssl_refresh")) class TestHiCacheArgs(unittest.TestCase): @@ -1328,10 +1402,17 @@ def _assert_hicache_fields( expected_mem_layout: str, expected_decode_backend: str | None = None, ): - self.assertEqual(args.hicache_io_backend, expected_io_backend) - self.assertEqual(args.hicache_mem_layout, expected_mem_layout) + self.assertEqual( + resolution_result(args, "hicache_io_backend"), expected_io_backend + ) + self.assertEqual( + resolution_result(args, "hicache_mem_layout"), expected_mem_layout + ) if expected_decode_backend is not None: - self.assertEqual(args.decode_attention_backend, expected_decode_backend) + self.assertEqual( + resolution_result(args, "decode_attention_backend"), + expected_decode_backend, + ) def test_hicache_io_backend_and_mem_layout_compatibility(self): cases = [ @@ -1409,9 +1490,9 @@ def test_hicache_kernel_keeps_implicit_fa3_decode_backend(self): ) args._handle_hicache() - self.assertEqual(args.hicache_io_backend, "kernel") - self.assertEqual(args.hicache_mem_layout, "page_first") - self.assertIsNone(args.decode_attention_backend) + self.assertEqual(resolution_result(args, "hicache_io_backend"), "kernel") + self.assertEqual(resolution_result(args, "hicache_mem_layout"), "page_first") + self.assertIsNone(resolution_result(args, "decode_attention_backend")) def test_decode_offload_rejects_host_pool_retraction(self): args = self._make_args( @@ -1494,11 +1575,19 @@ def test_decoupled_spec_cli_flags_round_trip(self): "/tmp/tr", ] ) - self.assertEqual(server_args.decoupled_spec_role, "verifier") - self.assertEqual(server_args.decoupled_spec_bind_endpoint, "ipc:///tmp/v") - self.assertEqual(server_args.decoupled_spec_connect_endpoints, ["ipc:///tmp/d"]) - self.assertEqual(server_args.decoupled_spec_rank, 0) - self.assertEqual(server_args.spec_trace_dir, "/tmp/tr") + self.assertEqual( + resolution_result(server_args, "decoupled_spec_role"), "verifier" + ) + self.assertEqual( + resolution_result(server_args, "decoupled_spec_bind_endpoint"), + "ipc:///tmp/v", + ) + self.assertEqual( + resolution_result(server_args, "decoupled_spec_connect_endpoints"), + ["ipc:///tmp/d"], + ) + self.assertEqual(resolution_result(server_args, "decoupled_spec_rank"), 0) + self.assertEqual(resolution_result(server_args, "spec_trace_dir"), "/tmp/tr") def test_decoupled_spec_role_rejects_invalid_choice(self): with self.assertRaises(SystemExit): @@ -1533,10 +1622,10 @@ def test_adaptive_defaults_to_config_step_when_spec_params_omitted(self): handle_speculative_decoding(args) - self.assertTrue(args.speculative_adaptive) - self.assertEqual(args.speculative_eagle_topk, 1) - self.assertEqual(args.speculative_num_steps, 3) - self.assertEqual(args.speculative_num_draft_tokens, 4) + self.assertTrue(resolution_result(args, "speculative_adaptive")) + self.assertEqual(resolution_result(args, "speculative_eagle_topk"), 1) + self.assertEqual(resolution_result(args, "speculative_num_steps"), 3) + self.assertEqual(resolution_result(args, "speculative_num_draft_tokens"), 4) class TestWaterfillArgs(CustomTestCase): @@ -1552,10 +1641,10 @@ def test_waterfill_enforces_shared_experts_fusion(self): from sglang.srt.arg_groups.overrides import resolved_view - # dual-apply retired: the fields stay pristine, the declarations win + # the fields stay pristine, the declarations win self.assertTrue(server_args.disable_shared_experts_fusion) self.assertFalse(resolved_view(server_args).disable_shared_experts_fusion) - self.assertTrue(server_args.enforce_shared_experts_fusion) + self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion")) def test_waterfill_overrides_moe_a2a_backend_to_deepep(self): server_args = ServerArgs( @@ -1570,7 +1659,7 @@ def test_waterfill_overrides_moe_a2a_backend_to_deepep(self): self.assertEqual(server_args.moe_a2a_backend, "none") # pristine self.assertEqual(resolved_view(server_args).moe_a2a_backend, "deepep") - self.assertTrue(server_args.enforce_shared_experts_fusion) + self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion")) def test_waterfill_keeps_megamoe_backend(self): server_args = ServerArgs( @@ -1586,7 +1675,7 @@ def test_waterfill_keeps_megamoe_backend(self): self.assertEqual(resolved_view(server_args).moe_a2a_backend, "megamoe") self.assertFalse(resolved_view(server_args).disable_shared_experts_fusion) - self.assertTrue(server_args.enforce_shared_experts_fusion) + self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion")) def test_waterfill_supports_deepep_low_latency_mode(self): server_args = ServerArgs( @@ -1598,9 +1687,9 @@ def test_waterfill_supports_deepep_low_latency_mode(self): # dummy-model path short-circuits __post_init__; invoke the handler directly. server_args._handle_a2a_moe() - self.assertEqual(server_args.deepep_mode, "low_latency") - self.assertFalse(server_args.disable_cuda_graph) - self.assertTrue(server_args.enforce_shared_experts_fusion) + self.assertEqual(resolution_result(server_args, "deepep_mode"), "low_latency") + self.assertFalse(resolution_result(server_args, "disable_cuda_graph")) + self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion")) class TestPrefillOnlyDisableKvCache(unittest.TestCase): @@ -1635,7 +1724,7 @@ def _validate_prefill_only_args(self, **overrides): def test_valid_minimal_config_constructs(self): sa = self._validate_prefill_only_args() - self.assertTrue(sa.prefill_only_disable_kv_cache) + self.assertTrue(resolution_result(sa, "prefill_only_disable_kv_cache")) def test_rejects_when_not_embedding(self): with self.assertRaisesRegex(ValueError, "requires --is-embedding"): @@ -1725,15 +1814,27 @@ def _handled_args(self, **overrides): def test_cuda_graph_prefill_role_defaults_disable_decode_graph(self): args = self._handled_args(disaggregation_mode="prefill") - self.assertFalse(args.disable_cuda_graph) - self.assertEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED) - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE) + self.assertFalse(resolution_result(args, "disable_cuda_graph")) + self.assertEqual( + resolution_result(args, "cuda_graph_config").decode.backend, + Backend.DISABLED, + ) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.BREAKABLE, + ) def test_cuda_graph_decode_role_defaults_disable_prefill_graph(self): args = self._handled_args(disaggregation_mode="decode") - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED) - self.assertNotEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.DISABLED, + ) + self.assertNotEqual( + resolution_result(args, "cuda_graph_config").decode.backend, + Backend.DISABLED, + ) def test_cuda_graph_global_disable_still_disables_both_phases_for_all_roles(self): for disaggregation_mode in ("prefill", "decode", "null"): @@ -1744,10 +1845,12 @@ def test_cuda_graph_global_disable_still_disables_both_phases_for_all_roles(self ) self.assertEqual( - args.cuda_graph_config.decode.backend, Backend.DISABLED + resolution_result(args, "cuda_graph_config").decode.backend, + Backend.DISABLED, ) self.assertEqual( - args.cuda_graph_config.prefill.backend, Backend.DISABLED + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.DISABLED, ) def test_cuda_graph_explicit_decode_backend_survives_prefill_role(self): @@ -1756,7 +1859,9 @@ def test_cuda_graph_explicit_decode_backend_survives_prefill_role(self): cuda_graph_backend_decode=Backend.FULL, ) - self.assertEqual(args.cuda_graph_config.decode.backend, Backend.FULL) + self.assertEqual( + resolution_result(args, "cuda_graph_config").decode.backend, Backend.FULL + ) self.assertIn((Phase.DECODE, "backend"), args._cuda_graph_config_locked) @@ -1782,12 +1887,18 @@ def _handled_args(self, **overrides): def test_enable_lora_keeps_breakable_prefill_graph(self): args = self._handled_args(enable_lora=True) - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.BREAKABLE, + ) def test_lora_paths_keep_breakable_prefill_graph(self): args = self._handled_args(lora_paths=["dummy/lora-adapter"]) - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.BREAKABLE, + ) def test_lora_still_disables_tc_piecewise_prefill_graph(self): # Pin the tc_piecewise LoRA rule itself, with the hardware rule @@ -1811,7 +1922,10 @@ def test_lora_still_disables_tc_piecewise_prefill_graph(self): ): args._disable_tc_piecewise_cudagraph_if_incompatible() - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.DISABLED, + ) class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase): @@ -1840,7 +1954,10 @@ def test_multimodal_arch_disables_prefill_breakable(self): is_multimodal=True, allowlisted=False, ) - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.DISABLED, + ) def test_allowlisted_multimodal_arch_keeps_prefill_breakable(self): args = self._handled_args( @@ -1848,7 +1965,10 @@ def test_allowlisted_multimodal_arch_keeps_prefill_breakable(self): is_multimodal=True, allowlisted=True, ) - self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.BREAKABLE, + ) def test_allowlist_membership(self): from sglang.srt.configs.model_config import ( @@ -2046,20 +2166,20 @@ def _args(**kwargs): def test_http_only_high_port_does_not_derive_grpc_port(self): sa = self._args(port=56000) sa._handle_deprecated_args() - self.assertIsNone(sa.grpc_port) + self.assertIsNone(resolution_result(sa, "grpc_port")) def test_grpc_port_enables_native_and_env_knobs(self): sa = self._args(grpc_port=50051) with envs.SGLANG_GRPC_WORKER_THREADS.override(8): sa._handle_deprecated_args() - self.assertEqual(sa.grpc_port, 50051) + self.assertEqual(resolution_result(sa, "grpc_port"), 50051) self.assertEqual(sa.grpc_worker_threads, 8) def test_env_grpc_port_enables_native(self): sa = self._args(port=30000) with envs.SGLANG_GRPC_PORT.override(45000): sa._handle_deprecated_args() - self.assertEqual(sa.grpc_port, 45000) + self.assertEqual(resolution_result(sa, "grpc_port"), 45000) @staticmethod def _sidecar_parser(): @@ -2196,20 +2316,20 @@ def test_sidecar_stop_uses_configured_shutdown_timeout(self): def test_legacy_smg_derives_grpc_port_from_http_port(self): sa = self._args(port=30000, smg_grpc_mode=True) sa._handle_deprecated_args() - self.assertEqual(sa.grpc_port, 40000) + self.assertEqual(resolution_result(sa, "grpc_port"), 40000) def test_grpc_mode_is_deprecated_alias_for_smg_grpc_mode(self): sa = self._args(grpc_mode=True) with self.assertLogs(server_args_module.logger, level="WARNING") as cm: sa._handle_deprecated_args() - self.assertTrue(sa.smg_grpc_mode) + self.assertTrue(resolution_result(sa, "smg_grpc_mode")) self.assertTrue(any("--grpc-mode is deprecated" in line for line in cm.output)) def test_legacy_smg_takes_precedence_over_grpc_port(self): sa = self._args(grpc_port=50051, smg_grpc_mode=True) sa._handle_deprecated_args() - self.assertTrue(sa.smg_grpc_mode) - self.assertEqual(sa.grpc_port, 50051) + self.assertTrue(resolution_result(sa, "smg_grpc_mode")) + self.assertEqual(resolution_result(sa, "grpc_port"), 50051) def test_native_grpc_rejects_multi_tokenizer(self): sa = self._args(grpc_port=40000, tokenizer_worker_num=2) @@ -2254,7 +2374,7 @@ def test_start_server_call_site_matches_native_signature(self): tokenizer_manager=MagicMock(), template_manager=MagicMock(), scheduler_info={}, - grpc_port=server_args.grpc_port, + grpc_port=resolution_result(server_args, "grpc_port"), ) self.assertEqual(handle, "handle") diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index c7d748f6a623..fbb40f2b4444 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -280,10 +280,16 @@ def test_dummy_fixture_publishes_the_object_it_resolved(self): set_global_server_args_for_scheduler(sa) self.assertIs(get_server_args(), sa) # Publishing is what resolved it; the handlers ahead of the dummy - # short-circuit still declare. + # short-circuit still declare. What they decided is the projection -- + # the fields keep what the caller passed. + from sglang.srt.arg_groups.overrides import resolution_result + + self.assertTrue(sa._resolved_overrides, "publishing declared nothing") for source, declared in sa._resolved_overrides: for field, value in declared.items(): - self.assertEqual(getattr(sa, field), value, f"{source}: {field}") + self.assertEqual( + resolution_result(sa, field), value, f"{source}: {field}" + ) class TestGoldenModelOverrides(_IsolatedPublish): diff --git a/test/registered/unit/test_runtime_context_override.py b/test/registered/unit/test_runtime_context_override.py index e7dfd144e4f4..623125a00ef2 100644 --- a/test/registered/unit/test_runtime_context_override.py +++ b/test/registered/unit/test_runtime_context_override.py @@ -31,11 +31,15 @@ def _publish(self): def test_override_writes_bag_not_server_args(self): sa = self._publish() - before = sa.hicache_ratio + # The published leaf, not the field: `hicache_ratio` is resolved by + # declaration, so the field still holds what the caller passed. + before = rc.get_memory().hicache_ratio + pristine = sa.hicache_ratio rc.get_context().override("test", hicache_ratio=before + 1.0) self.assertEqual(rc.get_memory().hicache_ratio, before + 1.0) - # server_args stays the pristine startup record. - self.assertEqual(sa.hicache_ratio, before) + # server_args stays the pristine startup record: the override does not + # touch it, and neither did resolution. + self.assertEqual(sa.hicache_ratio, pristine) def test_override_routes_across_namespaces(self): self._publish() diff --git a/test/registered/unit/test_server_args_migration.py b/test/registered/unit/test_server_args_migration.py index 1ea89e8c31b4..67b2d9eb72e8 100644 --- a/test/registered/unit/test_server_args_migration.py +++ b/test/registered/unit/test_server_args_migration.py @@ -96,8 +96,7 @@ def test_media_url_security_args(self): "32", ] ) - # The normalization is a declaration, so it lands in the - # resolution result rather than on the field. + # The normalization is a declaration. self.assertEqual( resolution_result(sa, "allowed_media_domains"), ["127.0.0.1", "media.example.com"], diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index 466e622b3a42..23e863670f17 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -134,10 +134,10 @@ # are step-12 exposure like any other pair. _PASSED = frozenset({"model_path", "device", "random_seed"}) +# The reads that still take a value off the supplied instance. `initialize_moe_config` +# is handed the record until the replay goes away; the rest are pre-publish launcher +# reads. _EXPOSED = { - ("dllm/config.py", "max_running_requests"), - ("dllm/config.py", "model_path"), - ("speculative/spec_registry.py", "disable_overlap_schedule"), ("disaggregation/encoder/server.py", "model_loader_extra_config"), ("layers/moe/utils.py", "deepep_mode"), ("layers/moe/utils.py", "disable_shared_experts_fusion"), @@ -150,38 +150,16 @@ ("configs/embedding_model_spec.py", "disable_radix_cache"), ("configs/embedding_model_spec.py", "is_embedding"), ("configs/embedding_model_spec.py", "prefill_only_disable_kv_cache"), - ("configs/model_config.py", "_speculative_draft_quantization_explicitly_set"), - ("configs/model_config.py", "disable_hybrid_swa_memory"), - ("configs/model_config.py", "dtype"), - ("configs/model_config.py", "enable_multi_layer_eagle"), - ("configs/model_config.py", "is_embedding"), - ("configs/model_config.py", "model_path"), - ("configs/model_config.py", "quantization"), - ("configs/model_config.py", "speculative_algorithm"), - ("configs/model_config.py", "speculative_draft_model_quantization"), - ("dllm/config.py", "max_running_requests"), - ("dllm/config.py", "model_path"), ("entrypoints/engine.py", "enable_symm_mem"), ("entrypoints/engine.py", "reasoning_parser"), ("entrypoints/engine.py", "tool_call_parser"), - ("layers/cp/base.py", "attn_cp_size"), - ("layers/cp/base.py", "cp_strategy"), - ("layers/cp/base.py", "enable_prefill_cp"), - ("layers/cp/bcg.py", "cp_strategy"), - ("layers/cp/bcg.py", "enable_prefill_cp"), ("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"), ("layers/moe/utils.py", "deepep_mode"), ("layers/moe/utils.py", "moe_a2a_backend"), ("layers/moe/utils.py", "moe_runner_backend"), ("layers/moe/utils.py", "quantization"), ("layers/moe/utils.py", "speculative_moe_runner_backend"), - ("lora/marlin_lora_temp/policy.py", "lora_paths"), - ("parser/template_detection.py", "model_path"), - ("speculative/adaptive_spec_params.py", "speculative_algorithm"), - ("speculative/adaptive_spec_params.py", "speculative_eagle_topk"), ("speculative/draft_worker_common.py", "speculative_draft_attention_backend"), - ("speculative/spec_info.py", "enable_multi_layer_eagle"), - ("speculative/spec_registry.py", "disable_overlap_schedule"), ("utils/common.py", "speculative_num_draft_tokens"), ("utils/common.py", "speculative_num_steps"), ("utils/hf_transformers/processor.py", "image_processor_backend"), @@ -214,9 +192,6 @@ # some code overrides post-publish. Each needs an ordering judgment, not a blanket # conversion; the list exists so a new one is a decision made when it is written. _OVERRIDDEN_AND_READ = { - ("configs/model_config.py", "dtype"), - ("configs/model_config.py", "model_path"), - ("dllm/config.py", "model_path"), ("entrypoints/engine.py", "reasoning_parser"), ("entrypoints/engine.py", "tool_call_parser"), ("weight_cache/daemon.py", "dp_size"), @@ -224,15 +199,12 @@ ("weight_cache/daemon.py", "ep_size"), ("weight_cache/daemon.py", "load_format"), ("weight_cache/daemon.py", "model_path"), - ("configs/model_config.py", "dtype"), - ("configs/model_config.py", "model_path"), ("mem_cache/pool_host/common.py", "hicache_storage_backend"), ("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"), ("mem_cache/unified_radix_cache.py", "hicache_storage_backend"), ("mem_cache/unified_radix_cache.py", "hicache_storage_backend_extra_config"), ("mem_cache/unified_radix_cache.py", "hicache_storage_prefetch_policy"), ("mem_cache/unified_radix_cache.py", "hicache_write_policy"), - ("parser/template_detection.py", "model_path"), ("utils/common.py", "speculative_num_draft_tokens"), ("utils/common.py", "speculative_num_steps"), ("weight_cache/daemon.py", "dp_size"), From 93f1eed21f3e60f249937225677f61112863c1d9 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 19:21:36 +0000 Subject: [PATCH 14/26] config: four more readers take the bag instead of a record `resolve_image_processor_backend` already had one caller passing `get_mm()` and three passing a record; all four pass the bag now, and the parameter is named for what it is. The FlashInfer all-reduce fusion resolver and the draft attention-backend fallback read their leaves directly and take no config at all -- the fusion one was reaching for `get_server_args()`, which is a global record read that only escaped the ratchet because it handed the whole object to a helper. `reserve_rope_cache_for_long_sequences` reads `model.context_length` and the two `spec` counts. The FlashInfer fusion test drives `_resolve_backend(backend, is_multi_node)` directly: the arch dispatch is what those cases are about, and the entry above it now takes no arguments. Five (file, field) pairs leave the supplied-instance exposure pin. --- .../disaggregation/encoder/preprocessor.py | 4 +- .../srt/disaggregation/encoder/receiver.py | 9 ++- .../srt/layers/flashinfer_comm_fusion.py | 17 ++--- .../sglang/srt/managers/tokenizer_manager.py | 2 +- .../sglang/srt/model_executor/model_runner.py | 1 - .../srt/speculative/draft_worker_common.py | 19 +++--- python/sglang/srt/utils/common.py | 18 +++--- .../srt/utils/hf_transformers/processor.py | 13 ++-- .../layers/test_flashinfer_comm_fusion.py | 62 ++++++++----------- 9 files changed, 78 insertions(+), 67 deletions(-) diff --git a/python/sglang/srt/disaggregation/encoder/preprocessor.py b/python/sglang/srt/disaggregation/encoder/preprocessor.py index 6d8e0561ad7b..ffa48a8b493a 100644 --- a/python/sglang/srt/disaggregation/encoder/preprocessor.py +++ b/python/sglang/srt/disaggregation/encoder/preprocessor.py @@ -127,7 +127,7 @@ def __init__( use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get() self.use_image_processor_gpu = ( use_image_processor_gpu - and resolve_image_processor_backend(server_args) != "pil" + and resolve_image_processor_backend(get_mm()) != "pil" ) self._load_mm_processor(server_args) @@ -158,7 +158,7 @@ def __init__( def _load_mm_processor(self, server_args: ServerArgs): from transformers import AutoImageProcessor, AutoVideoProcessor - image_processor_backend = resolve_image_processor_backend(server_args) + image_processor_backend = resolve_image_processor_backend(get_mm()) image_processor_kwargs = ( {} if image_processor_backend == "auto" diff --git a/python/sglang/srt/disaggregation/encoder/receiver.py b/python/sglang/srt/disaggregation/encoder/receiver.py index 50b75ccaea13..9e2d99a755f8 100644 --- a/python/sglang/srt/disaggregation/encoder/receiver.py +++ b/python/sglang/srt/disaggregation/encoder/receiver.py @@ -35,7 +35,12 @@ from sglang.srt.managers.schedule_batch import Modality, Req from sglang.srt.multimodal.cache import media_preprocess_kwargs from sglang.srt.multimodal.transport import determine_tensor_transport_mode -from sglang.srt.runtime_context import get_disagg, get_exec, get_serving +from sglang.srt.runtime_context import ( + get_disagg, + get_exec, + get_mm, + get_serving, +) from sglang.srt.server_args import ServerArgs from sglang.srt.utils import ImageData from sglang.srt.utils.common import safe_pickle_loads @@ -1838,7 +1843,7 @@ def _init_mm_processor( tokenizer_mode=server_args.tokenizer_mode, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, - image_processor_backend=resolve_image_processor_backend(server_args), + image_processor_backend=resolve_image_processor_backend(get_mm()), **extra_kwargs, ) diff --git a/python/sglang/srt/layers/flashinfer_comm_fusion.py b/python/sglang/srt/layers/flashinfer_comm_fusion.py index cd399649cfb0..8c1aa58a70d6 100644 --- a/python/sglang/srt/layers/flashinfer_comm_fusion.py +++ b/python/sglang/srt/layers/flashinfer_comm_fusion.py @@ -13,7 +13,7 @@ get_tp_group, ) from sglang.srt.distributed.parallel_state import in_the_same_node_as -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import ( ceil_align, get_cuda_driver_bindings, @@ -75,12 +75,16 @@ def _resolve_backend(backend: str, is_multi_node: bool = False) -> str: return backend -def resolve_flashinfer_allreduce_fusion_backend(server_args) -> Optional[str]: - backend = getattr(server_args, "flashinfer_allreduce_fusion_backend", None) +def resolve_flashinfer_allreduce_fusion_backend() -> Optional[str]: + """The fusion backend for this process, or None when fusion is off. + + Reads the published leaves (`exec.comm`, `parallel`): the backend is a + resolution decision, and the node count is launch topology. + """ + backend = get_exec().comm.flashinfer_allreduce_fusion_backend if backend is None: return None - is_multi_node = getattr(server_args, "nnodes", 1) > 1 - return _resolve_backend(backend, is_multi_node) + return _resolve_backend(backend, get_parallel().config.nnodes > 1) if is_flashinfer_available(): @@ -716,8 +720,7 @@ def ensure_workspace_initialized( token_num = token_num or max_token_num group_key = (device_group, cpu_group) effective_dtype = dtype or torch.bfloat16 - server_args = get_server_args() - backend = resolve_flashinfer_allreduce_fusion_backend(server_args) + backend = resolve_flashinfer_allreduce_fusion_backend() if backend is None: return False diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index ee97103c872d..6ec332baa001 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -3591,7 +3591,7 @@ def get_processor_wrapper(server_args): 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), + image_processor_backend=resolve_image_processor_backend(get_mm()), tokenizer_backend=get_serving().tokenizer_backend, model_name=get_model().model_path, ) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index f7947fb61a83..d4d6e3e54f24 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1202,7 +1202,6 @@ def load_model(self): # Pre-expand RoPE cache before CUDA Graph capture reserve_rope_cache_for_long_sequences( self.model, - self.server_args, self.model_config, logger, ) diff --git a/python/sglang/srt/speculative/draft_worker_common.py b/python/sglang/srt/speculative/draft_worker_common.py index c7c1c14fedcc..eb9ddb80631d 100644 --- a/python/sglang/srt/speculative/draft_worker_common.py +++ b/python/sglang/srt/speculative/draft_worker_common.py @@ -9,6 +9,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode +from sglang.srt.runtime_context import attention_backends, get_spec from sglang.srt.server_args import DRAFT_ATTENTION_BACKEND_CHOICES, ServerArgs from sglang.srt.speculative.dflash_info import DFlashVerifyInput from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 @@ -28,12 +29,16 @@ class DraftWorkerBundle(msgspec.Struct, frozen=True): resolved_attention_backend: str -def _resolve_draft_attention_backend_fallback( - *, server_args: ServerArgs, algo_label: str -) -> str: - draft_backend = server_args.speculative_draft_attention_backend +def _resolve_draft_attention_backend_fallback(*, algo_label: str) -> str: + """The draft's attention backend, from the published leaves. + + `spec.speculative_draft_attention_backend` when the operator named one, + otherwise the process's prefill backend. Both are resolution's answers, so + they come from the bags. + """ + draft_backend = get_spec().speculative_draft_attention_backend if draft_backend is None: - draft_backend, _ = server_args.get_attention_backends() + draft_backend, _ = attention_backends() if draft_backend is None: return "triton" if torch.version.hip else "flashinfer" if draft_backend not in DRAFT_ATTENTION_BACKEND_CHOICES: @@ -65,9 +70,7 @@ def build_draft_tp_worker( # validated (e.g. a self-drafting architecture); it skips the generic # supported-backend fallback below. draft_backend = attention_backend_override or ( - _resolve_draft_attention_backend_fallback( - server_args=server_args, algo_label=algo_label - ) + _resolve_draft_attention_backend_fallback(algo_label=algo_label) ) from sglang.srt.layers.moe.utils import draft_model_build_scope diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index c7972f9253a7..8f826768e8f3 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -4623,11 +4623,15 @@ def decorator(fn): return decorator -def reserve_rope_cache_for_long_sequences( - model, server_args, model_config, logger=None -): - """Pre-expand RoPE cache for long sequences and speculative decoding.""" +def reserve_rope_cache_for_long_sequences(model, model_config, logger=None): + """Pre-expand RoPE cache for long sequences and speculative decoding. + + Runs inside `ModelRunner`, past publish, so the three config inputs come + from the bags: the context length and the two speculative counts are + resolution's answers. + """ from sglang.srt.environ import envs + from sglang.srt.runtime_context import get_model, get_spec SAFETY_FACTOR = envs.SGLANG_SPEC_EXPANSION_SAFETY_FACTOR.get() MARGIN = envs.SGLANG_ROPE_CACHE_SAFETY_MARGIN.get() @@ -4635,7 +4639,7 @@ def reserve_rope_cache_for_long_sequences( # 1) Estimate base context upper bound base_ctx = ( - getattr(server_args, "context_length", None) + get_model().context_length or getattr(model_config, "context_len", None) or getattr(model_config, "max_model_len", None) or getattr(model_config.hf_text_config, "max_position_embeddings", None) @@ -4643,8 +4647,8 @@ def reserve_rope_cache_for_long_sequences( ) # 2) Speculative decoding expansion - steps = int(getattr(server_args, "speculative_num_steps", 0) or 0) - draft = int(getattr(server_args, "speculative_num_draft_tokens", 0) or 0) + steps = int(get_spec().speculative_num_steps or 0) + draft = int(get_spec().speculative_num_draft_tokens or 0) reserve = base_ctx + steps * draft * SAFETY_FACTOR + MARGIN # 3) Align to reduce reallocation frequency diff --git a/python/sglang/srt/utils/hf_transformers/processor.py b/python/sglang/srt/utils/hf_transformers/processor.py index de50416db576..c0c7a841161a 100644 --- a/python/sglang/srt/utils/hf_transformers/processor.py +++ b/python/sglang/srt/utils/hf_transformers/processor.py @@ -54,11 +54,16 @@ _IMAGE_PROCESSOR_BACKENDS = {"auto", "torchvision", "pil"} -def resolve_image_processor_backend(server_args) -> str: - """Resolve the new backend option while honoring the legacy disable flag.""" - if getattr(server_args, "disable_fast_image_processor", False): +def resolve_image_processor_backend(mm_config) -> str: + """Resolve the new backend option while honoring the legacy disable flag. + + Takes the `mm` config bag (`get_mm()`): both leaves are resolved config, and + every caller is past publish. `getattr` with a default keeps it working for a + stand-in that carries only one of the two. + """ + if getattr(mm_config, "disable_fast_image_processor", False): return "pil" - return getattr(server_args, "image_processor_backend", "auto") + return getattr(mm_config, "image_processor_backend", "auto") def _normalize_image_processor_backend( diff --git a/test/registered/unit/layers/test_flashinfer_comm_fusion.py b/test/registered/unit/layers/test_flashinfer_comm_fusion.py index 9eaf5c8fc27a..5fdc935b81c8 100644 --- a/test/registered/unit/layers/test_flashinfer_comm_fusion.py +++ b/test/registered/unit/layers/test_flashinfer_comm_fusion.py @@ -1,5 +1,4 @@ import contextlib -import types import unittest from unittest.mock import MagicMock, patch @@ -90,23 +89,24 @@ def _torch_allreduce_residual_rmsnorm_baseline( class TestFlashInferCommFusion(CustomTestCase): + """The arch dispatch is `_resolve_backend(backend, is_multi_node)`. + + The public entry above it takes no arguments -- it reads + `exec.comm.flashinfer_allreduce_fusion_backend` and `parallel.nnodes` off the + published bags -- so the cases here drive the dispatch directly. + """ + def test_auto_backend_resolves_by_arch(self): - single_node = types.SimpleNamespace( - flashinfer_allreduce_fusion_backend="auto", nnodes=1 - ) - multi_node = types.SimpleNamespace( - flashinfer_allreduce_fusion_backend="auto", nnodes=2 - ) + single_node = ("auto", False) + multi_node = ("auto", True) # Blackwell: mnnvl on both single-node and multi-node. with patch.object(fusion, "is_sm100_supported", return_value=True): self.assertEqual( - fusion.resolve_flashinfer_allreduce_fusion_backend(single_node), + fusion._resolve_backend(*single_node), "mnnvl", ) - self.assertEqual( - fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node), "mnnvl" - ) + self.assertEqual(fusion._resolve_backend(*multi_node), "mnnvl") # SM90: auto uses trtllm on single-node, multi-node is unsupported. with ( @@ -114,11 +114,11 @@ def test_auto_backend_resolves_by_arch(self): patch.object(fusion, "is_sm90_supported", return_value=True), ): self.assertEqual( - fusion.resolve_flashinfer_allreduce_fusion_backend(single_node), + fusion._resolve_backend(*single_node), "trtllm", ) with self.assertRaises(ValueError): - fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node) + fusion._resolve_backend(*multi_node) # Architectures outside SM90/SM10X are unsupported. Both pre-SM90 # and post-SM10X devices (e.g. SM120) must fail closed. @@ -129,48 +129,40 @@ def test_auto_backend_resolves_by_arch(self): patch.object(fusion, "is_sm90_supported", return_value=False), ): with self.assertRaises(ValueError): - fusion.resolve_flashinfer_allreduce_fusion_backend(single_node) + fusion._resolve_backend(*single_node) with self.assertRaises(ValueError): - fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node) + fusion._resolve_backend(*multi_node) def test_explicit_backend_validation(self): - single_node_mnnvl = types.SimpleNamespace( - flashinfer_allreduce_fusion_backend="mnnvl", nnodes=1 - ) - multi_node_mnnvl = types.SimpleNamespace( - flashinfer_allreduce_fusion_backend="mnnvl", nnodes=2 - ) - single_node_trtllm = types.SimpleNamespace( - flashinfer_allreduce_fusion_backend="trtllm", nnodes=1 - ) - multi_node_trtllm = types.SimpleNamespace( - flashinfer_allreduce_fusion_backend="trtllm", nnodes=2 - ) + single_node_mnnvl = ("mnnvl", False) + multi_node_mnnvl = ("mnnvl", True) + single_node_trtllm = ("trtllm", False) + multi_node_trtllm = ("trtllm", True) with ( patch.object(fusion, "is_sm100_supported", return_value=False), patch.object(fusion, "is_sm90_supported", return_value=True), ): self.assertEqual( - fusion.resolve_flashinfer_allreduce_fusion_backend(single_node_mnnvl), + fusion._resolve_backend(*single_node_mnnvl), "mnnvl", ) self.assertEqual( - fusion.resolve_flashinfer_allreduce_fusion_backend(single_node_trtllm), + fusion._resolve_backend(*single_node_trtllm), "trtllm", ) with self.assertRaises(ValueError): - fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_mnnvl) + fusion._resolve_backend(*multi_node_mnnvl) with self.assertRaises(ValueError): - fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_trtllm) + fusion._resolve_backend(*multi_node_trtllm) with patch.object(fusion, "is_sm100_supported", return_value=True): self.assertEqual( - fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_mnnvl), + fusion._resolve_backend(*multi_node_mnnvl), "mnnvl", ) with self.assertRaises(ValueError): - fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_trtllm) + fusion._resolve_backend(*multi_node_trtllm) for arch in ("pre_sm90", "post_sm10x"): with ( @@ -184,9 +176,9 @@ def test_explicit_backend_validation(self): single_node_trtllm, multi_node_trtllm, ): - with self.subTest(backend=args.flashinfer_allreduce_fusion_backend): + with self.subTest(backend=args[0], multi_node=args[1]): with self.assertRaises(ValueError): - fusion.resolve_flashinfer_allreduce_fusion_backend(args) + fusion._resolve_backend(*args) def test_allreduce_fusion_backends_match_torch_baseline(self): fake_comm = _FakeFlashInferComm() From fd78e623a244782f17a3b4f8418bbd37b046e44c Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 19:38:18 +0000 Subject: [PATCH 15/26] config: the CP strategy binder takes the three values it needs `init_cp_strategy(server_args)` was called from two places that cannot read the same source: resolution calls it inside `__post_init__`, where the bags do not exist yet, and `get_cp_strategy` calls it lazily in a worker process, where the record is not where the resolved sizes live -- that path was reaching for `get_server_args()` and handing the whole object over, which is how a global record read escapes the ratchet. It takes `enable_prefill_cp`, `cp_size` and `cp_strategy` now. Resolution passes them off its view, the lazy path off `get_parallel().config`, and the unit tests pass them directly instead of building a `SimpleNamespace` per case. --- python/sglang/srt/layers/cp/base.py | 34 +++++++------ python/sglang/srt/server_args.py | 6 ++- test/registered/cp/test_cp_strategy_unit.py | 50 +++++++------------ .../unit/test_global_config_read_ratchet.py | 4 ++ 4 files changed, 47 insertions(+), 47 deletions(-) diff --git a/python/sglang/srt/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index c57e33e97c8d..225fc4fdb399 100644 --- a/python/sglang/srt/layers/cp/base.py +++ b/python/sglang/srt/layers/cp/base.py @@ -33,7 +33,6 @@ if TYPE_CHECKING: from sglang.srt.model_executor.forward_batch_info import ForwardBatch - from sglang.srt.server_args import ServerArgs class ContextParallelStrategyKind(IntEnum): @@ -237,23 +236,27 @@ def _is_dsa_active() -> bool: _STRATEGY: Optional[ContextParallelStrategy] = None -def init_cp_strategy(server_args: ServerArgs) -> None: - """Bind the configured CP strategy for this process.""" - from sglang.srt.arg_groups.overrides import resolving_view +def init_cp_strategy( + *, enable_prefill_cp: bool, cp_size: int, cp_strategy: str +) -> None: + """Bind the CP strategy for this process. - cfg = resolving_view(server_args) + Takes the three values: resolution calls this from inside `__post_init__`, + where the bags do not exist yet, and `get_cp_strategy` calls it lazily in a + worker, which reads them off the published bags. Each caller reads from the + source it has. + """ global _STRATEGY - if not cfg.enable_prefill_cp: + if not enable_prefill_cp: _STRATEGY = None return - cp_size = cfg.attn_cp_size if cp_size <= 1: _STRATEGY = None return - kind = ContextParallelStrategyKind.from_string(cfg.cp_strategy) + kind = ContextParallelStrategyKind.from_string(cp_strategy) if kind == ContextParallelStrategyKind.ZIGZAG: from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy @@ -264,8 +267,7 @@ def init_cp_strategy(server_args: ServerArgs) -> None: _STRATEGY = InterleaveCPStrategy(cp_size=cp_size) else: raise ValueError( - f"Unsupported cp_strategy kind {kind} for " - f"cp_strategy={cfg.cp_strategy!r}" + f"Unsupported cp_strategy kind {kind} for cp_strategy={cp_strategy!r}" ) @@ -280,14 +282,16 @@ def get_cp_strategy() -> Optional[ContextParallelStrategy]: global _STRATEGY if _STRATEGY is None: - from sglang.srt.runtime_context import get_server_args - try: - server_args = get_server_args() + parallel = get_parallel().config except ValueError: return None - if server_args is not None and get_parallel().config.enable_prefill_cp: - init_cp_strategy(server_args) + if parallel.enable_prefill_cp: + init_cp_strategy( + enable_prefill_cp=True, + cp_size=parallel.attn_cp_size, + cp_strategy=parallel.cp_strategy, + ) return _STRATEGY diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 17a4afa4942f..a90733df68a8 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -7100,7 +7100,11 @@ def _handle_context_parallelism(self): from sglang.srt.layers.cp.base import init_cp_strategy - init_cp_strategy(self) + init_cp_strategy( + enable_prefill_cp=bool(cfg.enable_prefill_cp), + cp_size=cfg.attn_cp_size, + cp_strategy=cfg.cp_strategy, + ) def _handle_dwdp(self): cfg = resolving_view(self) diff --git a/test/registered/cp/test_cp_strategy_unit.py b/test/registered/cp/test_cp_strategy_unit.py index 4b8d6f8bd21d..53cc4fa2ee28 100644 --- a/test/registered/cp/test_cp_strategy_unit.py +++ b/test/registered/cp/test_cp_strategy_unit.py @@ -60,7 +60,7 @@ def all_gather_into_tensor(self, output, input_tensor): class TestCPStrategyUnit(CustomTestCase): def tearDown(self): - init_cp_strategy(SimpleNamespace(enable_prefill_cp=False)) + init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="zigzag") def test_strategy_kind_maps_cli_values(self): self.assertEqual(ContextParallelStrategyKind.NONE.value, 0) @@ -77,11 +77,9 @@ def test_strategy_kind_maps_cli_values(self): def test_init_cp_strategy_binds_zigzag_strategy(self): init_cp_strategy( - SimpleNamespace( - enable_prefill_cp=True, - cp_strategy="zigzag", - attn_cp_size=4, - ) + enable_prefill_cp=True, + cp_size=4, + cp_strategy="zigzag", ) self.assertTrue(is_cp_enabled()) @@ -91,11 +89,9 @@ def test_init_cp_strategy_binds_zigzag_strategy(self): def test_get_cp_strategy_is_initialized_under_cp_v2(self): init_cp_strategy( - SimpleNamespace( - enable_prefill_cp=True, - cp_strategy="interleave", - attn_cp_size=4, - ) + enable_prefill_cp=True, + cp_size=4, + cp_strategy="interleave", ) with patch( @@ -108,7 +104,7 @@ def test_get_cp_strategy_is_initialized_under_cp_v2(self): class TestPrefillCPBCGReplay(CustomTestCase): def tearDown(self): - init_cp_strategy(SimpleNamespace(enable_prefill_cp=False)) + init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="zigzag") def _make_runner(self): runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner) @@ -141,11 +137,9 @@ def _make_forward_batch(self): def _enable_zigzag(self): init_cp_strategy( - SimpleNamespace( - enable_prefill_cp=True, - cp_strategy="zigzag", - attn_cp_size=4, - ) + enable_prefill_cp=True, + cp_size=4, + cp_strategy="zigzag", ) def test_local_capacity_overflow_uses_next_capture_bucket(self): @@ -267,16 +261,13 @@ def fill_from(self, _source, **kwargs): class TestCPZigzagStrategy(CustomTestCase): def setUp(self): init_cp_strategy( - SimpleNamespace( - enable_prefill_cp=True, - cp_strategy="zigzag", - attn_cp_size=4, - attention_backend="fa3", - ) + enable_prefill_cp=True, + cp_size=4, + cp_strategy="zigzag", ) def tearDown(self): - init_cp_strategy(SimpleNamespace(enable_prefill_cp=False)) + init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="zigzag") def _metadata_for_rank(self, rank, *, cp_size, seq_lens, extend_seq_lens): strategy = ZigzagCPStrategy(cp_size=cp_size) @@ -809,16 +800,13 @@ def reference_attention( class TestCPInterleaveStrategy(CustomTestCase): def setUp(self): init_cp_strategy( - SimpleNamespace( - enable_prefill_cp=True, - cp_strategy="interleave", - attn_cp_size=4, - attention_backend="fa3", - ) + enable_prefill_cp=True, + cp_size=4, + cp_strategy="interleave", ) def tearDown(self): - init_cp_strategy(SimpleNamespace(enable_prefill_cp=False)) + init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="zigzag") def _metadata_for_rank(self, rank, *, cp_size, seq_lens, extend_seq_lens): strategy = InterleaveCPStrategy(cp_size=cp_size) diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index 6b840f97057a..4f01eb65122c 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -53,6 +53,10 @@ # The test below asserts this map is exactly the set of such reads, so the # reasons cannot drift away from the code. _CONFIGURED_SIZE_CALL_SITES = { + ("srt/layers/cp/base.py", "attn_cp_size"): ( + "the lazy strategy bind in a worker: the CP group is what the strategy " + "is being built for, and the configured width is what describes it" + ), ("srt/entrypoints/engine.py", "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" From 9ab252dff4f37d25a710c7ec639fd896e79f362d Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 19:45:14 +0000 Subject: [PATCH 16/26] config: the embedding plan reports the resolved configuration `resolved_embedding_plan` is the `/server_info` and gRPC readback of the embedding runtime knobs, and it read `cuda_graph_config`, `chunked_prefill_size`, `disable_radix_cache`, `is_embedding` and `prefill_only_disable_kv_cache` off the record -- which now holds the raw input, so the plan would have reported `None` for the graph config of a server running one. The two callers pass `resolving_view(record)`, and the parameter is `config` rather than `server_args`: the function's contract is "something that answers with the resolved configuration", which is why it was duck-typed to begin with. Five (file, field) pairs leave the exposure pin, which is down to four -- the launcher's pre-publish env setup and the auto-parser late resolution, both of which read the record because that is the only thing that exists at those points. --- .../srt/configs/embedding_model_spec.py | 21 ++++++++----------- python/sglang/srt/entrypoints/grpc_bridge.py | 3 ++- python/sglang/srt/entrypoints/http_server.py | 5 ++++- .../unit/configs/test_embedding_model_spec.py | 2 +- 4 files changed, 16 insertions(+), 15 deletions(-) diff --git a/python/sglang/srt/configs/embedding_model_spec.py b/python/sglang/srt/configs/embedding_model_spec.py index c529e19993de..a621a9d0d04f 100644 --- a/python/sglang/srt/configs/embedding_model_spec.py +++ b/python/sglang/srt/configs/embedding_model_spec.py @@ -221,17 +221,18 @@ def _native_embedding_spec( def resolved_embedding_plan( - spec: EmbeddingModelSpec, *, server_args: Any, model_config: Any + spec: EmbeddingModelSpec, *, config: Any, model_config: Any ) -> dict[str, Any]: """Combine static capabilities with the effective server configuration. This boundary deliberately accepts duck-typed arguments so the declarative registry remains independent of ServerArgs and ModelConfig import cycles. + `config` must answer with the *resolved* configuration -- the readback + callers pass `resolving_view(record)`, since the record's fields are the raw + input. """ - prefill_graph = getattr( - getattr(server_args, "cuda_graph_config", None), "prefill", None - ) + prefill_graph = getattr(getattr(config, "cuda_graph_config", None), "prefill", None) backend = getattr(prefill_graph, "backend", None) backend_value = getattr(backend, "value", backend) capture_sizes = getattr(prefill_graph, "bs", None) or [] @@ -239,7 +240,7 @@ def resolved_embedding_plan( return { **spec.as_dict(), - "enabled": bool(getattr(server_args, "is_embedding", False)), + "enabled": bool(getattr(config, "is_embedding", False)), "supports_dimensions": bool(getattr(model_config, "is_matryoshka", False)), "matryoshka_dimensions": list( getattr(model_config, "matryoshka_dimensions", None) or [] @@ -252,14 +253,10 @@ def resolved_embedding_plan( }, "cache": { "kv_cache_disabled": bool( - getattr(server_args, "prefill_only_disable_kv_cache", False) - ), - "radix_cache_disabled": bool( - getattr(server_args, "disable_radix_cache", False) + getattr(config, "prefill_only_disable_kv_cache", False) ), - "chunked_prefill_disabled": getattr( - server_args, "chunked_prefill_size", None - ) + "radix_cache_disabled": bool(getattr(config, "disable_radix_cache", False)), + "chunked_prefill_disabled": getattr(config, "chunked_prefill_size", None) == -1, }, } diff --git a/python/sglang/srt/entrypoints/grpc_bridge.py b/python/sglang/srt/entrypoints/grpc_bridge.py index 9349d3f92358..4adf1d737091 100644 --- a/python/sglang/srt/entrypoints/grpc_bridge.py +++ b/python/sglang/srt/entrypoints/grpc_bridge.py @@ -14,6 +14,7 @@ from pydantic import ValidationError +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.configs.embedding_model_spec import resolved_embedding_plan from sglang.srt.runtime_context import ( get_lora, @@ -416,7 +417,7 @@ def get_model_info(self) -> str: if embedding_model_spec is not None: result["embedding"] = resolved_embedding_plan( embedding_model_spec, - server_args=self.server_args, + config=resolving_view(self.server_args), model_config=model_config, ) return json.dumps(result, default=str) diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 13817d468720..e992b9af3868 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -63,6 +63,7 @@ from fastapi.responses import ORJSONResponse, Response, StreamingResponse from fastapi.routing import APIRoute +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.configs.embedding_model_spec import resolved_embedding_plan from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode @@ -771,7 +772,9 @@ async def model_info(): if embedding_model_spec is not None: result["embedding"] = resolved_embedding_plan( embedding_model_spec, - server_args=_global_state.tokenizer_manager.server_args, + # Through the declarations: the plan reports the *effective* config, + # and the fields hold the raw input. + config=resolving_view(_global_state.tokenizer_manager.server_args), model_config=model_config, ) return result diff --git a/test/registered/unit/configs/test_embedding_model_spec.py b/test/registered/unit/configs/test_embedding_model_spec.py index 6e1d5018e39c..bd8850c4b9ed 100644 --- a/test/registered/unit/configs/test_embedding_model_spec.py +++ b/test/registered/unit/configs/test_embedding_model_spec.py @@ -71,7 +71,7 @@ def test_resolved_plan_reports_effective_runtime_knobs(self): ) plan = resolved_embedding_plan( spec, - server_args=SimpleNamespace( + config=SimpleNamespace( is_embedding=True, cuda_graph_config=SimpleNamespace( prefill=SimpleNamespace( From 4cf15d7d3b5f2b487dfcaaf48c8683ae93aba11c Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 19:46:24 +0000 Subject: [PATCH 17/26] config: say why the four remaining supplied-instance reads stay All four read the record because the record is the only thing that exists where they run: the NCCL environment setup and the auto-parser late resolution both happen in the launcher before the publish. Written down next to the pin so the next person does not have to re-derive it, and so a new entry has to come with the same kind of reason. --- .../test_supplied_instance_exposure_ratchet.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index 23e863670f17..bee5673ed539 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -134,9 +134,18 @@ # are step-12 exposure like any other pair. _PASSED = frozenset({"model_path", "device", "random_seed"}) -# The reads that still take a value off the supplied instance. `initialize_moe_config` -# is handed the record until the replay goes away; the rest are pre-publish launcher -# reads. +# What is left reads the record because the record is the only thing that exists +# at that point in the process. Everything else moved to the bags or to +# `resolving_view`; a new entry here needs the same kind of reason. +# +# engine.py / enable_symm_mem +# `_set_envs_and_config` sets NCCL environment variables before anything +# publishes -- there is no bag to read yet. +# engine.py / reasoning_parser, tool_call_parser +# template_detection.py / model_path +# the auto-parser detection is late resolution: it runs in the launcher's +# validation stage, decides these fields and writes them *in place*, +# before the publish that would give it a bag. _EXPOSED = { ("disaggregation/encoder/server.py", "model_loader_extra_config"), ("layers/moe/utils.py", "deepep_mode"), From d2f7bdd543f05cf695472347837ae90ba8485710 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 20:09:16 +0000 Subject: [PATCH 18/26] config: the encoder, HiCache and metrics readers take the bags The biggest remaining clusters of "read a field off a handed record", all of them past their process entry's publish: * the encoder's five modules read `get_parallel().config.tp_size` and `get_serving().host` / `.port` -- `runtime.py` was already mixing `get_parallel().config.dp_size` with `server_args.tp_size` on one line; * `get_allocator_type()` reads the two HiCache leaves and takes no config, which drops the parameter from `_get_allocator_type` and its 14 call sites in the hybrid pool assembler; * the unified radix cache's write and prefetch policies, the detokenizer and gRPC metrics flags, the XPU and runner-backend memory-saver checks, the CP DSA split, the FP4 GEMM backend and the EP redundant-expert count. Both exposure pins are now at their floor: four entries in `_EXPOSED` and three in `_OVERRIDDEN_AND_READ`, all of them the launcher's pre-publish env setup and the auto-parser late resolution. --- .../srt/disaggregation/encoder/grpc_server.py | 15 ++++++--- .../srt/disaggregation/encoder/http_server.py | 12 +++---- .../disaggregation/encoder/preprocessor.py | 6 ++-- .../srt/disaggregation/encoder/receiver.py | 12 ++++--- .../srt/disaggregation/encoder/runtime.py | 4 +-- .../srt/disaggregation/encoder/server.py | 15 +++++---- python/sglang/srt/entrypoints/grpc_server.py | 13 ++++++-- .../xpu/graph_runner/xpu_graph_runner.py | 3 +- python/sglang/srt/layers/cp/utils.py | 2 +- .../srt/layers/quantization/fp4_utils.py | 3 +- .../srt/managers/detokenizer_manager.py | 12 +++++-- python/sglang/srt/mem_cache/hiradix_cache.py | 2 +- .../hybrid_cache/hybrid_pool_assembler.py | 32 +++++++++---------- .../sglang/srt/mem_cache/pool_host/common.py | 11 ++++--- .../srt/mem_cache/unified_radix_cache.py | 9 +++--- .../model_runner_components/moe_ep_setup.py | 3 +- .../model_executor/runner_backend/utils.py | 6 ++-- .../unit/test_global_config_read_ratchet.py | 13 ++++++++ ...test_supplied_instance_exposure_ratchet.py | 19 ----------- 19 files changed, 109 insertions(+), 83 deletions(-) diff --git a/python/sglang/srt/disaggregation/encoder/grpc_server.py b/python/sglang/srt/disaggregation/encoder/grpc_server.py index 91024e8594ef..58004e0eb918 100644 --- a/python/sglang/srt/disaggregation/encoder/grpc_server.py +++ b/python/sglang/srt/disaggregation/encoder/grpc_server.py @@ -24,7 +24,12 @@ from sglang.srt.disaggregation.encoder.server import MMEncoder, launch_encoder from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle from sglang.srt.managers.schedule_batch import Modality -from sglang.srt.runtime_context import get_disagg, publish +from sglang.srt.runtime_context import ( + get_disagg, + get_parallel, + get_serving, + publish, +) from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import random_uuid from sglang.srt.utils.network import NetworkAddress, get_zmq_socket @@ -212,11 +217,11 @@ async def serve_grpc_encoder(server_args: ServerArgs): 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() send_sockets: List[zmq.Socket] = [] - for rank in range(1, server_args.tp_size): + for rank in range(1, get_parallel().config.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) @@ -254,7 +259,9 @@ async def serve_grpc_encoder(server_args: ServerArgs): ) reflection.enable_server_reflection(SERVICE_NAMES, server) - listen_addr = NetworkAddress(server_args.host, server_args.port).to_host_port_str() + listen_addr = NetworkAddress( + get_serving().host, get_serving().port + ).to_host_port_str() server.add_insecure_port(listen_addr) await server.start() diff --git a/python/sglang/srt/disaggregation/encoder/http_server.py b/python/sglang/srt/disaggregation/encoder/http_server.py index e204048c69ce..4970e3d49258 100644 --- a/python/sglang/srt/disaggregation/encoder/http_server.py +++ b/python/sglang/srt/disaggregation/encoder/http_server.py @@ -113,11 +113,11 @@ def _register_encoder_url_with_bootstrap(server_args: ServerArgs): instead of serialising sleeps in a single thread. """ - host = server_args.host + host = get_serving().host if not host or host in ("0.0.0.0", "::"): - host = get_local_ip_auto(server_args.host) + host = get_local_ip_auto(get_serving().host) scheme = "https" if server_args.ssl_certfile else "http" - encoder_url = NetworkAddress(host, server_args.port).to_url(scheme) + encoder_url = NetworkAddress(host, get_serving().port).to_url(scheme) payload = {"url": encoder_url} bootstrap_urls = list(server_args.encoder_register_urls) if not bootstrap_urls: @@ -174,11 +174,11 @@ def _worker(): def _unregister_encoder_url_from_bootstrap(server_args: ServerArgs): - host = server_args.host + host = get_serving().host if not host or host in ("0.0.0.0", "::"): - host = get_local_ip_auto(server_args.host) + host = get_local_ip_auto(get_serving().host) scheme = "https" if server_args.ssl_certfile else "http" - encoder_url = NetworkAddress(host, server_args.port).to_url(scheme) + encoder_url = NetworkAddress(host, get_serving().port).to_url(scheme) payload = {"url": encoder_url} for bootstrap_url in server_args.encoder_register_urls: diff --git a/python/sglang/srt/disaggregation/encoder/preprocessor.py b/python/sglang/srt/disaggregation/encoder/preprocessor.py index ffa48a8b493a..bd6d3f06e847 100644 --- a/python/sglang/srt/disaggregation/encoder/preprocessor.py +++ b/python/sglang/srt/disaggregation/encoder/preprocessor.py @@ -167,7 +167,7 @@ def _load_mm_processor(self, server_args: ServerArgs): try: self.image_processor = AutoImageProcessor.from_pretrained( get_serving().tokenizer_path or get_model().model_path, - trust_remote_code=server_args.trust_remote_code, + trust_remote_code=get_model().trust_remote_code, revision=server_args.revision, **image_processor_kwargs, ) @@ -178,7 +178,7 @@ def _load_mm_processor(self, server_args: ServerArgs): try: self.video_processor = AutoVideoProcessor.from_pretrained( get_serving().tokenizer_path or get_model().model_path, - trust_remote_code=server_args.trust_remote_code, + trust_remote_code=get_model().trust_remote_code, revision=server_args.revision, ) except Exception as e: @@ -188,7 +188,7 @@ def _load_mm_processor(self, server_args: ServerArgs): try: _audio_proc = AutoProcessor.from_pretrained( get_serving().tokenizer_path or get_model().model_path, - trust_remote_code=server_args.trust_remote_code, + trust_remote_code=get_model().trust_remote_code, revision=server_args.revision, ) if not hasattr(_audio_proc, "feature_extractor"): diff --git a/python/sglang/srt/disaggregation/encoder/receiver.py b/python/sglang/srt/disaggregation/encoder/receiver.py index 9e2d99a755f8..ace0bcd8da19 100644 --- a/python/sglang/srt/disaggregation/encoder/receiver.py +++ b/python/sglang/srt/disaggregation/encoder/receiver.py @@ -39,6 +39,8 @@ get_disagg, get_exec, get_mm, + get_model, + get_parallel, get_serving, ) from sglang.srt.server_args import ServerArgs @@ -1726,13 +1728,13 @@ def __init__( # When None (e.g. in a scheduler subprocess that has no in-process # bootstrap), fall back to a snapshot of the static --encoder-urls. self.encode_urls: List[str] = ( - encode_urls if encode_urls is not None else list(server_args.encoder_urls) + encode_urls if encode_urls is not None else list(get_disagg().encoder_urls) ) self.recv_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get() - self.host = get_local_ip_auto(server_args.host) + self.host = get_local_ip_auto(get_serving().host) self.pp_rank = pp_rank self.tp_rank = tp_rank - self.tp_size = server_args.tp_size + self.tp_size = get_parallel().config.tp_size self.tp_group = tp_group self.nnodes = server_args.nnodes self.hostname = get_local_ip_auto() @@ -1841,7 +1843,7 @@ def _init_mm_processor( _processor = get_processor( get_serving().tokenizer_path, tokenizer_mode=server_args.tokenizer_mode, - trust_remote_code=server_args.trust_remote_code, + trust_remote_code=get_model().trust_remote_code, revision=server_args.revision, image_processor_backend=resolve_image_processor_backend(get_mm()), **extra_kwargs, @@ -2664,7 +2666,7 @@ def create_mm_receiver( transport_mode = envs.SGLANG_ENCODER_MM_RECEIVER_MODE.get() logger.debug(f"MMReceiver transport_mode from env: {transport_mode}") - _validate_transport_mode(transport_mode, encode_urls or server_args.encoder_urls) + _validate_transport_mode(transport_mode, encode_urls or get_disagg().encoder_urls) logger.info(f"EPD MMReceiver: using transport_mode={transport_mode}") receiver_cls = _MM_RECEIVER_BY_MODE.get(transport_mode) diff --git a/python/sglang/srt/disaggregation/encoder/runtime.py b/python/sglang/srt/disaggregation/encoder/runtime.py index 433f0e9fc73d..a4d88cd9ea9a 100644 --- a/python/sglang/srt/disaggregation/encoder/runtime.py +++ b/python/sglang/srt/disaggregation/encoder/runtime.py @@ -1572,10 +1572,10 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher: HTTP uses this entry point today. gRPC can reuse it later without importing HTTP application state or Uvicorn. """ - if get_parallel().config.dp_size <= 1 or server_args.tp_size != 1: + if get_parallel().config.dp_size <= 1 or get_parallel().config.tp_size != 1: raise ValueError( "Encoder DP mode requires --dp-size > 1 and --tp-size 1; got " - f"dp_size={get_parallel().config.dp_size}, tp_size={server_args.tp_size}." + f"dp_size={get_parallel().config.dp_size}, tp_size={get_parallel().config.tp_size}." ) dp_size = get_parallel().config.dp_size logger.info(f"Launching encoder in DP mode: dp_size={dp_size}") diff --git a/python/sglang/srt/disaggregation/encoder/server.py b/python/sglang/srt/disaggregation/encoder/server.py index cef8e884ba9d..34b1386db723 100644 --- a/python/sglang/srt/disaggregation/encoder/server.py +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -63,6 +63,7 @@ get_exec, get_mm, get_model, + get_parallel, publish, ) from sglang.srt.server_args import ServerArgs @@ -450,7 +451,7 @@ def __init__( this instance's value, not a config change, so it travels as an argument.""" assert_published(server_args, role="encoder") - logger.info(f"init MMEncoder {rank}/{server_args.tp_size}") + logger.info(f"init MMEncoder {rank}/{get_parallel().config.tp_size}") self.server_args = server_args configure_media_url_security( get_mm().allowed_media_domains, @@ -470,7 +471,7 @@ def __init__( self.load_config = LoadConfig( load_format=get_model().load_format, download_dir=server_args.download_dir, - model_loader_extra_config=server_args.model_loader_extra_config, + model_loader_extra_config=get_model().model_loader_extra_config, remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip, remote_instance_weight_loader_seed_instance_service_port=server_args.remote_instance_weight_loader_seed_instance_service_port, remote_instance_weight_loader_send_weights_group_ports=server_args.remote_instance_weight_loader_send_weights_group_ports, @@ -491,12 +492,14 @@ def __init__( init_distributed_environment( backend=get_default_distributed_backend(self.device), - world_size=server_args.tp_size, + world_size=get_parallel().config.tp_size, rank=rank, distributed_init_method=dist_init_method, local_rank=rank, ) - initialize_model_parallel(tensor_model_parallel_size=server_args.tp_size) + initialize_model_parallel( + tensor_model_parallel_size=get_parallel().config.tp_size + ) initialize_dp_attention(server_args, self.model_config) self.model = load_model( @@ -554,7 +557,7 @@ def __init__( ) self.mm_global_cache = EmbeddingCacheController( rank, - server_args.tp_size, + get_parallel().config.tp_size, embedding_store=embedding_store, hidden_dims=self._embedding_dims, tp_group=get_tp_group().cpu_group, @@ -1032,7 +1035,7 @@ async def _prepare_encode_context( ) def _broadcast_global_cache_mask(self, mask_tensor: torch.Tensor): - if self.server_args.tp_size > 1: + if get_parallel().config.tp_size > 1: torch.distributed.broadcast( mask_tensor, src=0, diff --git a/python/sglang/srt/entrypoints/grpc_server.py b/python/sglang/srt/entrypoints/grpc_server.py index c3e22762c76f..b78e24341137 100644 --- a/python/sglang/srt/entrypoints/grpc_server.py +++ b/python/sglang/srt/entrypoints/grpc_server.py @@ -165,6 +165,15 @@ async def serve_grpc(server_args, model_info=None): "version mismatch — see the chained exception above for details." ) from e + from sglang.srt.arg_groups.overrides import resolving_view + + # The integrated servicer builds an `Engine`, which validates and publishes + # on its own. Validating here would run `check_server_args` twice, and the + # LoRA normalization is not idempotent -- the second pass sees the `LoRARef` + # objects the first one declared and rejects them. So this entry reads the + # declarations for what it needs before the engine exists. + cfg = resolving_view(server_args) + sidecar_app = web.Application() sidecar_runner = None sidecar_port = ( @@ -176,7 +185,7 @@ async def serve_grpc(server_args, model_info=None): # Metrics setup: must set PROMETHEUS_MULTIPROC_DIR before scheduler # processes import prometheus_client, since the env var is inherited # at fork time. - if server_args.enable_metrics: + if cfg.enable_metrics: try: from sglang.srt.observability.func_timer import enable_func_timer from sglang.srt.utils import set_prometheus_multiproc_dir @@ -232,7 +241,7 @@ async def _on_request_manager_ready(request_manager, srv_args, sched_info): ) if sidecar_supported: serve_kwargs["on_request_manager_ready"] = _on_request_manager_ready - elif server_args.enable_metrics: + elif cfg.enable_metrics: # User explicitly asked for metrics but the installed servicer can't # start the sidecar that serves them — fail loud rather than silently # produce a server with no /metrics endpoint. diff --git a/python/sglang/srt/hardware_backend/xpu/graph_runner/xpu_graph_runner.py b/python/sglang/srt/hardware_backend/xpu/graph_runner/xpu_graph_runner.py index b8c7fa592c03..bcc8b4e62292 100644 --- a/python/sglang/srt/hardware_backend/xpu/graph_runner/xpu_graph_runner.py +++ b/python/sglang/srt/hardware_backend/xpu/graph_runner/xpu_graph_runner.py @@ -22,6 +22,7 @@ from torch.profiler import ProfilerActivity, profile from sglang.srt.model_executor.runner import DecodeCudaGraphRunner +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import register_xpu_device_properties_for_dynamo logger = logging.getLogger(__name__) @@ -116,7 +117,7 @@ def _apply_xpu_compile_config() -> None: def __init__(self, model_runner: ModelRunner): assert ( - not model_runner.server_args.enable_memory_saver + not get_exec().features.enable_memory_saver ), "XPUGraphRunner does not support Torch Memory Saver yet." register_fake_ops() self._apply_xpu_compile_config() diff --git a/python/sglang/srt/layers/cp/utils.py b/python/sglang/srt/layers/cp/utils.py index b397c4fed19e..be2ff1757fdc 100644 --- a/python/sglang/srt/layers/cp/utils.py +++ b/python/sglang/srt/layers/cp/utils.py @@ -63,7 +63,7 @@ def is_glm_dsa_cache_layer_split_enabled(model_runner: "ModelRunner") -> bool: return ( not model_runner.is_draft_worker - and model_runner.server_args.enable_dsa_cache_layer_split + and get_parallel().config.enable_dsa_cache_layer_split and model_runner.use_mla_backend and is_deepseek_dsa(model_runner.model_config.hf_config) ) diff --git a/python/sglang/srt/layers/quantization/fp4_utils.py b/python/sglang/srt/layers/quantization/fp4_utils.py index 1b5858d71fc7..5cb7f097adef 100644 --- a/python/sglang/srt/layers/quantization/fp4_utils.py +++ b/python/sglang/srt/layers/quantization/fp4_utils.py @@ -6,6 +6,7 @@ import torch +from sglang.srt.runtime_context import get_exec from sglang.srt.utils.common import ( get_device_capability, is_cuda, @@ -146,7 +147,7 @@ def initialize_fp4_gemm_config(server_args: ServerArgs) -> None: """Initialize FP4 GEMM configuration from server args.""" global FP4_GEMM_RUNNER_BACKEND - backend = server_args.fp4_gemm_runner_backend + backend = get_exec().kernel.fp4_gemm_runner_backend if backend == "auto": if is_sm100_supported(): backend = "flashinfer_cutedsl" diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index 13d65d612077..485e33cd972f 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -39,7 +39,13 @@ ) 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 get_device, get_serving, publish +from sglang.srt.runtime_context import ( + get_device, + get_model, + get_observability, + 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 @@ -130,7 +136,7 @@ def init_tokenizer(self, server_args: ServerArgs): self.tokenizer = get_tokenizer( get_serving().tokenizer_path, tokenizer_mode=server_args.tokenizer_mode, - trust_remote_code=server_args.trust_remote_code, + trust_remote_code=get_model().trust_remote_code, revision=server_args.revision, tokenizer_backend=server_args.tokenizer_backend, ) @@ -151,7 +157,7 @@ def init_running_status(self, server_args: ServerArgs): test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(), ) - if server_args.enable_metrics: + if get_observability().enable_metrics: start_cpu_monitor_thread("detokenizer") def init_request_dispatcher(self): diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index c91fe88af441..3435712254f9 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -86,7 +86,7 @@ def __init__(self, params: CacheInitParams, server_args: ServerArgs): self.page_size = params.page_size self.kv_cache = params.token_to_kv_pool_allocator.get_kvcache() - allocator_type = get_allocator_type(server_args) + allocator_type = get_allocator_type() if isinstance(self.kv_cache, MHATokenToKVPool): self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)( diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index e29de3a0152a..2b38326b0111 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -44,8 +44,8 @@ logger = logging.getLogger(__name__) -def _get_allocator_type(server_args: ServerArgs) -> str: - return get_allocator_type(server_args) +def _get_allocator_type() -> str: + return get_allocator_type() def _evict_swa_for_device_alloc(cache: UnifiedRadixCache, required_size: int) -> None: @@ -126,7 +126,7 @@ def build_kv_host_pool( get_memory().hicache_size if host_size is None else host_size, page_size, get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), pool_label=pool_label, **kwargs, ) @@ -540,7 +540,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=swa_num_host_pages, slot_page_size=kvcache.swa_page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator entries.append( @@ -567,7 +567,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=num_host_pages, slot_page_size=page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) c4_indexer_host_pool = DeepSeekV4PagedHostPool( pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER), @@ -579,7 +579,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=num_host_pages, slot_page_size=page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) entries.extend( [ @@ -610,7 +610,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=swa_num_host_pages, swa_page_size=kvcache.swa_page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) c4_indexer_state_host_pool = DeepSeekV4StateHostPool( pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER_STATE), @@ -621,7 +621,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=swa_num_host_pages, swa_page_size=kvcache.swa_page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) entries.extend( [ @@ -653,7 +653,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=num_host_pages, slot_page_size=page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) # C128 state pool is intentionally not registered with hicache. # page_size=256 % 128 == 0, so state pool is not consumed on load. @@ -739,7 +739,7 @@ def build_hybrid_mamba_stack( mamba_pool, get_memory().hicache_ratio, mamba_host_size, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), layout=get_memory().hicache_mem_layout, ) entries = [ @@ -1038,7 +1038,7 @@ def build_full_draft_pools( host_to_device_ratio=host_pool_group.logical_size / pool.size, page_size=controller.page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), pool_label="draft", ) draft_layer_mapping = {i: i for i in range(pool.layer_num)} @@ -1064,7 +1064,7 @@ def build_full_draft_pools( pool, draft_host_pool, get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) specs.append( SidecarPoolSpec( @@ -1111,7 +1111,7 @@ def build_swa_draft_pools( num_host_pages=target_swa_host_pool.num_host_pages, slot_page_size=draft_swa_pool.page_size, layout=target_swa_host_pool.layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) else: host_pool = _build_mha_mla_host_pool( @@ -1119,7 +1119,7 @@ def build_swa_draft_pools( host_to_device_ratio=target_swa_host_pool.size / draft_swa_pool.size, page_size=target_swa_host_pool.page_size, layout=target_swa_host_pool.layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), pool_label="draft_swa", ) @@ -1513,7 +1513,7 @@ def build( full_kv_pool, kv_host_pool, get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ), prefetch_threshold=prefetch_threshold, model_name=model_name, @@ -1951,7 +1951,7 @@ def attach_hybrid_dsa_pool_to_hiradix_cache( kv, kv_host_pool, get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ), model_name=get_serving().served_model_name, storage_backend_extra_config=extra_config, diff --git a/python/sglang/srt/mem_cache/pool_host/common.py b/python/sglang/srt/mem_cache/pool_host/common.py index 7701356542e7..d781c20958b8 100644 --- a/python/sglang/srt/mem_cache/pool_host/common.py +++ b/python/sglang/srt/mem_cache/pool_host/common.py @@ -99,14 +99,15 @@ def get_allocator_from_storage(allocator_type): return HostTensorAllocator() -def get_allocator_type(server_args) -> str: - backend = getattr(server_args, "hicache_storage_backend", None) +def get_allocator_type() -> str: + """The host-allocator kind the published HiCache configuration asks for.""" + from sglang.srt.runtime_context import get_memory + + backend = get_memory().hicache_storage_backend if backend == "shm": return "shm" if backend == "dynamic": - extra_config_str = getattr( - server_args, "hicache_storage_backend_extra_config", None - ) + extra_config_str = get_memory().hicache_storage_backend_extra_config if extra_config_str: try: config = json.loads(extra_config_str) diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 4f9bd12403c3..1a7429e84bdd 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -81,6 +81,7 @@ StorageMetrics, StorageMetricsCollector, ) +from sglang.srt.runtime_context import get_memory from sglang.srt.session.streaming_session import StreamingSession from sglang.srt.utils.common import ceil_align @@ -385,7 +386,7 @@ def init_hicache(self, server_args: ServerArgs, params: CacheInitParams) -> None self.extra_metric_labels = server_args.extra_metric_labels # Parse storage config once, share with assembler and tree - storage_backend = server_args.hicache_storage_backend + storage_backend = get_memory().hicache_storage_backend storage_extra_config = None storage_prefetch_threshold = 256 prefetch_timeout_base = 1.0 @@ -399,7 +400,7 @@ def init_hicache(self, server_args: ServerArgs, params: CacheInitParams) -> None prefetch_timeout_per_ki_token, hicache_storage_pass_prefix_keys, ) = HybridCacheController.parse_storage_backend_extra_config( - server_args.hicache_storage_backend_extra_config + get_memory().hicache_storage_backend_extra_config ) attach_hybrid_pool_to_unified_cache( @@ -442,7 +443,7 @@ def init_hicache(self, server_args: ServerArgs, params: CacheInitParams) -> None # State initialization self.write_through_threshold = ( - 1 if server_args.hicache_write_policy == "write_through" else 2 + 1 if get_memory().hicache_write_policy == "write_through" else 2 ) self.is_write_back = ( self.cache_controller is not None @@ -457,7 +458,7 @@ def init_hicache(self, server_args: ServerArgs, params: CacheInitParams) -> None pool=_COMPONENT_POOL_LABEL[ct], ) self.load_back_threshold = 10 - self.prefetch_stop_policy = server_args.hicache_storage_prefetch_policy + self.prefetch_stop_policy = get_memory().hicache_storage_prefetch_policy # Runtime attach/detach of the L3 backend (startup, admin API, atexit). self._storage_attachment = StorageAttachment(self) diff --git a/python/sglang/srt/model_executor/model_runner_components/moe_ep_setup.py b/python/sglang/srt/model_executor/model_runner_components/moe_ep_setup.py index 27c9951cc65b..25cdea784111 100644 --- a/python/sglang/srt/model_executor/model_runner_components/moe_ep_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/moe_ep_setup.py @@ -13,6 +13,7 @@ ) from sglang.srt.layers.moe.hash_topk import HashTopK from sglang.srt.layers.moe.topk import TopK +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import get_bool_env_var, is_hip, log_info_on_rank0 if TYPE_CHECKING: @@ -54,7 +55,7 @@ def prepare_moe_topk( # Redundant experts therefore need to be included in the per-rank # expert count used for Waterfill's shared-expert slot remapping. num_physical_routed_experts = ( - num_routed_experts + server_args.ep_num_redundant_experts + num_routed_experts + get_exec().moe.ep_num_redundant_experts ) if isinstance(module, TopK): routed_scaling_factor = module.topk_config.routed_scaling_factor diff --git a/python/sglang/srt/model_executor/runner_backend/utils.py b/python/sglang/srt/model_executor/runner_backend/utils.py index a07ec247ac71..63c98b45c6ba 100644 --- a/python/sglang/srt/model_executor/runner_backend/utils.py +++ b/python/sglang/srt/model_executor/runner_backend/utils.py @@ -64,7 +64,7 @@ def resolve_decode_backend( cfg = get_exec().graph.cuda_graph_config backend_name = cfg.decode.backend if cfg is not None else Backend.FULL - enable_memory_saver = model_runner.server_args.enable_memory_saver + enable_memory_saver = get_exec().features.enable_memory_saver if model_runner.device == "npu": from sglang.srt.hardware_backend.npu.graph_runner.npu_cudagraph_backend import ( @@ -115,13 +115,13 @@ def resolve_prefill_backend( if backend_name == Backend.BREAKABLE: return BreakableCudaGraphBackend( cuda_graph_runner, - enable_memory_saver=model_runner.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, debug_eager=get_exec().graph.debug_cuda_graph, ) if backend_name == Backend.FULL: return FullCudaGraphBackend( cuda_graph_runner, - enable_memory_saver=model_runner.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, ) # Default: tc_piecewise. return TcPiecewiseCudaGraphBackend(cuda_graph_runner) diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index 4f01eb65122c..4b3f97992245 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -216,6 +216,19 @@ "the encode server's launch entry sizes its workers before it has " "spawned any of them" ), + ("srt/disaggregation/encoder/grpc_server.py", "tp_size"): ( + "the same worker-count arithmetic on the gRPC entry: it spawns the TP " + "workers, so their groups do not exist yet" + ), + ("srt/disaggregation/encoder/server.py", "tp_size"): ( + "`MMEncoder` builds its own TP group from this size -- " + "`initialize_model_parallel` is the call being handed it, so there is " + "nothing live to ask" + ), + ("srt/disaggregation/encoder/receiver.py", "tp_size"): ( + "the receiver labels and shards by the launch width; it runs in the " + "tokenizer process, which holds no encoder groups" + ), ("srt/utils/common.py", "tp_size"): ( "the require_*_tp_gather predicates compared the configured tp_size " "when they read the record; the live property answers a different " diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index bee5673ed539..50d8cc893f01 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -147,31 +147,20 @@ # validation stage, decides these fields and writes them *in place*, # before the publish that would give it a bag. _EXPOSED = { - ("disaggregation/encoder/server.py", "model_loader_extra_config"), ("layers/moe/utils.py", "deepep_mode"), ("layers/moe/utils.py", "disable_shared_experts_fusion"), ("layers/moe/utils.py", "moe_a2a_backend"), ("layers/moe/utils.py", "moe_runner_backend"), ("layers/moe/utils.py", "quantization"), ("layers/moe/utils.py", "speculative_moe_runner_backend"), - ("configs/embedding_model_spec.py", "chunked_prefill_size"), - ("configs/embedding_model_spec.py", "cuda_graph_config"), - ("configs/embedding_model_spec.py", "disable_radix_cache"), - ("configs/embedding_model_spec.py", "is_embedding"), - ("configs/embedding_model_spec.py", "prefill_only_disable_kv_cache"), ("entrypoints/engine.py", "enable_symm_mem"), ("entrypoints/engine.py", "reasoning_parser"), ("entrypoints/engine.py", "tool_call_parser"), - ("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"), ("layers/moe/utils.py", "deepep_mode"), ("layers/moe/utils.py", "moe_a2a_backend"), ("layers/moe/utils.py", "moe_runner_backend"), ("layers/moe/utils.py", "quantization"), ("layers/moe/utils.py", "speculative_moe_runner_backend"), - ("speculative/draft_worker_common.py", "speculative_draft_attention_backend"), - ("utils/common.py", "speculative_num_draft_tokens"), - ("utils/common.py", "speculative_num_steps"), - ("utils/hf_transformers/processor.py", "image_processor_backend"), ("weight_cache/daemon.py", "attn_cp_size"), ("weight_cache/daemon.py", "deepep_mode"), ("weight_cache/daemon.py", "dp_size"), @@ -208,14 +197,6 @@ ("weight_cache/daemon.py", "ep_size"), ("weight_cache/daemon.py", "load_format"), ("weight_cache/daemon.py", "model_path"), - ("mem_cache/pool_host/common.py", "hicache_storage_backend"), - ("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"), - ("mem_cache/unified_radix_cache.py", "hicache_storage_backend"), - ("mem_cache/unified_radix_cache.py", "hicache_storage_backend_extra_config"), - ("mem_cache/unified_radix_cache.py", "hicache_storage_prefetch_policy"), - ("mem_cache/unified_radix_cache.py", "hicache_write_policy"), - ("utils/common.py", "speculative_num_draft_tokens"), - ("utils/common.py", "speculative_num_steps"), ("weight_cache/daemon.py", "dp_size"), ("weight_cache/daemon.py", "dtype"), ("weight_cache/daemon.py", "ep_size"), From 2377f2f8a5894f4de81e2bb45f002fad6e8845b4 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 20:16:45 +0000 Subject: [PATCH 19/26] config: the HTTP entry reads the serving and observability bags `_setup_and_run_http_server` runs after `Engine._launch_subprocesses` has published, so its host, port, log level and metrics flag come from `get_serving()` / `get_observability()` -- 28 reads, including the two `enable_metrics` gates in the app setup. The DP controller keeps its record reads: its declared namespace set is `{exec, parallel, device, disagg}`, so reading `serving` or `observability` there would be refused under `SGLANG_ROLE_NAMESPACES=enforce`. Narrowing that set was the point of declaring it, and widening it to move a `host` read is the wrong trade. --- python/sglang/srt/entrypoints/http_server.py | 48 +++++++++++-------- .../unit/server_args/test_server_args.py | 13 +++-- 2 files changed, 36 insertions(+), 25 deletions(-) diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index e992b9af3868..be7a5943c83c 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -285,7 +285,7 @@ async def lifespan(fast_api_app: FastAPI): thread_label = f"MultiTokenizer-{_global_state.tokenizer_manager.worker_id}" # Add prometheus middleware - if server_args.enable_metrics: + if get_observability().enable_metrics: add_prometheus_middleware(app) enable_func_timer() @@ -489,6 +489,7 @@ async def custom_handler(request: Request): get_exec, get_lora, get_model, + get_observability, get_parallel, get_serving, publish, @@ -2543,7 +2544,7 @@ def _setup_and_run_http_server( if tokenizer_manager is not None: tokenizer_manager._subprocess_watchdog = subprocess_watchdog - if server_args.enable_metrics: + if get_observability().enable_metrics: add_prometheus_track_response_middleware(app) # Pass additional arguments to the lifespan function. @@ -2605,12 +2606,13 @@ def _setup_and_run_http_server( if server_args.enable_http2: logger.info( f"Starting embedded Granian HTTP/2 server on " - f"{server_args.host}:{server_args.port}" + f"{get_serving().host}:{get_serving().port}" ) _run_granian_server( - host=server_args.host, - port=server_args.port, - log_level=server_args.log_level_http or server_args.log_level, + host=get_serving().host, + port=get_serving().port, + log_level=get_observability().log_level_http + or get_observability().log_level, http2_max_concurrent_streams=( server_args.http2_max_concurrent_streams ), @@ -2624,10 +2626,11 @@ def _setup_and_run_http_server( # Use Config/Server API for access to the SSLContext. config = uvicorn.Config( app, - host=server_args.host, - port=server_args.port, + host=get_serving().host, + port=get_serving().port, root_path=server_args.fastapi_root_path, - log_level=server_args.log_level_http or server_args.log_level, + log_level=get_observability().log_level_http + or get_observability().log_level, timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(), loop="uvloop", ssl_keyfile=server_args.ssl_keyfile, @@ -2661,10 +2664,11 @@ async def _run_with_ssl_refresh(): # Default case, one tokenizer process uvicorn.run( app, - host=server_args.host, - port=server_args.port, + host=get_serving().host, + port=get_serving().port, root_path=server_args.fastapi_root_path, - log_level=server_args.log_level_http or server_args.log_level, + log_level=get_observability().log_level_http + or get_observability().log_level, timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(), loop="uvloop", ssl_keyfile=server_args.ssl_keyfile, @@ -2692,12 +2696,13 @@ async def _run_with_ssl_refresh(): if server_args.enable_http2: logger.info( f"Starting embedded Granian HTTP/2 server on " - f"{server_args.host}:{server_args.port}" + f"{get_serving().host}:{get_serving().port}" ) _run_granian_server( - host=server_args.host, - port=server_args.port, - log_level=server_args.log_level_http or server_args.log_level, + host=get_serving().host, + port=get_serving().port, + log_level=get_observability().log_level_http + or get_observability().log_level, http2_max_concurrent_streams=( server_args.http2_max_concurrent_streams ), @@ -2710,10 +2715,11 @@ async def _run_with_ssl_refresh(): else: uvicorn.run( "sglang.srt.entrypoints.http_server:app", - host=server_args.host, - port=server_args.port, + host=get_serving().host, + port=get_serving().port, root_path=server_args.fastapi_root_path, - log_level=server_args.log_level_http or server_args.log_level, + log_level=get_observability().log_level_http + or get_observability().log_level, timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(), timeout_worker_healthcheck=envs.SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT.get(), loop="uvloop", @@ -2751,12 +2757,12 @@ def _start_native_grpc_server_for_runtime( ) grpc_handle = grpc_native.start_server( - host=server_args.host, + host=get_serving().host, 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}:{grpc_port}") + logger.info(f"Native gRPC server started on {get_serving().host}:{grpc_port}") return grpc_handle diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 234d268f52eb..fc2f4c4f5697 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -27,6 +27,7 @@ Phase, PhaseConfig, ) +from sglang.srt.runtime_context import get_context, get_serving from sglang.srt.server_args import PortArgs, ServerArgs, prepare_server_args from sglang.srt.server_args_config_parser import ConfigArgumentMerger from sglang.test.ci.ci_register import register_cpu_ci @@ -2357,9 +2358,13 @@ def test_start_server_call_site_matches_native_signature(self): fake_core = SimpleNamespace(start_server=MagicMock(return_value="handle")) fake_bridge = SimpleNamespace(RuntimeHandle=MagicMock(return_value="rt")) - server_args = SimpleNamespace( - host="127.0.0.1", grpc_port=50051, grpc_worker_threads=4 - ) + # The host comes from the `serving` bag; `grpc_worker_threads` is not a + # field (resolution sets it from the environment), so it stays on the + # stand-in the call site is handed. + override = get_context().override_server_args(host="127.0.0.1", grpc_port=50051) + override.install() + self.addCleanup(override.restore) + server_args = SimpleNamespace(grpc_worker_threads=4) with ( patch( "sglang.srt.rust_extensions.load_rust_extension", @@ -2374,7 +2379,7 @@ def test_start_server_call_site_matches_native_signature(self): tokenizer_manager=MagicMock(), template_manager=MagicMock(), scheduler_info={}, - grpc_port=resolution_result(server_args, "grpc_port"), + grpc_port=get_serving().grpc_port, ) self.assertEqual(handle, "handle") From d4ff8bb0fb7990a9d2fd85308f72b695e7e13c9f Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 20:26:21 +0000 Subject: [PATCH 20/26] config: the last tp_size reads leave the record Six sites, split by what they are actually asking. The dual-chunk attention backend shards over the *live* group, so it reads `get_parallel().tp_size` like the rest of the head-count arithmetic in the tree. The runner windows, the `/v1/loads` accelerator count, the NIXL rank arithmetic and the tokenizer's worker division all want the launch width in a process that holds no model groups, so they read `get_parallel().config.tp_size` and are registered with that reason. --- python/sglang/srt/disaggregation/nixl/conn.py | 5 +++-- python/sglang/srt/entrypoints/v1_loads.py | 2 +- .../dual_chunk_flashattention_backend.py | 4 ++-- .../srt/managers/tokenizer_control_mixin.py | 4 ++-- .../srt/model_executor/cpu_graph_runner.py | 2 +- .../srt/model_executor/runner/base_runner.py | 2 +- .../unit/test_global_config_read_ratchet.py | 20 +++++++++++++++++++ 7 files changed, 30 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index fb0c3460cc1d..f72fb25fc97f 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -46,7 +46,7 @@ resolve_dcp_dst_entry_indices, ) from sglang.srt.environ import envs -from sglang.srt.runtime_context import get_schedule +from sglang.srt.runtime_context import get_parallel, get_schedule from sglang.srt.server_args import ServerArgs try: @@ -405,7 +405,8 @@ def __init__( ): super().__init__(args, disaggregation_mode, server_args, is_mla_backend) self.transfer_source_rank = ( - self.kv_args.pp_rank * self.server_args.tp_size + self.kv_args.engine_rank + self.kv_args.pp_rank * get_parallel().config.tp_size + + self.kv_args.engine_rank ) self.kv_args.kv_data_mem_kinds = _normalize_kv_mem_kinds( getattr(self.kv_args, "kv_data_mem_kinds", None), diff --git a/python/sglang/srt/entrypoints/v1_loads.py b/python/sglang/srt/entrypoints/v1_loads.py index 93f77bd3c92b..3591fe79fcd9 100644 --- a/python/sglang/srt/entrypoints/v1_loads.py +++ b/python/sglang/srt/entrypoints/v1_loads.py @@ -146,7 +146,7 @@ async def get_loads( "version": __version__, "accelerator": _accelerator_name(), "num_accelerators": _num_accelerators_per_dp_rank( - tokenizer_manager.server_args.tp_size, + get_parallel().config.tp_size, get_parallel().config.pp_size, get_parallel().config.dp_size, get_parallel().config.enable_dp_attention, diff --git a/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py b/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py index 4c9c8ea46a80..83c22746ea20 100644 --- a/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py +++ b/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py @@ -112,10 +112,10 @@ def __init__( self.device = model_runner.device self.max_context_len = model_runner.model_config.context_len self.num_heads = model_runner.model_config.get_num_attention_heads( - model_runner.server_args.tp_size + get_parallel().tp_size ) self.num_kv_heads = model_runner.model_config.get_num_kv_heads( - model_runner.server_args.tp_size + get_parallel().tp_size ) self.head_size = model_runner.model_config.head_dim diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index 0860a5e9fb83..41923444b8dd 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -179,8 +179,8 @@ def update_control_communicator_fan_out(self: TokenizerManager, worker_count: in ) if primary_group_control: control_fan_out = ( - worker_count + self.server_args.tp_size - 1 - ) // self.server_args.tp_size + worker_count + get_parallel().config.tp_size - 1 + ) // get_parallel().config.tp_size else: control_fan_out = worker_count diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 7a1b716bf4ce..6d7843dc01b4 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -609,7 +609,7 @@ def __init__(self, model_runner: ModelRunner): self.enable_profile_cuda_graph = ( model_runner.server_args.enable_profile_cuda_graph ) - self.tp_size = model_runner.server_args.tp_size + self.tp_size = get_parallel().config.tp_size self.dp_size = get_parallel().config.dp_size self.pp_size = get_parallel().config.pp_size diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 1ebc26310e6d..180c1795e58e 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -216,7 +216,7 @@ def __init__(self, model_runner: ModelRunner) -> None: self.model_runner = model_runner self.device = model_runner.device self.device_module = torch.get_device_module(self.device) - self.tp_size = model_runner.server_args.tp_size + self.tp_size = get_parallel().config.tp_size # elastic-EP scale-up rewrites dp_size on the published config self.dp_size = get_parallel().config.dp_size self.pp_size = get_parallel().config.pp_size diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index 4b3f97992245..39e684f3a4fe 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -119,6 +119,26 @@ ("srt/managers/scheduler.py", "dcp_size"): ( "same pre-distributed-init arithmetic in configure_scheduler_process" ), + ("srt/model_executor/runner/base_runner.py", "tp_size"): ( + "the same window as the stage count next to it: a draft runner shares " + "the target's groups, so the live property would answer for the wrong " + "runner" + ), + ("srt/model_executor/cpu_graph_runner.py", "tp_size"): ( + "the same window, on the CPU graph path" + ), + ("srt/entrypoints/v1_loads.py", "tp_size"): ( + "the accelerator count is arithmetic over the launch shape, reported " + "from the tokenizer process, which holds no model groups" + ), + ("srt/disaggregation/nixl/conn.py", "tp_size"): ( + "the NIXL rank arithmetic runs on the transfer path, which the CPU-only " + "conn tests exercise without starting torch.distributed" + ), + ("srt/managers/tokenizer_control_mixin.py", "tp_size"): ( + "the tokenizer divides its worker count by the launch width; it holds " + "no model groups" + ), ("srt/model_executor/runner/base_runner.py", "pp_size"): ( "the runner's layer window is arithmetic over the configured stage " "count; a draft runner shares the target's groups, so the live " From 9beba7b4b36d2d14c8e269dedef45af5c4234898 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Mon, 24 Aug 2026 05:32:39 +0000 Subject: [PATCH 21/26] config: the Ray driver sizes its actors from the published configuration `RayEngine` publishes as part of `Engine._launch_subprocesses` and lays the actors out afterwards, so all 22 record reads in the two driver modules were reading the raw input where the bag was already available -- and both files were already mixing the two, `_compute_world_size` multiplying `get_parallel().config.pp_size` by `server_args.tp_size` on one line. A launch that leaves `dp_size` to resolution would have sized the placement group from `None`. Both modules are at zero record reads now, with a local `parallel = get_parallel().config` where a function reads several. `_compute_world_size` takes no argument. The four new configured-size reads are registered with their reason: the driver is sizing the actors that will hold the process groups, so there is nothing live to ask. The Ray path has no CI coverage (`test/manual/test_ray_engine.py` boots a real cluster), so `test_ray_driver_reads_the_bags` pins it three ways: the world-size arithmetic against a published config, the same arithmetic following a post-publish `override` -- which is what separates a bag read from a record read -- and a file-scoped check that neither module reads a field off an instance. It reports all 22 reads on the pre-conversion tree. Verified against a real cluster with a cached model: `TestRayEngineOfflineTP1` and `TestRayEngineOfflineTP2` pass (5 tests), and the custom-placement-group case launches, serves and shuts down -- it dies afterwards on a Ray GCS teardown timeout that `origin/main` hits identically. --- .../srt/ray/data_parallel_controller.py | 28 ++-- python/sglang/srt/ray/engine.py | 60 ++++----- .../unit/test_global_config_read_ratchet.py | 7 + .../unit/test_ray_driver_reads_the_bags.py | 121 ++++++++++++++++++ 4 files changed, 170 insertions(+), 46 deletions(-) create mode 100644 test/registered/unit/test_ray_driver_reads_the_bags.py diff --git a/python/sglang/srt/ray/data_parallel_controller.py b/python/sglang/srt/ray/data_parallel_controller.py index 4bbd15f681e8..89f86453b061 100644 --- a/python/sglang/srt/ray/data_parallel_controller.py +++ b/python/sglang/srt/ray/data_parallel_controller.py @@ -88,7 +88,7 @@ def launch_dp_schedulers(self, server_args: ServerArgs, port_args: PortArgs): dp_port_args_list.append(tmp_port_args) # Create ZMQ PUSH socket for this DP rank (controller → scheduler) - if server_args.node_rank == 0: + if get_parallel().config.node_rank == 0: self.workers[dp_rank] = get_zmq_socket( self.context, zmq.PUSH, @@ -139,7 +139,7 @@ def _launch_ray_tp_group( dp_rank: DP rank for regular DP; None for DP attention (derived from tp_rank). worker_ports: Pre-allocated ports for DP attention; None for regular DP. """ - nnodes = server_args.nnodes + nnodes = get_parallel().config.nnodes batch_start_idx = len(self.scheduler_actors) if not self.is_custom_pg: @@ -148,7 +148,7 @@ def _launch_ray_tp_group( pp_range, tp_range, pp_per_node, tp_per_node = _calculate_rank_ranges( nnodes, get_parallel().config.pp_size, - server_args.tp_size, + get_parallel().config.tp_size, node_rank=node_idx, ) for pp_rank in pp_range: @@ -160,13 +160,14 @@ def _launch_ray_tp_group( tp_rank % tp_per_node ) - if get_parallel().config.enable_dp_attention: + parallel = get_parallel().config + if parallel.enable_dp_attention: _, _, actual_dp_rank, _ = compute_dp_attention_world_info( - get_parallel().config.enable_dp_attention, + parallel.enable_dp_attention, tp_rank, - server_args.tp_size, - get_parallel().config.dp_size, - get_parallel().config.attn_cp_size, + parallel.tp_size, + parallel.dp_size, + parallel.attn_cp_size, ) rank_port_args = PortArgs.init_new( server_args, actual_dp_rank, worker_ports @@ -204,10 +205,11 @@ def _launch_ray_tp_group( self.scheduler_actors.append(actor) else: - world_size = _compute_world_size(server_args) + world_size = _compute_world_size() bundle_indices = _resolve_bundle_indices(self.pg, world_size) - ranks_per_tp_group = server_args.tp_size * get_parallel().config.pp_size + parallel = get_parallel().config + ranks_per_tp_group = parallel.tp_size * parallel.pp_size if dp_rank is not None: start_rank = dp_rank * ranks_per_tp_group end_rank = start_rank + ranks_per_tp_group @@ -224,8 +226,8 @@ def _launch_ray_tp_group( for global_rank in range(start_rank, end_rank): local_rank = global_rank % ranks_per_tp_group - pp_rank = local_rank // server_args.tp_size - tp_rank = local_rank % server_args.tp_size + pp_rank = local_rank // parallel.tp_size + tp_rank = local_rank % parallel.tp_size rank_port_args = port_args actual_dp_rank = dp_rank @@ -235,7 +237,7 @@ def _launch_ray_tp_group( _, _, actual_dp_rank, _ = compute_dp_attention_world_info( get_parallel().config.enable_dp_attention, tp_rank, - server_args.tp_size, + get_parallel().config.tp_size, get_parallel().config.dp_size, get_parallel().config.attn_cp_size, ) diff --git a/python/sglang/srt/ray/engine.py b/python/sglang/srt/ray/engine.py index 3dbf3e57b506..7be4b194aba0 100644 --- a/python/sglang/srt/ray/engine.py +++ b/python/sglang/srt/ray/engine.py @@ -105,18 +105,17 @@ def get_node_ip(): ) -def _compute_world_size(server_args: ServerArgs) -> int: +def _compute_world_size() -> int: """Compute world_size (total number of scheduler actors/GPUs needed). Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size. + Reads the published parallel leaves: the driver is sizing the actors that + will hold the process groups, so there is nothing live to ask. """ - if get_parallel().config.enable_dp_attention: - return server_args.tp_size * get_parallel().config.pp_size - return ( - get_parallel().config.dp_size - * server_args.tp_size - * get_parallel().config.pp_size - ) + parallel = get_parallel().config + if parallel.enable_dp_attention: + return parallel.tp_size * parallel.pp_size + return parallel.dp_size * parallel.tp_size * parallel.pp_size def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]: @@ -274,16 +273,13 @@ def _launch_scheduler_processes( placement_group as create_placement_group, ) - if get_parallel().config.enable_dp_attention: - total_gpus = server_args.tp_size * get_parallel().config.pp_size + parallel = get_parallel().config + if parallel.enable_dp_attention: + total_gpus = parallel.tp_size * parallel.pp_size else: - total_gpus = ( - get_parallel().config.dp_size - * server_args.tp_size - * get_parallel().config.pp_size - ) + total_gpus = parallel.dp_size * parallel.tp_size * parallel.pp_size - nnodes = server_args.nnodes + nnodes = parallel.nnodes gpus_per_node = total_gpus // nnodes strategy = "STRICT_PACK" if nnodes == 1 else "SPREAD" @@ -300,8 +296,8 @@ def _launch_scheduler_processes( ray.get(pg.ready()) is_custom_pg = placement_group is not None - nnodes = server_args.nnodes - world_size = _compute_world_size(server_args) + nnodes = get_parallel().config.nnodes + world_size = _compute_world_size() if not is_custom_pg: engine_bundle, engine_ip = _find_engine_bundle(pg, nnodes) @@ -341,7 +337,7 @@ def _launch_scheduler_processes( _calculate_rank_ranges( nnodes, get_parallel().config.pp_size, - server_args.tp_size, + get_parallel().config.tp_size, node_rank=node_idx, ) ) @@ -377,9 +373,10 @@ def _launch_scheduler_processes( f"bundle_indices={bundle_indices}" ) + tp_size = get_parallel().config.tp_size for rank in range(world_size): - pp_rank = rank // server_args.tp_size - tp_rank = rank % server_args.tp_size + pp_rank = rank // tp_size + tp_rank = rank % tp_size bundle_idx = bundle_indices[rank] actor = _create_scheduler_actor( @@ -455,21 +452,18 @@ def _launch_dp_scheduler_processes( RayDataParallelController, ) - if get_parallel().config.enable_dp_attention: + parallel = get_parallel().config + if parallel.enable_dp_attention: # DP attention folds DP into TP — total GPUs = tp_size * pp_size - total_gpus = server_args.tp_size * get_parallel().config.pp_size + total_gpus = parallel.tp_size * parallel.pp_size else: - total_gpus = ( - get_parallel().config.dp_size - * server_args.tp_size - * get_parallel().config.pp_size - ) - gpus_per_node = total_gpus // server_args.nnodes + total_gpus = parallel.dp_size * parallel.tp_size * parallel.pp_size + gpus_per_node = total_gpus // parallel.nnodes logger.info( - f"Ray DP cluster: {server_args.nnodes} nodes, " - f"{gpus_per_node} GPUs/node, dp_size={get_parallel().config.dp_size}, " - f"tp_size={server_args.tp_size}, pp_size={get_parallel().config.pp_size}, " - f"enable_dp_attention={get_parallel().config.enable_dp_attention}" + f"Ray DP cluster: {parallel.nnodes} nodes, " + f"{gpus_per_node} GPUs/node, dp_size={parallel.dp_size}, " + f"tp_size={parallel.tp_size}, pp_size={parallel.pp_size}, " + f"enable_dp_attention={parallel.enable_dp_attention}" ) # Set dist_init_addr on server_args so PortArgs.init_new() can compute diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index 39e684f3a4fe..de01ae6191fb 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -72,6 +72,13 @@ "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/engine.py", "tp_size"): ( + "the same placement arithmetic as the stage count: the driver sizes " + "the actors that will hold the process groups" + ), + ("srt/ray/data_parallel_controller.py", "tp_size"): ( + "the same arithmetic on the DP path, also in the driver" + ), ("srt/ray/data_parallel_controller.py", "pp_size"): ( "same placement arithmetic on the DP path -- ranks per TP group, " "computed in the driver before the actors start" diff --git a/test/registered/unit/test_ray_driver_reads_the_bags.py b/test/registered/unit/test_ray_driver_reads_the_bags.py new file mode 100644 index 000000000000..948a61be403d --- /dev/null +++ b/test/registered/unit/test_ray_driver_reads_the_bags.py @@ -0,0 +1,121 @@ +"""The Ray driver sizes its actors from the published configuration. + +`RayEngine` publishes as part of `Engine._launch_subprocesses` and *then* lays +out the actors, so the placement arithmetic reads the `parallel` bag rather than +the record it was constructed from. That matters because the record holds the +raw input: a launch that leaves `dp_size` to resolution would size the placement +group from `None`. + +There is no CI coverage of the Ray path (`test/manual/test_ray_engine.py` boots a +real cluster), so these cases drive the two pure helpers directly against a +published config -- including the override direction, which is what tells a bag +read from a record read. +""" + +import importlib.util +import unittest + +from sglang.srt.runtime_context import get_context, get_parallel +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=6, suite="base-a-test-cpu") + +# `sglang.srt.ray.engine` imports `ray` at module scope, and the CPU runner has +# no ray wheel. The file-scoped source scan below is the part that has to run +# everywhere; the three arithmetic cases need the import. +_HAS_RAY = importlib.util.find_spec("ray") is not None +_needs_ray = unittest.skipUnless(_HAS_RAY, "ray is not installed") + + +class TestRayDriverReadsTheBags(CustomTestCase): + def _publish(self, **fields): + override = get_context().override_server_args(**fields) + override.install() + self.addCleanup(override.restore) + + @_needs_ray + def test_world_size_multiplies_the_published_sizes(self): + from sglang.srt.ray.engine import _compute_world_size + + self._publish(tp_size=2, pp_size=3, dp_size=4, enable_dp_attention=False) + self.assertEqual(_compute_world_size(), 24) + + @_needs_ray + def test_dp_attention_folds_dp_into_tp(self): + from sglang.srt.ray.engine import _compute_world_size + + self._publish(tp_size=4, pp_size=2, dp_size=4, enable_dp_attention=True) + # DP attention folds DP into TP, so dp_size drops out of the product. + self.assertEqual(_compute_world_size(), 8) + + @_needs_ray + def test_the_world_size_follows_a_post_publish_override(self): + """The direction that separates a bag read from a record read. + + `override` writes the bag and never the record, so a driver still + reading `server_args.tp_size` would keep answering with the old size. + """ + from sglang.srt.ray.engine import _compute_world_size + + self._publish(tp_size=2, pp_size=1, dp_size=1, enable_dp_attention=False) + self.assertEqual(_compute_world_size(), 2) + get_context().override("test.ray_driver", tp_size=8) + self.assertEqual(get_parallel().config.tp_size, 8) + self.assertEqual(_compute_world_size(), 8) + + def test_the_driver_modules_read_no_field_off_a_record(self): + """File-scoped: neither Ray driver module reads a config field off an + instance any more. + + The Ray path has no CI coverage, so this is what keeps a new + `server_args.tp_size` from appearing in it -- the placement arithmetic + runs after the publish, and the bags are the surface that carries what + resolution decided. + """ + import ast + import dataclasses + import pathlib + + import sglang + from sglang.srt.server_args import ServerArgs + + fields = {field.name for field in dataclasses.fields(ServerArgs)} + srt = pathlib.Path(sglang.__file__).resolve().parent / "srt" + offenders = [] + for rel in ("ray/engine.py", "ray/data_parallel_controller.py"): + tree = ast.parse((srt / rel).read_text(encoding="utf-8-sig")) + holders = {"server_args", "sa"} + for node in ast.walk(tree): + if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + continue + for arg in list(node.args.args) + list(node.args.kwonlyargs): + if arg.annotation is not None and "ServerArgs" in ast.dump( + arg.annotation + ): + holders.add(arg.arg) + for node in ast.walk(tree): + if ( + isinstance(node, ast.Attribute) + and node.attr in fields + and isinstance(node.ctx, ast.Load) + and ( + (isinstance(node.value, ast.Name) and node.value.id in holders) + or ( + isinstance(node.value, ast.Attribute) + and node.value.attr == "server_args" + ) + ) + ): + offenders.append(f"{rel}:{node.lineno} reads .{node.attr}") + self.assertEqual( + offenders, + [], + "the Ray driver reads a config field off a record; the driver runs " + "after the publish, so read `get_parallel().config`:\n " + + "\n ".join(offenders), + ) + + +if __name__ == "__main__": + unittest.main() From 11ecaf3af954513e0dcf86ef986e3c50439c137f Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Mon, 24 Aug 2026 05:49:33 +0000 Subject: [PATCH 22/26] config: the last convertible readers take the bags Nine reads in eight files, each already mixing bag reads with a record read: the KV configurator's memory-saver flag, the dist-init host in `bootstrap`, the gRPC and sidecar hosts, the rust server's transport width and hosts, the in-process HTTP engine adapter, and the MindSpore runner's host. `initialize_moe_config` and `initialize_fp4_gemm_config` go with them, and they change signature: both were handed a record they read resolution's answers off (`moe_a2a_backend`, `deepep_mode`, `quantization`, the speculative pair) or, in the fp4 case, no longer read at all. They take no argument now and read `exec.moe` / `spec` / `model` / `exec.kernel`, which every caller has published by the time it calls -- the scheduler, the weight-cache daemon, and the `one_batch` work function. The five `layers/moe/utils.py` exposure pins go with the conversion, and the three tests that used to hand it a stand-in publish one. `configure_logger` keeps its record read, and this is the reason: it runs before the publish in the launcher and in the encoder HTTP entry, it is called with stand-ins, and `multimodal_gen` calls it with a *different* `ServerArgs` class that has no bags at all. A bag read there would raise on three separate paths. What is left of the 172 resolution-named instance reads this clearing started from is 39, in five places, all structural: the launcher before its publish (19), the DP controller whose declared namespace set excludes `serving` and `observability` (14), the auto-parser late resolution (4), the multimodal processor's per-instance `base_gpu_id`/`tp_size` (engines sharing a process each have their own), and `configure_logger`. --- .../skills/sglang-runtime-context/SKILL.md | 20 ++-- python/sglang/benchmark/offline_throughput.py | 15 ++- python/sglang/benchmark/one_batch.py | 107 +++++++++++------- python/sglang/benchmark/one_batch_server.py | 4 +- python/sglang/compile_deep_gemm.py | 41 ++++--- .../sglang/lang/backend/runtime_endpoint.py | 10 +- python/sglang/launch_server.py | 10 +- .../srt/configs/embedding_model_spec.py | 3 +- python/sglang/srt/distributed/bootstrap.py | 3 +- .../triton_symm_mem_ag.py | 18 +-- python/sglang/srt/entrypoints/engine.py | 43 ++++--- python/sglang/srt/entrypoints/grpc_server.py | 2 +- python/sglang/srt/entrypoints/http_server.py | 2 - python/sglang/srt/entrypoints/sidecar.py | 2 +- python/sglang/srt/eplb/expert_distribution.py | 6 +- .../attention/minimax_sparse_backend.py | 5 +- python/sglang/srt/layers/moe/utils.py | 55 +++++---- .../srt/layers/quantization/fp4_utils.py | 9 +- .../srt/layers/quantization/fp8_utils.py | 26 ++--- python/sglang/srt/managers/rust_server.py | 15 ++- python/sglang/srt/managers/scheduler.py | 6 +- .../srt/mem_cache/kv_cache_configurator.py | 5 +- .../srt/model_executor/mindspore_runner.py | 3 +- python/sglang/srt/weight_cache/daemon.py | 70 ++++++------ .../test/scripted_runtime/scheduler_hook.py | 7 +- test/manual/ep/test_flashinfer_dispatcher.py | 4 +- .../spec/test_draft_construction_isolation.py | 8 +- .../unit/test_global_config_read_ratchet.py | 16 +++ .../unit/test_ray_driver_reads_the_bags.py | 8 +- test/registered/unit/test_runtime_context.py | 9 +- ...test_supplied_instance_exposure_ratchet.py | 64 ++--------- 31 files changed, 319 insertions(+), 277 deletions(-) diff --git a/.claude/skills/sglang-runtime-context/SKILL.md b/.claude/skills/sglang-runtime-context/SKILL.md index 0eea90569526..77e93c66f771 100644 --- a/.claude/skills/sglang-runtime-context/SKILL.md +++ b/.claude/skills/sglang-runtime-context/SKILL.md @@ -139,9 +139,9 @@ bag to override at all. `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 + stays legal, unit tests rely on it). Nothing in those files reads the instance + any more -- the exposure ratchet's pin set is empty, so the next such read is a + new entry that has to argue for itself. 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). @@ -149,18 +149,17 @@ bag to override at all. supplied-instance contract; don't rewrite the parameter reads unless the field is runtime-mutated (see the elastic-EP `ep_size` case in `eplb/expert_location.py`) — **or the field is one that resolution fills in - and the callee runs in a process that has published.** That second case is - pinned debt, not a style question: the record is destined to carry the - user's raw input, so `server_args.page_size` inside a runner-owned - constructor will read the raw pre-resolution value instead of the effective - one. Debt means a decision, not automatically a bag read: pick where the + and the callee runs in a process that has published.** That second case is a + decision, not a style question: the record carries the user's raw input, so a + resolution-filled field read off it inside a runner-owned constructor answers + with the pre-resolution value instead of the effective one. Debt means a decision, not automatically a bag read: pick where the value should come from — usually the `get_*()` bag, sometimes a runner stamp or a constructor argument (the per-mode attention pair and the encode-server `gpu_id` above are dispositions of exactly this debt). The per-instance boundaries above are **not** exempt from this unless-clause (the multi-Engine exemption is retracted); each one gets its own disposition. `test_supplied_instance_exposure_ratchet.py` - pins the remaining set — three spellings of the read: `server_args.field`, + pins that set (empty today) — three spellings of the read: `server_args.field`, literal-name `getattr(server_args, "field", default)`, and the parked form (`self.x = server_args` in a method that takes the parameter, read as `self.x.field` anywhere in the class) — and fails on a new one, so the @@ -411,7 +410,8 @@ probes, swappable ACTIVE values. Not for config mirrors (read the bag leaf inste - Groups are typed dataclasses on `Flags` (`capture` / `moe` / `dp`): typo-safe writes, transactional test-only `override(**kw)` context manager. -- `flags.moe` is materialized by `initialize_moe_config(server_args)` at scheduler init; +- `flags.moe` is materialized by `initialize_moe_config()` at scheduler init (it + reads `exec.moe` / `spec` / `model`, and takes no record); accessors (`get_moe_a2a_backend` etc.) are thin shims with lazy defaults. The speculative contexts (`speculative_moe_backend_context`) swap the ACTIVE leaves around draft forwards. - `flags.dp` is materialized by `initialize_dp_attention`; `is_dp_attention_enabled()` is a diff --git a/python/sglang/benchmark/offline_throughput.py b/python/sglang/benchmark/offline_throughput.py index dd8037b38393..37acaad6ff2f 100644 --- a/python/sglang/benchmark/offline_throughput.py +++ b/python/sglang/benchmark/offline_throughput.py @@ -30,6 +30,7 @@ from sglang.benchmark.datasets.random import sample_random_requests from sglang.benchmark.utils import get_tokenizer, set_ulimit from sglang.lang.backend.runtime_endpoint import Runtime +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.entrypoints.engine import Engine from sglang.srt.server_args import ServerArgs @@ -366,6 +367,7 @@ def _create_ray_engine_backend(server_args: ServerArgs): RayEngine requires a placement group, so we launch it inside a Ray actor and return a lightweight proxy that forwards calls via ray.get(). """ + cfg = resolving_view(server_args) import ray from ray.runtime_env import RuntimeEnv from ray.util.placement_group import placement_group @@ -377,7 +379,7 @@ def _create_ray_engine_backend(server_args: ServerArgs): if not ray.is_initialized(): ray.init(runtime_env=RuntimeEnv(env_vars=env_vars)) - total_gpus = server_args.tp_size * server_args.pp_size + total_gpus = cfg.tp_size * cfg.pp_size pg = placement_group([{"CPU": 1, "GPU": total_gpus}], strategy="STRICT_PACK") ray.get(pg.ready()) @@ -398,7 +400,7 @@ def call(self, method, **kwargs): placement_group=pg, placement_group_bundle_index=0, ), - ).remote(**dict(server_args._raw_input)) + ).remote(**dict(cfg._raw_input)) class _Proxy: """Forwards method calls to the remote RayEngine actor.""" @@ -434,20 +436,21 @@ def throughput_test( ): # A programmatic caller may hand over a freshly constructed record, and # the backends below read the resolved paths and the raw snapshot. - server_args.resolve_once() + cfg = resolving_view(server_args) + cfg.resolve_once() if bench_args.backend == "engine": - if server_args.use_ray: + if cfg.use_ray: backend = _create_ray_engine_backend(server_args) else: backend = Engine(server_args=server_args) if not backend: raise ValueError("Please provide valid engine arguments") elif bench_args.backend == "runtime": - backend = Runtime(**dict(server_args._raw_input)) + backend = Runtime(**dict(cfg._raw_input)) else: raise ValueError('Please set backend to either "engine" or "runtime"') - tokenizer_id = server_args.tokenizer_path or server_args.model_path + tokenizer_id = cfg.tokenizer_path or cfg.model_path tokenizer = get_tokenizer(tokenizer_id) # Set global environments diff --git a/python/sglang/benchmark/one_batch.py b/python/sglang/benchmark/one_batch.py index d6d12a499f1e..55df78a61b98 100644 --- a/python/sglang/benchmark/one_batch.py +++ b/python/sglang/benchmark/one_batch.py @@ -64,6 +64,7 @@ import torch import torch.distributed as dist +from sglang.srt.arg_groups.overrides import resolution_result, resolving_view from sglang.srt.configs.model_config import ModelConfig from sglang.srt.distributed.parallel_state import ( destroy_distributed_environment, @@ -79,7 +80,11 @@ from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.managers.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw from sglang.srt.mem_cache.base_prefix_cache import EvictParams -from sglang.srt.model_executor.cuda_graph_config import Phase, cuda_graph_fully_disabled +from sglang.srt.model_executor.cuda_graph_config import ( + CudaGraphConfig, + Phase, + cuda_graph_fully_disabled, +) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.runtime_context import get_parallel, get_schedule, publish @@ -297,44 +302,45 @@ def from_cli_args(cls, args: argparse.Namespace): def load_model(server_args, port_args, gpu_id, tp_rank): + cfg = resolving_view(server_args) suppress_other_loggers() rank_print = print if tp_rank == 0 else lambda *args, **kwargs: None - moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size) + moe_ep_rank = tp_rank // (cfg.tp_size // cfg.ep_size) model_config = ModelConfig.from_server_args(server_args) attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = ( compute_dp_attention_world_info( - server_args.enable_dp_attention, + cfg.enable_dp_attention, tp_rank, - server_args.tp_size, - server_args.dp_size, - server_args.attn_cp_size, + cfg.tp_size, + cfg.dp_size, + cfg.attn_cp_size, ) ) ps = ParallelState( tp_rank=tp_rank, - tp_size=server_args.tp_size, + tp_size=cfg.tp_size, pp_rank=0, pp_size=1, dp_rank=None, - dp_size=server_args.dp_size, + dp_size=cfg.dp_size, attn_tp_rank=attn_tp_rank, attn_tp_size=attn_tp_size, attn_cp_rank=0, - attn_cp_size=server_args.attn_cp_size, - attn_dcp_rank=tp_rank % server_args.dcp_size, - attn_dcp_size=server_args.dcp_size, + attn_cp_size=cfg.attn_cp_size, + attn_dcp_rank=tp_rank % cfg.dcp_size, + attn_dcp_size=cfg.dcp_size, attn_dp_rank=attn_dp_rank, attn_dp_size=attn_dp_size, moe_ep_rank=moe_ep_rank, - moe_ep_size=server_args.ep_size, + moe_ep_size=cfg.ep_size, moe_dp_rank=None, - moe_dp_size=server_args.moe_dp_size, + moe_dp_size=cfg.moe_dp_size, gpu_id=gpu_id, ) runner_kwargs = dict( model_config=model_config, - mem_fraction_static=server_args.mem_fraction_static, + mem_fraction_static=cfg.mem_fraction_static, gpu_id=gpu_id, ps=ps, nccl_port=port_args.nccl_port, @@ -350,20 +356,20 @@ def load_model(server_args, port_args, gpu_id, tp_rank): model_runner = MlxModelRunnerStub(**runner_kwargs) else: model_runner = ModelRunner(**runner_kwargs) - if server_args.is_startup_weight_load_overlap: + if cfg.is_startup_weight_load_overlap: model_runner.start_startup_weight_load() model_runner.alloc_memory_pool() model_runner.init_attention_backends() model_runner.init_cuda_graphs() - if server_args.is_startup_weight_load_overlap: + if cfg.is_startup_weight_load_overlap: model_runner.finalize_startup_weight_load() rank_print(f"max_total_num_tokens={model_runner.max_total_num_tokens}") tokenizer = get_tokenizer( - server_args.tokenizer_path, - tokenizer_mode=server_args.tokenizer_mode, - trust_remote_code=server_args.trust_remote_code, + cfg.tokenizer_path, + tokenizer_mode=cfg.tokenizer_mode, + trust_remote_code=cfg.trust_remote_code, ) - if server_args.tp_size > 1: + if cfg.tp_size > 1: dist.barrier() if _use_mlx: @@ -584,19 +590,20 @@ class _MlxBenchRunner: """Wraps MlxModelRunner for the MLX benchmark path.""" def __init__(self, model_runner, server_args): + cfg = resolving_view(server_args) from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner # Radix cache requires the scheduler's allocator/trie; disable in # standalone bench mode where no scheduler is present. init_kwargs = dict( - model_path=server_args.model_path, - trust_remote_code=server_args.trust_remote_code, + model_path=cfg.model_path, + trust_remote_code=cfg.trust_remote_code, disable_radix_cache=True, - mem_fraction_static=server_args.mem_fraction_static, - quantization=server_args.quantization, + mem_fraction_static=cfg.mem_fraction_static, + quantization=cfg.quantization, ) - if server_args.max_total_tokens is not None: - init_kwargs["pool_size"] = server_args.max_total_tokens + if cfg.max_total_tokens is not None: + init_kwargs["pool_size"] = cfg.max_total_tokens self.mlx_runner = MlxModelRunner(**init_kwargs) self.mlx_runner.init_cache_pools(req_to_token_pool=None) self.fake_torch_runner = model_runner @@ -883,17 +890,18 @@ def latency_test( gpu_id, tp_rank, ): + cfg = resolving_view(server_args) # `main` runs this inline for tp_size == 1 and spawns it per rank otherwise; # a spawned child arrives with nothing published. publish(server_args, role="scheduler") - initialize_moe_config(server_args) - initialize_fp8_gemm_config(server_args) - initialize_fp4_gemm_config(server_args) + initialize_moe_config() + initialize_fp8_gemm_config() + initialize_fp4_gemm_config() - # Set CPU affinity if get_bool_env_var("SGLANG_SET_CPU_AFFINITY"): + parallel = get_parallel().config set_gpu_proc_affinity( - server_args.pp_size, server_args.tp_size, server_args.nnodes, tp_rank + parallel.pp_size, parallel.tp_size, parallel.nnodes, tp_rank ) # Configure the logger @@ -988,22 +996,45 @@ def latency_test( for result in result_list: fout.write(json.dumps(result) + "\n") - if server_args.tp_size > 1: + if cfg.tp_size > 1: destroy_model_parallel() destroy_distributed_environment() def main(server_args, bench_args): + # The decode phase has to capture the batch sizes this run benchmarks, and + # the per-phase convenience knob loses to an explicit --cuda-graph-config + # JSON (resolution applies that last), so the size is merged into that JSON. + if getattr(server_args, "_declarations_materialized", False): + # A record the caller already resolved: nothing will parse a raw dict + # again, so the declaration has to be the finished typed config. + merged = resolution_result(server_args, "cuda_graph_config") + merged = ( + merged.to_dict() + if isinstance(merged, CudaGraphConfig) + else dict(merged or {}) + ) + decode = dict(merged.get(Phase.DECODE) or {}) + decode["max_bs"] = max(bench_args.batch_size) + merged[Phase.DECODE] = decode + graph_config = CudaGraphConfig.from_dict(merged) + else: + explicit = server_args.cuda_graph_config + if isinstance(explicit, CudaGraphConfig): + explicit = explicit.to_dict() + graph_config = dict(explicit or {}) + decode = dict(graph_config.get(Phase.DECODE) or {}) + decode["max_bs"] = max(bench_args.batch_size) + graph_config[Phase.DECODE] = decode + server_args = server_args.replace_resolved( + "benchmark.one_batch", cuda_graph_config=graph_config + ) server_args.resolve_once() - - # The legacy cuda_graph_max_bs_decode field does not propagate; set the - # decode phase. - if server_args.cuda_graph_config is not None: - server_args.cuda_graph_config[Phase.DECODE].max_bs = max(bench_args.batch_size) + cfg = resolving_view(server_args) _set_envs_and_config(server_args) - if server_args.model_path: + if cfg.model_path: if bench_args.correctness_test: work_func = correctness_test else: diff --git a/python/sglang/benchmark/one_batch_server.py b/python/sglang/benchmark/one_batch_server.py index 76a4a0314bbd..3ea5f80b9158 100644 --- a/python/sglang/benchmark/one_batch_server.py +++ b/python/sglang/benchmark/one_batch_server.py @@ -33,6 +33,7 @@ from sglang.benchmark.endpoint import acquire_endpoint from sglang.benchmark.utils import get_processor, get_tokenizer from sglang.profiler import run_profile +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.server_args import ServerArgs @@ -1234,6 +1235,7 @@ def run_benchmark_internal( def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs): + cfg = resolving_view(server_args) results, server_info = run_benchmark_internal(server_args, bench_args) # Save results as pydantic models in the JSON format @@ -1241,7 +1243,7 @@ def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs): save_results_as_pydantic_models( results, pydantic_result_filename=bench_args.pydantic_result_filename, - model_path=server_args.model_path, + model_path=cfg.model_path, server_args=bench_args.server_args_for_metrics, ) diff --git a/python/sglang/compile_deep_gemm.py b/python/sglang/compile_deep_gemm.py index ab5ac06f38a7..bb07cc5d051f 100644 --- a/python/sglang/compile_deep_gemm.py +++ b/python/sglang/compile_deep_gemm.py @@ -17,13 +17,19 @@ import requests +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.entrypoints.warmup import warmup from sglang.srt.environ import envs from sglang.srt.managers.io_struct import GenerateReqInput from sglang.srt.managers.tokenizer_manager import TokenizerManager -from sglang.srt.model_executor.cuda_graph_config import Backend, Phase +from sglang.srt.model_executor.cuda_graph_config import ( + Backend, + CudaGraphConfig, + Phase, +) +from sglang.srt.runtime_context import get_parallel from sglang.srt.server_args import ServerArgs from sglang.srt.utils import kill_process_tree @@ -59,8 +65,7 @@ async def warm_up_compile( disaggregation_mode: str, tokenizer_manager: TokenizerManager ): print("\nGenerate warm up request for compiling DeepGEMM...\n") - server_args = tokenizer_manager.server_args - dp_size = server_args.dp_size + dp_size = get_parallel().config.dp_size base_ids = [0, 1, 2, 3] sampling_params = { "temperature": 0.0, @@ -76,7 +81,8 @@ async def warm_up_compile( ) generate_req_input.bootstrap_host = [FAKE_BOOTSTRAP_HOST] * dp_size generate_req_input.bootstrap_room = [ - i * (2**63 // dp_size) + (i % server_args.tp_size) for i in range(dp_size) + i * (2**63 // dp_size) + (i % get_parallel().config.tp_size) + for i in range(dp_size) ] else: input_ids = ( @@ -105,6 +111,7 @@ def launch_server_process_and_send_one_request( # Keeps the device probe out of the fork below, for a caller that reaches # this without resolving first. server_args.resolve_once() + cfg = resolving_view(server_args) proc = multiprocessing.Process(target=launch_server_internal, args=(server_args,)) proc.start() @@ -125,7 +132,7 @@ def launch_server_process_and_send_one_request( if response.status_code == 200: # Rank-0 node send a request to sync with other node and then return. if server_args.node_rank == 0: - dp_size = server_args.dp_size + dp_size = cfg.dp_size base_ids = [0, 1, 2, 3] payload = { "sampling_params": { @@ -133,11 +140,11 @@ def launch_server_process_and_send_one_request( "temperature": 0, }, } - if server_args.disaggregation_mode != "null": + if cfg.disaggregation_mode != "null": payload["input_ids"] = [list(base_ids) for _ in range(dp_size)] payload["bootstrap_host"] = [FAKE_BOOTSTRAP_HOST] * dp_size payload["bootstrap_room"] = [ - i * (2**63 // dp_size) + (i % server_args.tp_size) + i * (2**63 // dp_size) + (i % cfg.tp_size) for i in range(dp_size) ] else: @@ -177,16 +184,24 @@ def compile_server_args(args, compile_args: CompileArgs) -> ServerArgs: """The config this script serves with: no cuda graph, no torch compile, and a watchdog that outlives the compilation.""" args.enable_torch_compile = False + # The convenience flags lose to an explicit --cuda-graph-config JSON, which + # resolution applies last, so this tool's "no cuda graph" guarantee is + # merged into that JSON instead -- an operator serving with their own config + # still compiles without capture. + explicit = args.cuda_graph_config + if isinstance(explicit, CudaGraphConfig): + explicit = explicit.to_dict() + explicit = dict(explicit or {}) + for phase in (Phase.DECODE, Phase.PREFILL): + phase_config = dict(explicit.get(phase) or {}) + phase_config["backend"] = Backend.DISABLED + explicit[phase] = phase_config + args.cuda_graph_config = explicit # Watchdog timeout follows compile_args.timeout because compilation takes long. args.watchdog_timeout = compile_args.timeout args.warmups = "compile-deep-gemm" - server_args = ServerArgs.from_cli_args(args) - # `cuda_graph_config` is None until resolution parses it. - server_args.resolve_once() - server_args.cuda_graph_config[Phase.DECODE].backend = Backend.DISABLED - server_args.cuda_graph_config[Phase.PREFILL].backend = Backend.DISABLED print(f"Disable CUDA Graph and Torch Compile to save time...") - return server_args + return ServerArgs.from_cli_args(args) def run_compile(server_args: ServerArgs, compile_args: CompileArgs): diff --git a/python/sglang/lang/backend/runtime_endpoint.py b/python/sglang/lang/backend/runtime_endpoint.py index 0e74efd5202b..e0e3a4c82438 100644 --- a/python/sglang/lang/backend/runtime_endpoint.py +++ b/python/sglang/lang/backend/runtime_endpoint.py @@ -456,13 +456,15 @@ def cache_prefix(self, prefix: str): self.endpoint.cache_prefix(prefix) def get_tokenizer(self): + from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.utils.hf_transformers_utils import get_tokenizer + cfg = resolving_view(self.server_args) return get_tokenizer( - self.server_args.tokenizer_path or self.server_args.model_path, - tokenizer_mode=self.server_args.tokenizer_mode, - trust_remote_code=self.server_args.trust_remote_code, - revision=self.server_args.revision, + cfg.tokenizer_path or cfg.model_path, + tokenizer_mode=cfg.tokenizer_mode, + trust_remote_code=cfg.trust_remote_code, + revision=cfg.revision, ) async def async_generate( diff --git a/python/sglang/launch_server.py b/python/sglang/launch_server.py index 0ee16c0ca480..1ddbb4e68876 100644 --- a/python/sglang/launch_server.py +++ b/python/sglang/launch_server.py @@ -5,6 +5,7 @@ import sys import warnings +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.plugins import load_plugins from sglang.srt.server_args import prepare_server_args from sglang.srt.utils import kill_process_tree @@ -18,10 +19,11 @@ def run_server(server_args): # The flags dispatched on below are decided by resolution (`--grpc-mode` # folds into `smg_grpc_mode`), and `prepare_server_args` returns raw input. server_args.resolve_once() + cfg = resolving_view(server_args) - if server_args.encoder_only: + if cfg.encoder_only: # For encoder disaggregation - if server_args.smg_grpc_mode or server_args.grpc_mode: + if cfg.smg_grpc_mode or cfg.grpc_mode: from sglang.srt.disaggregation.encoder.grpc_server import ( serve_grpc_encoder, ) @@ -31,7 +33,7 @@ def run_server(server_args): from sglang.srt.disaggregation.encoder.http_server import launch_server launch_server(server_args) - elif server_args.smg_grpc_mode: + elif cfg.smg_grpc_mode: # Legacy SMG gRPC server (--smg-grpc-mode, or the deprecated --grpc-mode # which __post_init__ folds into smg_grpc_mode). The native Rust gRPC # server is a separate path, enabled by --grpc-port, that starts @@ -39,7 +41,7 @@ def run_server(server_args): from sglang.srt.entrypoints.grpc_server import serve_grpc asyncio.run(serve_grpc(server_args)) - elif server_args.use_ray: + elif cfg.use_ray: # Ray mode: HTTP mode with Ray backend. try: from sglang.srt.ray.http_server import launch_server diff --git a/python/sglang/srt/configs/embedding_model_spec.py b/python/sglang/srt/configs/embedding_model_spec.py index a621a9d0d04f..6bcf5a326cec 100644 --- a/python/sglang/srt/configs/embedding_model_spec.py +++ b/python/sglang/srt/configs/embedding_model_spec.py @@ -228,8 +228,7 @@ def resolved_embedding_plan( This boundary deliberately accepts duck-typed arguments so the declarative registry remains independent of ServerArgs and ModelConfig import cycles. `config` must answer with the *resolved* configuration -- the readback - callers pass `resolving_view(record)`, since the record's fields are the raw - input. + callers pass `resolving_view(record)`, which is where a decision lives. """ prefill_graph = getattr(getattr(config, "cuda_graph_config", None), "prefill", None) diff --git a/python/sglang/srt/distributed/bootstrap.py b/python/sglang/srt/distributed/bootstrap.py index 623b9aa63bc8..f5a4c438786c 100644 --- a/python/sglang/srt/distributed/bootstrap.py +++ b/python/sglang/srt/distributed/bootstrap.py @@ -31,6 +31,7 @@ from sglang.srt.runtime_context import ( get_exec, get_parallel, + get_serving, ) from sglang.srt.server_args import ServerArgs from sglang.srt.utils import ( @@ -187,7 +188,7 @@ def _resolve_dist_init_method(*, server_args: ServerArgs, dist_port: int) -> str dist_init_method = na.to_tcp() else: dist_init_method = NetworkAddress( - server_args.host or "127.0.0.1", dist_port + get_serving().host or "127.0.0.1", dist_port ).to_tcp() return dist_init_method diff --git a/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py b/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py index 74e82d73a623..0e7c6a7a0ba2 100644 --- a/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py +++ b/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py @@ -423,19 +423,19 @@ def recommended_max_tokens(include_prefill: bool, floor: int = 0) -> int: NCCL. Covers the spec-decode batch plus, if ``include_prefill``, a prefill chunk. Returns ``floor`` if server args are unavailable.""" try: - from sglang.srt.runtime_context import get_server_args + from sglang.srt.runtime_context import get_schedule, get_spec - sa = get_server_args() + def g(value) -> int: + return value if isinstance(value, int) and value > 0 else 0 - def g(name: str) -> int: - v = getattr(sa, name, 0) - return v if isinstance(v, int) and v > 0 else 0 - - tokens = g("max_running_requests") * max( - g("speculative_num_draft_tokens"), g("speculative_eagle_topk"), 1 + schedule, spec = get_schedule(), get_spec() + tokens = g(schedule.max_running_requests) * max( + g(spec.speculative_num_draft_tokens), g(spec.speculative_eagle_topk), 1 ) if include_prefill: - tokens = max(tokens, g("chunked_prefill_size"), g("max_prefill_tokens")) + tokens = max( + tokens, g(schedule.chunked_prefill_size), g(schedule.max_prefill_tokens) + ) return max(tokens, floor) except Exception: return floor diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index e4e625c6062d..670831be1d17 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -48,6 +48,7 @@ import uvloop import zmq +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.elastic_ep.expert_backup_manager import run_expert_backup_manager from sglang.srt.entrypoints.engine_info_bootstrap_server import ( EngineInfoBootstrapServer, @@ -98,6 +99,7 @@ from sglang.srt.parser.template_manager import TemplateManager from sglang.srt.plugins import load_plugins from sglang.srt.runtime_context import ( + get_disagg, get_exec, get_model, get_parallel, @@ -309,9 +311,9 @@ def __init__(self, **kwargs): trace_modules=server_args.trace_modules, ) thread_label = "Tokenizer" - if server_args.disaggregation_mode == "prefill": + if get_disagg().disaggregation_mode == "prefill": thread_label = "Prefill Tokenizer" - elif server_args.disaggregation_mode == "decode": + elif get_disagg().disaggregation_mode == "decode": thread_label = "Decode Tokenizer" trace_set_thread_info(thread_label) @@ -1056,10 +1058,8 @@ def _launch_subprocesses( # Needs a tokenizer and a chat template, so it cannot live in the # pipeline; after the plugins, which may register the parser detected. - if ( - server_args.reasoning_parser == "auto" - or server_args.tool_call_parser == "auto" - ): + parsers = resolving_view(server_args) + if parsers.reasoning_parser == "auto" or parsers.tool_call_parser == "auto": resolve_auto_parsers(server_args) # This publish replaces whatever was published before it, so the @@ -1627,27 +1627,28 @@ def save_sharded_model(self, **kwargs): def _set_envs_and_config(server_args: ServerArgs): + cfg = resolving_view(server_args) # Set global environments # MNNVL fabric (GB200/GB300) multi-node: cross-node NVLink needs NCCL's # cuMem-based buffers and MNNVL transport. Default them on (user-set # values win; the symm-mem override below only fires when unset). - if server_args.nnodes > 1 and is_mnnvl_fabric_device(): + if cfg.nnodes > 1 and is_mnnvl_fabric_device(): os.environ.setdefault("NCCL_CUMEM_ENABLE", "1") os.environ.setdefault("NCCL_MNNVL_ENABLE", "1") - if "NCCL_CUMEM_ENABLE" not in os.environ or server_args.enable_symm_mem: - os.environ["NCCL_CUMEM_ENABLE"] = str(int(server_args.enable_symm_mem)) + if "NCCL_CUMEM_ENABLE" not in os.environ or cfg.enable_symm_mem: + os.environ["NCCL_CUMEM_ENABLE"] = str(int(cfg.enable_symm_mem)) if ( "NCCL_NVLS_ENABLE" not in os.environ - or server_args.enable_nccl_nvls - or server_args.enable_symm_mem + or cfg.enable_nccl_nvls + or cfg.enable_symm_mem ): os.environ["NCCL_NVLS_ENABLE"] = str( - int(server_args.enable_nccl_nvls or server_args.enable_symm_mem) + int(cfg.enable_nccl_nvls or cfg.enable_symm_mem) ) - if "NCCL_GRAPH_MIXING_SUPPORT" not in os.environ or server_args.enable_symm_mem: + if "NCCL_GRAPH_MIXING_SUPPORT" not in os.environ or cfg.enable_symm_mem: # Note(wh): NCCL_GRAPH_MIXING_SUPPORT=0 can help improve performance for symmetric kernels. # details in https://github.com/NVIDIA/nccl-tests/issues/333#issuecomment-3103636985 - if server_args.dcp_size > 1: + if cfg.dcp_size > 1: os.environ["NCCL_GRAPH_MIXING_SUPPORT"] = "0" os.environ["CUDA_DEVICE_MAX_CONNECTIONS"] = "8" @@ -1669,7 +1670,7 @@ def _set_envs_and_config(server_args: ServerArgs): ) # Set prometheus env vars - if server_args.enable_metrics: + if cfg.enable_metrics: set_prometheus_multiproc_dir() # Set ulimit @@ -1677,7 +1678,7 @@ def _set_envs_and_config(server_args: ServerArgs): # Check flashinfer version if not get_bool_env_var("SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK"): - if "flashinfer" in server_args.get_attention_backends(): + if "flashinfer" in cfg.get_attention_backends(): assert_pkg_version( "flashinfer_python", "0.6.17", @@ -1694,7 +1695,7 @@ def _set_envs_and_config(server_args: ServerArgs): # Signal handlers can only be registered from the main thread. if threading.current_thread() is threading.main_thread(): - if server_args.custom_sigquit_handler is None: + if cfg.custom_sigquit_handler is None: # Register the signal handler. # The child processes will send SIGQUIT to this process when any error happens # This process then clean up the whole process tree @@ -1709,10 +1710,8 @@ def launch_phase_sigquit_handler(signum, frame): signal.signal(signal.SIGQUIT, launch_phase_sigquit_handler) else: # Allow users to register a custom SIGQUIT handler for things like crash dump - logger.error( - f"Using custom SIGQUIT handler: {server_args.custom_sigquit_handler}" - ) - signal.signal(signal.SIGQUIT, server_args.custom_sigquit_handler) + logger.error(f"Using custom SIGQUIT handler: {cfg.custom_sigquit_handler}") + signal.signal(signal.SIGQUIT, cfg.custom_sigquit_handler) else: logger.warning( "Signal handler is not added because the engine is not in the " @@ -1724,7 +1723,7 @@ def launch_phase_sigquit_handler(signum, frame): mp.set_start_method("spawn", force=True) # Set gc threshold - if gc_threshold := server_args.gc_threshold: + if gc_threshold := cfg.gc_threshold: gc.set_threshold(*gc_threshold) _log_legacy_kernel_cache_dirs() diff --git a/python/sglang/srt/entrypoints/grpc_server.py b/python/sglang/srt/entrypoints/grpc_server.py index b78e24341137..0c2a8c559316 100644 --- a/python/sglang/srt/entrypoints/grpc_server.py +++ b/python/sglang/srt/entrypoints/grpc_server.py @@ -213,7 +213,7 @@ async def _on_request_manager_ready(request_manager, srv_args, sched_info): ) try: sidecar_runner = await _start_sidecar_server( - server_args.host, sidecar_port, sidecar_app + cfg.host, sidecar_port, sidecar_app ) except OSError as e: logger.error( diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index be7a5943c83c..d56aa71e39b5 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -773,8 +773,6 @@ async def model_info(): if embedding_model_spec is not None: result["embedding"] = resolved_embedding_plan( embedding_model_spec, - # Through the declarations: the plan reports the *effective* config, - # and the fields hold the raw input. config=resolving_view(_global_state.tokenizer_manager.server_args), model_config=model_config, ) diff --git a/python/sglang/srt/entrypoints/sidecar.py b/python/sglang/srt/entrypoints/sidecar.py index d7107a975803..754950af26fb 100644 --- a/python/sglang/srt/entrypoints/sidecar.py +++ b/python/sglang/srt/entrypoints/sidecar.py @@ -118,7 +118,7 @@ def start_sidecar(server_args) -> Sidecar: module_name = server_args.sidecar assert module_name is not None sidecar_args, shutdown_timeout = _parse_sidecar_args(server_args.sidecar_args) - endpoint = build_sidecar_endpoint(server_args.host, get_serving().grpc_port) + endpoint = build_sidecar_endpoint(get_serving().host, get_serving().grpc_port) proc = mp.get_context("spawn").Process( name=f"sglang_sidecar_{module_name}", target=_run_sidecar, diff --git a/python/sglang/srt/eplb/expert_distribution.py b/python/sglang/srt/eplb/expert_distribution.py index ec6133956a7f..458ecfc75d80 100644 --- a/python/sglang/srt/eplb/expert_distribution.py +++ b/python/sglang/srt/eplb/expert_distribution.py @@ -764,7 +764,7 @@ def _append_utilization_rate( single_pass_global_physical_count, num_gpu=self._expert_location_metadata.ep_size, ) - gpu_physical_count = gpu_physical_count.to(self._server_args.device) + gpu_physical_count = gpu_physical_count.to(get_device_namespace().device) torch.distributed.reduce( gpu_physical_count, dst=0, op=torch.distributed.ReduceOp.SUM ) @@ -898,9 +898,9 @@ def __init__(self, *args, **kwargs): # Cannot use local_physical_count to support select_experts self._expert_location_metadata.num_physical_experts, ), - buffer_size=self._server_args.expert_distribution_recorder_buffer_size, + buffer_size=get_exec().moe.expert_distribution_recorder_buffer_size, dtype=torch.int32, - device=self._server_args.device, + device=get_device_namespace().device, ) self._first_dump = True diff --git a/python/sglang/srt/layers/attention/minimax_sparse_backend.py b/python/sglang/srt/layers/attention/minimax_sparse_backend.py index def19cbc655b..49e617e39232 100644 --- a/python/sglang/srt/layers/attention/minimax_sparse_backend.py +++ b/python/sglang/srt/layers/attention/minimax_sparse_backend.py @@ -7,6 +7,7 @@ import torch +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.configs.model_config import ( get_minimax_sparse_attention_config, get_minimax_sparse_disable_value_layer_ids, @@ -116,7 +117,9 @@ def __init__(self, runner: ModelRunner): self.max_context_len = int(runner.model_config.context_len) # Per-forward cache for the native decode block table (rebuilt each forward). self._native_decode_bt: dict = {} - self.fp8_attn_gemm = m3_fp8_attn_gemm_enabled(runner.server_args) + self.fp8_attn_gemm = m3_fp8_attn_gemm_enabled( + resolving_view(runner.server_args) + ) if self.fp8_attn_gemm: assert self.kv_pool.main_pool.dtype == torch.float8_e4m3fn, ( "fp8 attn-GEMM mode requires an fp8_e4m3fn main KV pool, got " diff --git a/python/sglang/srt/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py index 7f4e510241e0..11f8a4b156f1 100644 --- a/python/sglang/srt/layers/moe/utils.py +++ b/python/sglang/srt/layers/moe/utils.py @@ -4,7 +4,6 @@ import os from contextlib import contextmanager from enum import Enum, IntEnum -from typing import TYPE_CHECKING import torch @@ -12,14 +11,18 @@ from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) -from sglang.srt.runtime_context import get_exec, get_flags, get_forward, get_parallel +from sglang.srt.runtime_context import ( + get_exec, + get_flags, + get_forward, + get_model, + get_parallel, + get_spec, +) from sglang.srt.utils import is_cuda, is_npu _is_npu = is_npu() -if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs - from sglang.srt.runtime_context import get_server_args from sglang.srt.utils.common import log_info_on_rank0 @@ -308,37 +311,47 @@ def get_ascend_dispatcher_output_dtype(dispatcher): return DispatcherOutputDtype.BF16 -def initialize_moe_config(server_args: ServerArgs): +def initialize_moe_config(): + """Seed the MoE runtime flags from the published configuration. + + Reads the bags: `moe_a2a_backend` and its siblings are resolution's + answers, and the record carries them only while declarations materialize + onto it. Called once per process after publish + (scheduler init, the benchmark work functions). + """ + exec_moe = get_exec().moe + overlap = get_exec().overlap + spec = get_spec() moe = get_flags().moe - moe.a2a_backend = MoeA2ABackend(server_args.moe_a2a_backend) - moe.runner_backend = MoeRunnerBackend(server_args.moe_runner_backend) + moe.a2a_backend = MoeA2ABackend(exec_moe.moe_a2a_backend) + moe.runner_backend = MoeRunnerBackend(exec_moe.moe_runner_backend) moe.speculative_runner_backend = ( - MoeRunnerBackend(server_args.speculative_moe_runner_backend) - if server_args.speculative_moe_runner_backend is not None + MoeRunnerBackend(spec.speculative_moe_runner_backend) + if spec.speculative_moe_runner_backend is not None else moe.runner_backend ) moe.speculative_a2a_backend = ( - MoeA2ABackend(server_args.speculative_moe_a2a_backend) - if server_args.speculative_moe_a2a_backend is not None + MoeA2ABackend(spec.speculative_moe_a2a_backend) + if spec.speculative_moe_a2a_backend is not None else moe.a2a_backend ) - moe.deepep_mode = DeepEPMode(server_args.deepep_mode) - moe.deepep_config = server_args.deepep_config or "" - moe.tbo_enabled = server_args.enable_two_batch_overlap - moe.sbo_enabled = server_args.enable_single_batch_overlap + moe.deepep_mode = DeepEPMode(exec_moe.deepep_mode) + moe.deepep_config = exec_moe.deepep_config or "" + moe.tbo_enabled = overlap.enable_two_batch_overlap + moe.sbo_enabled = overlap.enable_single_batch_overlap if moe.sbo_enabled and is_cuda(): if torch.cuda.get_device_capability()[0] == 9: raise ValueError( "SBO (single batch overlap) is not supported on SM90 GPUs with latest sgl-deep-gemm wheel. Please try removing --enable-single-batch-overlap argument." ) - moe.tbo_token_distribution_threshold = server_args.tbo_token_distribution_threshold - moe.disable_fp4_allgather = server_args.disable_flashinfer_cutlass_moe_fp4_allgather - moe.quantization = server_args.quantization + moe.tbo_token_distribution_threshold = overlap.tbo_token_distribution_threshold + moe.disable_fp4_allgather = exec_moe.disable_flashinfer_cutlass_moe_fp4_allgather + moe.quantization = get_model().quantization # Seeded with the user's intent; each model's gate refines the ACTIVE # value for its own build (install_shared_experts_fusion_decision). - moe.disable_shared_experts_fusion = server_args.disable_shared_experts_fusion + moe.disable_shared_experts_fusion = exec_moe.disable_shared_experts_fusion moe.speculative_disable_shared_experts_fusion = ( - server_args.disable_shared_experts_fusion + exec_moe.disable_shared_experts_fusion ) diff --git a/python/sglang/srt/layers/quantization/fp4_utils.py b/python/sglang/srt/layers/quantization/fp4_utils.py index 5cb7f097adef..52588dbbf726 100644 --- a/python/sglang/srt/layers/quantization/fp4_utils.py +++ b/python/sglang/srt/layers/quantization/fp4_utils.py @@ -2,7 +2,7 @@ import logging from enum import Enum -from typing import TYPE_CHECKING, Optional +from typing import Optional import torch @@ -14,9 +14,6 @@ ) from sglang.srt.utils.custom_op import register_custom_op_from_extern -if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs - logger = logging.getLogger(__name__) @@ -143,8 +140,8 @@ def get_flashinfer_backend(self) -> str: FP4_GEMM_RUNNER_BACKEND: Fp4GemmRunnerBackend | None = None -def initialize_fp4_gemm_config(server_args: ServerArgs) -> None: - """Initialize FP4 GEMM configuration from server args.""" +def initialize_fp4_gemm_config() -> None: + """Initialize the FP4 GEMM backend from the published configuration.""" global FP4_GEMM_RUNNER_BACKEND backend = get_exec().kernel.fp4_gemm_runner_backend diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index f71ba47844e2..267bf3b976a6 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -3,23 +3,10 @@ import logging from enum import Enum from functools import lru_cache, partial -from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union +from typing import Callable, List, Optional, Tuple, Union import torch -from sglang.kernels.ops.quantization.fp8_kernel import ( - sglang_per_token_group_quant_fp8, - sglang_per_token_group_quant_fp8_row_padded, -) -from sglang.srt.environ import envs -from sglang.srt.layers import deep_gemm_wrapper -from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil -from sglang.srt.runtime_context import get_exec, get_parallel -from sglang.srt.utils.common import torch_release - -if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs - from sglang.kernels.ops.quantization.fp8_kernel import ( fp8_dtype, fp8_max, @@ -28,12 +15,18 @@ is_fp8_fnuz, per_token_group_quant_fp8, scaled_fp8_quant, + sglang_per_token_group_quant_fp8, + sglang_per_token_group_quant_fp8_row_padded, sglang_per_token_quant_fp8, static_quant_fp8, triton_scaled_mm, w8a8_block_fp8_matmul_deepgemm, w8a8_block_fp8_matmul_triton, ) +from sglang.srt.environ import envs +from sglang.srt.layers import deep_gemm_wrapper +from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import ( ceil_align, ceil_div, @@ -54,6 +47,7 @@ is_xpu, offloader, ) +from sglang.srt.utils.common import torch_release from sglang.srt.utils.custom_op import register_custom_op logger = logging.getLogger(__name__) @@ -799,11 +793,11 @@ def _dispatch_auto_backend() -> Callable: return triton_w8a8_block_fp8_linear -def initialize_fp8_gemm_config(server_args: ServerArgs) -> None: +def initialize_fp8_gemm_config() -> None: """Initialize FP8 GEMM configuration.""" global FP8_GEMM_RUNNER_BACKEND - backend = server_args.fp8_gemm_runner_backend + backend = get_exec().kernel.fp8_gemm_runner_backend if backend == "auto" and is_sm120_supported(): backend = "cutlass" diff --git a/python/sglang/srt/managers/rust_server.py b/python/sglang/srt/managers/rust_server.py index 79c1a48aaf4b..4251f15257b2 100644 --- a/python/sglang/srt/managers/rust_server.py +++ b/python/sglang/srt/managers/rust_server.py @@ -20,6 +20,7 @@ import msgspec +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.environ import envs from sglang.srt.managers.io_struct import TokenizedGenerateReqInput from sglang.srt.managers.utils import ( @@ -29,6 +30,8 @@ ) from sglang.srt.runtime_context import ( get_mm, + get_observability, + get_parallel, get_serving, ) from sglang.srt.utils.flatten import ( @@ -272,7 +275,7 @@ def _use_feature_shm(self) -> bool: ) return ( - self.server_args.tp_size > 1 + get_parallel().config.tp_size > 1 and determine_tensor_transport_mode() != "default" and not self.server_args.skip_tokenizer_init ) @@ -395,13 +398,13 @@ def launch(cls, scheduler: Scheduler) -> RustServer: "ingress has no equivalent). Launch without SGLANG_RUST_SERVER, or " "drop --preferred-sampling-params and send those values per request." ) - http_addr = f"{server_args.host}:{server_args.port}" + http_addr = f"{get_serving().host}:{server_args.port}" # Per-DP-rank HTTP port with client load balancing. `None` when DP is off, # so the rank is not conflated with rank 0 of a one-rank group. dp_rank = scheduler.ps.attn_dp_rank if scheduler.ps.dp_size > 1 else None if dp_rank is not None: - http_addr = f"{server_args.host}:{server_args.port + dp_rank}" + http_addr = f"{get_serving().host}:{server_args.port + dp_rank}" launch_cores, server_cores = cls._partition_cores( mm_workers=( @@ -754,7 +757,7 @@ def _build_server_args(scheduler: Scheduler) -> ServerArgs: ext = load_rust_extension("sglang.srt.rust_extensions._server") - sa = scheduler.server_args + sa = resolving_view(scheduler.server_args) mc = scheduler.model_config disaggregation_mode = { "null": ext.DisaggregationMode.Null, @@ -768,9 +771,9 @@ def _build_server_args(scheduler: Scheduler) -> ServerArgs: revision=sa.revision, load_format=sa.load_format, weight_version=sa.weight_version, - host=sa.host, + host=get_serving().host, port=sa.port, - log_level=sa.log_level, + log_level=get_observability().log_level, log_level_http=sa.log_level_http, chat_template=sa.chat_template, tool_call_parser=sa.tool_call_parser, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 24b2689c9ca5..fd3a27bf5b45 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -903,11 +903,11 @@ def init_moe_gemm_config(self): "moe_topk", ) if any(hasattr(config_to_check, attr) for attr in moe_topk_attrs): - initialize_moe_config(self.server_args) + initialize_moe_config() # Initialize GEMM-related configuration for FP8 and FP4 backends. - initialize_fp8_gemm_config(self.server_args) - initialize_fp4_gemm_config(self.server_args) + initialize_fp8_gemm_config() + initialize_fp4_gemm_config() initialize_bf16_gemm_config(self.server_args) # This must be called after initialize_moe_config diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index ed0215bfffdf..75821be4df43 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -8,6 +8,7 @@ import msgspec import torch +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.configs.hybrid_arch import ( hybrid_gdn_config, kimi_linear_config, @@ -1262,7 +1263,7 @@ def _build_ascend_minimax_sparse_kv_pool( sparse_layer_ids=sparse_layer_ids, disable_value_sparse_layer_ids=disable_value_sparse_layer_ids, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, ) @@ -1533,7 +1534,7 @@ def _build_minimax_sparse_kv_pool(self, *, max_total_num_tokens: int) -> KVCache # with the widening-dequant contract. index_dtype=( self.kv_cache_dtype - if m3_fp8_attn_gemm_enabled(self.server_args) + if m3_fp8_attn_gemm_enabled(resolving_view(self.server_args)) else self.model_dtype ), head_num=self.model_config.get_num_kv_heads( diff --git a/python/sglang/srt/model_executor/mindspore_runner.py b/python/sglang/srt/model_executor/mindspore_runner.py index 4cdcaed505b6..6d3a4d06aa21 100644 --- a/python/sglang/srt/model_executor/mindspore_runner.py +++ b/python/sglang/srt/model_executor/mindspore_runner.py @@ -14,6 +14,7 @@ from mindspore.communication import create_group from sglang.srt.distributed.parallel_state import _groups +from sglang.srt.runtime_context import get_serving logger = logging.getLogger(__name__) @@ -109,7 +110,7 @@ def init_ms_distributed(world_size, rank, local_rank, server_args, port): if server_args.dist_init_addr: dist_init_method = f"tcp://{server_args.dist_init_addr}" else: - dist_init_method = f"tcp://{server_args.host}:{port}" + dist_init_method = f"tcp://{get_serving().host}:{port}" set_ms_parallel_env(rank, local_rank, world_size, dist_init_method) ms.set_context(infer_boost="on", jit_level="O0") diff --git a/python/sglang/srt/weight_cache/daemon.py b/python/sglang/srt/weight_cache/daemon.py index d14d79075fbc..8769061c9af3 100644 --- a/python/sglang/srt/weight_cache/daemon.py +++ b/python/sglang/srt/weight_cache/daemon.py @@ -47,6 +47,7 @@ import torch import torch.distributed as dist +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.configs.load_config import LoadConfig from sglang.srt.platforms import current_platform from sglang.srt.runtime_context import get_parallel, publish @@ -148,27 +149,29 @@ def __init__( dist_init_method: Optional[str] = None, ): self.server_args = server_args - self.model_path = server_args.model_path + # The daemon publishes further down, in `load()`. + cfg = resolving_view(server_args) + self.model_path = cfg.model_path self.gpu_id = gpu_id - self.tp_size = server_args.tp_size + self.tp_size = cfg.tp_size self.tp_rank = tp_rank - self.pp_size = server_args.pp_size + self.pp_size = cfg.pp_size self.pp_rank = pp_rank - self.dp_size = server_args.dp_size - self.ep_size = server_args.ep_size - self.moe_dp_size = server_args.moe_dp_size - self.enable_dp_attention = server_args.enable_dp_attention - self.enable_dp_lm_head = server_args.enable_dp_lm_head - self.attn_cp_size = server_args.attn_cp_size - self.moe_dense_tp_size = server_args.moe_dense_tp_size - self.moe_a2a_backend = server_args.moe_a2a_backend - self.deepep_mode = server_args.deepep_mode - self.load_format = server_args.load_format - self.dtype = server_args.dtype - self.quantization = server_args.quantization - self.model_loader_extra_config = server_args.model_loader_extra_config - self.trust_remote_code = server_args.trust_remote_code - self.revision = server_args.revision + self.dp_size = cfg.dp_size + self.ep_size = cfg.ep_size + self.moe_dp_size = cfg.moe_dp_size + self.enable_dp_attention = cfg.enable_dp_attention + self.enable_dp_lm_head = cfg.enable_dp_lm_head + self.attn_cp_size = cfg.attn_cp_size + self.moe_dense_tp_size = cfg.moe_dense_tp_size + self.moe_a2a_backend = cfg.moe_a2a_backend + self.deepep_mode = cfg.deepep_mode + self.load_format = cfg.load_format + self.dtype = cfg.dtype + self.quantization = cfg.quantization + self.model_loader_extra_config = cfg.model_loader_extra_config + self.trust_remote_code = cfg.trust_remote_code + self.revision = cfg.revision self.dist_init_method = dist_init_method self.socket_path = get_socket_path( @@ -223,7 +226,7 @@ def _init_distributed(self, server_args, model_config): distributed_init_method=self.dist_init_method, local_rank=self.gpu_id, backend=current_platform.get_torch_distributed_backend_str(), - moe_a2a_backend=server_args.moe_a2a_backend, + moe_a2a_backend=self.moe_a2a_backend, ) initialize_model_parallel( @@ -281,7 +284,7 @@ def load(self): from sglang.srt.layers.moe import initialize_moe_config - initialize_moe_config(server_args) + initialize_moe_config() # Initialize distributed backend for model loading # (must be done after server_args and model_config are available) @@ -672,23 +675,24 @@ def launch_weight_cache_daemons( --nnodes 2 --node-rank 1 \\ --dist-init-method tcp://node0-ip:29500 """ + cfg = resolving_view(server_args) import socket as sock_mod # Replicate _calculate_rank_ranges logic from engine.py - pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1) - nnodes_per_pp_rank = max(server_args.nnodes // server_args.pp_size, 1) + pp_size_per_node = max(cfg.pp_size // cfg.nnodes, 1) + nnodes_per_pp_rank = max(cfg.nnodes // cfg.pp_size, 1) pp_rank_range = range( - pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank), - pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank + 1), + pp_size_per_node * (cfg.node_rank // nnodes_per_pp_rank), + pp_size_per_node * (cfg.node_rank // nnodes_per_pp_rank + 1), ) nnodes_per_tp_group = nnodes_per_pp_rank - tp_size_per_node = server_args.tp_size // nnodes_per_tp_group + tp_size_per_node = cfg.tp_size // nnodes_per_tp_group tp_rank_range = range( - tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group), - tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1), + tp_size_per_node * (cfg.node_rank % nnodes_per_tp_group), + tp_size_per_node * (cfg.node_rank % nnodes_per_tp_group + 1), ) - if server_args.nnodes > 1 and dist_init_method is None: + if cfg.nnodes > 1 and dist_init_method is None: raise ValueError( "dist_init_method is required for multi-node weight cache daemons. " "Use --dist-init-method tcp://: to specify the " @@ -705,7 +709,7 @@ def launch_weight_cache_daemons( # Validate and clean up stale .ready/.sock files from prior runs. for pp_rank in pp_rank_range: for tp_rank in tp_rank_range: - global_rank = compute_global_rank(server_args.tp_size, pp_rank, tp_rank) + global_rank = compute_global_rank(cfg.tp_size, pp_rank, tp_rank) cleanup_stale_daemon_files(global_rank, force=force) procs = [] @@ -716,8 +720,8 @@ def launch_weight_cache_daemons( tp_rank, pp_size_per_node, tp_size_per_node, - base_gpu_id=server_args.base_gpu_id, - gpu_id_step=server_args.gpu_id_step, + base_gpu_id=cfg.base_gpu_id, + gpu_id_step=cfg.gpu_id_step, ) proc = spawn_weight_cache_daemon( server_args, @@ -738,7 +742,7 @@ def launch_weight_cache_daemons( start_time = time.time() for pp_rank in pp_rank_range: for tp_rank in tp_rank_range: - global_rank = compute_global_rank(server_args.tp_size, pp_rank, tp_rank) + global_rank = compute_global_rank(cfg.tp_size, pp_rank, tp_rank) ready_path = get_ready_path(global_rank) while not os.path.exists(ready_path): time.sleep(check_interval) @@ -772,7 +776,7 @@ def launch_weight_cache_daemons( ) logger.info( - f"All {num_daemons} weight cache daemons on node {server_args.node_rank} are ready " + f"All {num_daemons} weight cache daemons on node {cfg.node_rank} are ready " 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})" diff --git a/python/sglang/test/scripted_runtime/scheduler_hook.py b/python/sglang/test/scripted_runtime/scheduler_hook.py index 58253cfcb8d5..51462b5d43ea 100644 --- a/python/sglang/test/scripted_runtime/scheduler_hook.py +++ b/python/sglang/test/scripted_runtime/scheduler_hook.py @@ -10,6 +10,7 @@ import zmq +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.environ import envs from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle from sglang.srt.utils.network import get_zmq_socket @@ -56,8 +57,8 @@ def _drive_engine_through_warmup(ctx: ScriptedContext) -> Generator: """Run the engine until the server warmup request has been received and fully processed, so scripts never observe foreign warmup traffic.""" scheduler = ctx.scheduler - server_args = scheduler.server_args - if server_args.skip_server_warmup: + cfg = resolving_view(scheduler.server_args) + if cfg.skip_server_warmup: logger.info("scripted_runtime: skip_server_warmup set, not driving warmup") return @@ -67,7 +68,7 @@ def _drive_engine_through_warmup(ctx: ScriptedContext) -> Generator: # is_fully_idle() can transiently report idle while a PP microbatch result # is still in flight, so require it to hold for two full microbatch # rotations after the warmup request was observed on the recv socket. - quiesce_iters = 2 * (server_args.pp_size + server_args.pp_async_batch_depth) + quiesce_iters = 2 * (cfg.pp_size + cfg.pp_async_batch_depth) proxy = ctx._tokenizer_recv_proxy deadline = start_time + WARMUP_DRIVE_TIMEOUT_S diff --git a/test/manual/ep/test_flashinfer_dispatcher.py b/test/manual/ep/test_flashinfer_dispatcher.py index cbb6bccdfa89..8bb540f6646d 100644 --- a/test/manual/ep/test_flashinfer_dispatcher.py +++ b/test/manual/ep/test_flashinfer_dispatcher.py @@ -10,6 +10,7 @@ from sglang.srt.layers.dp_attention import set_dp_buffer_len from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatcher from sglang.srt.layers.moe.utils import initialize_moe_config +from sglang.srt.runtime_context import publish from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler from sglang.test.test_utils import CustomTestCase @@ -22,7 +23,8 @@ def setUpClass(cls): server_args.moe_runner_backend = "flashinfer_cutlass" server_args.moe_a2a_backend = "flashinfer" set_global_server_args_for_scheduler(server_args) - initialize_moe_config(server_args) + publish(server_args, role="scheduler") + initialize_moe_config() init_distributed_environment( world_size=-1, # Auto-detect from environment diff --git a/test/registered/unit/spec/test_draft_construction_isolation.py b/test/registered/unit/spec/test_draft_construction_isolation.py index 0469ed4e6955..5609009d90a7 100644 --- a/test/registered/unit/spec/test_draft_construction_isolation.py +++ b/test/registered/unit/spec/test_draft_construction_isolation.py @@ -19,7 +19,7 @@ speculative_moe_backend_context, ) from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_context, get_flags, get_model +from sglang.srt.runtime_context import get_context, get_flags, get_model, publish from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -145,9 +145,11 @@ def test_initialize_moe_config_seeds_both_leaves(self): from sglang.srt.server_args import ServerArgs self._seed() - initialize_moe_config( - ServerArgs(model_path="dummy", disable_shared_experts_fusion=True) + publish( + ServerArgs(model_path="dummy", disable_shared_experts_fusion=True), + role="scheduler", ) + initialize_moe_config() moe = get_flags().moe self.assertTrue(moe.disable_shared_experts_fusion) self.assertTrue(moe.speculative_disable_shared_experts_fusion) diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index de01ae6191fb..b183374b63cf 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -57,6 +57,14 @@ "the lazy strategy bind in a worker: the CP group is what the strategy " "is being built for, and the configured width is what describes it" ), + ("benchmark/one_batch.py", "pp_size"): ( + "CPU affinity for this rank, computed right after the work function " + "publishes and before dist init, so the groups do not exist yet" + ), + ("benchmark/one_batch.py", "tp_size"): ( + "the same affinity computation: the layout is the configured one, and " + "the live group is not up at this point in the work function" + ), ("srt/entrypoints/engine.py", "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" @@ -256,6 +264,14 @@ "the receiver labels and shards by the launch width; it runs in the " "tokenizer process, which holds no encoder groups" ), + ("srt/managers/rust_server.py", "tp_size"): ( + "the rust server decides its transport from the launch width, in the " + "tokenizer process, which holds no model groups" + ), + ("compile_deep_gemm.py", "tp_size"): ( + "the warm-up request fans bootstrap rooms across the launch's ranks; it " + "runs in the tokenizer process, which holds no model groups" + ), ("srt/utils/common.py", "tp_size"): ( "the require_*_tp_gather predicates compared the configured tp_size " "when they read the record; the live property answers a different " diff --git a/test/registered/unit/test_ray_driver_reads_the_bags.py b/test/registered/unit/test_ray_driver_reads_the_bags.py index 948a61be403d..7532334e7294 100644 --- a/test/registered/unit/test_ray_driver_reads_the_bags.py +++ b/test/registered/unit/test_ray_driver_reads_the_bags.py @@ -1,10 +1,10 @@ """The Ray driver sizes its actors from the published configuration. `RayEngine` publishes as part of `Engine._launch_subprocesses` and *then* lays -out the actors, so the placement arithmetic reads the `parallel` bag rather than -the record it was constructed from. That matters because the record holds the -raw input: a launch that leaves `dp_size` to resolution would size the placement -group from `None`. +out the actors, so the placement arithmetic reads the `parallel` bag. That is +where a resolution decision lives: a launch that leaves `dp_size` to resolution +has it in the `parallel` bag, and the override case below is what tells the two +apart. There is no CI coverage of the Ray path (`test/manual/test_ray_engine.py` boots a real cluster), so these cases drive the two pure helpers directly against a diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index 11d4fd836a92..fbcc472f1c2e 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -519,8 +519,6 @@ class TestMoeFlagsGroup(_IsolatedServerArgs): swap under the speculative contexts and restore on exit.""" def _init(self, **kw): - from types import SimpleNamespace - from sglang.srt.layers.moe.utils import initialize_moe_config defaults = dict( @@ -538,7 +536,12 @@ def _init(self, **kw): disable_shared_experts_fusion=False, ) defaults.update(kw) - initialize_moe_config(SimpleNamespace(**defaults)) + # The flags are seeded from the bags, so the test publishes a config + # carrying these values. + override = get_context().override_server_args(**defaults) + override.install() + self.addCleanup(override.restore) + initialize_moe_config() def test_lazy_defaults_before_initialize(self): from sglang.srt.layers.moe.utils import ( diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index 50d8cc893f01..95107549f366 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -134,49 +134,12 @@ # are step-12 exposure like any other pair. _PASSED = frozenset({"model_path", "device", "random_seed"}) -# What is left reads the record because the record is the only thing that exists -# at that point in the process. Everything else moved to the bags or to -# `resolving_view`; a new entry here needs the same kind of reason. -# -# engine.py / enable_symm_mem -# `_set_envs_and_config` sets NCCL environment variables before anything -# publishes -- there is no bag to read yet. -# engine.py / reasoning_parser, tool_call_parser -# template_detection.py / model_path -# the auto-parser detection is late resolution: it runs in the launcher's -# validation stage, decides these fields and writes them *in place*, -# before the publish that would give it a bag. -_EXPOSED = { - ("layers/moe/utils.py", "deepep_mode"), - ("layers/moe/utils.py", "disable_shared_experts_fusion"), - ("layers/moe/utils.py", "moe_a2a_backend"), - ("layers/moe/utils.py", "moe_runner_backend"), - ("layers/moe/utils.py", "quantization"), - ("layers/moe/utils.py", "speculative_moe_runner_backend"), - ("entrypoints/engine.py", "enable_symm_mem"), - ("entrypoints/engine.py", "reasoning_parser"), - ("entrypoints/engine.py", "tool_call_parser"), - ("layers/moe/utils.py", "deepep_mode"), - ("layers/moe/utils.py", "moe_a2a_backend"), - ("layers/moe/utils.py", "moe_runner_backend"), - ("layers/moe/utils.py", "quantization"), - ("layers/moe/utils.py", "speculative_moe_runner_backend"), - ("weight_cache/daemon.py", "attn_cp_size"), - ("weight_cache/daemon.py", "deepep_mode"), - ("weight_cache/daemon.py", "dp_size"), - ("weight_cache/daemon.py", "dtype"), - ("weight_cache/daemon.py", "enable_dp_attention"), - ("weight_cache/daemon.py", "enable_dp_lm_head"), - ("weight_cache/daemon.py", "ep_size"), - ("weight_cache/daemon.py", "load_format"), - ("weight_cache/daemon.py", "model_loader_extra_config"), - ("weight_cache/daemon.py", "model_path"), - ("weight_cache/daemon.py", "moe_a2a_backend"), - ("weight_cache/daemon.py", "moe_dense_tp_size"), - ("weight_cache/daemon.py", "moe_dp_size"), - ("weight_cache/daemon.py", "pp_size"), - ("weight_cache/daemon.py", "quantization"), -} +# Empty. A pair belongs here when a reader has no bag to read -- it runs before +# its process publishes -- and cannot use `resolving_view` either. The launcher's +# pre-publish reads (`_set_envs_and_config`, the auto-parser gate) and the +# late-resolution detection it calls all read the declarations now, so nothing +# qualifies. A new entry needs that kind of reason next to it. +_EXPOSED: frozenset = frozenset() # Pairs whose resolution write only happens on a CUDA host (capability or # `is_cuda()` gated): asserted on the CUDA registration, invisible to the CPU @@ -189,20 +152,7 @@ # Axis two: (file, field) pairs where a supplied-instance read names a field that # some code overrides post-publish. Each needs an ordering judgment, not a blanket # conversion; the list exists so a new one is a decision made when it is written. -_OVERRIDDEN_AND_READ = { - ("entrypoints/engine.py", "reasoning_parser"), - ("entrypoints/engine.py", "tool_call_parser"), - ("weight_cache/daemon.py", "dp_size"), - ("weight_cache/daemon.py", "dtype"), - ("weight_cache/daemon.py", "ep_size"), - ("weight_cache/daemon.py", "load_format"), - ("weight_cache/daemon.py", "model_path"), - ("weight_cache/daemon.py", "dp_size"), - ("weight_cache/daemon.py", "dtype"), - ("weight_cache/daemon.py", "ep_size"), - ("weight_cache/daemon.py", "load_format"), - ("weight_cache/daemon.py", "model_path"), -} +_OVERRIDDEN_AND_READ: frozenset = frozenset() def _expanded_override_keys(rel, tree, call, kw) -> set: From fc344489a41e091ef304c7592df5915caabd0bc6 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 10:41:37 +0000 Subject: [PATCH 23/26] config: declarations stay in the stash instead of being replayed onto the fields Resolution ended by applying the whole declaration stash back onto the record, so the fields carried the resolved configuration and a post-resolution reader could read either. That replay is what made "server_args holds the raw input" untrue for the resolvers that only declare -- the model-specific overrides and the registry entries, which write nothing themselves. The replay is gone. Measured across 18 launch shapes, seven fields change hands because of it (`attention_backend`, `disable_overlap_schedule`, `enable_tp_lm_head_all_to_all`, `moe_a2a_backend`, `page_size`, `sampling_backend`, `speculative_moe_runner_backend`): the record now answers with what the operator passed and the bags answer with what resolution decided. Every other field is unchanged, because `declare_resolution` still writes as it declares. The readers that were getting the resolved value off the record follow: * `KVCacheConfigurator` reads `get_schedule().page_size`, which is what the rest of that file already did. * `check_server_args`, `check_lora_server_args`, `_check_two_batch_overlap`, `describe_kv_events_publisher` and `get_attention_backends` read through `resolved_view(self)` -- a validator or a derived member has to answer for the configuration resolution decided, not for the fields. * `compute_world_size` takes the resolved topology (the `parallel` bag) and the `/get_internal_state` readback hands it one. `enable_dp_attention` and `dp_size` are both resolution's answers, so a raw read reported `dp_size * tp * pp` for a DeepSeek MLA context-parallel server that runs `tp * pp`. The Ray driver's copy of the formula delegates to it. * The MiniMax sparse backend reads `get_spec()` for the draft-token count and the speculative algorithm. The count is auto-filled by resolution (EAGLE `steps + 1`, ngram 12), so a raw read left it `None` and the NPU verify-metadata cache silently skipped graph capture. `_declarations_materialized` is now `_resolution_finished`: it still arms the read-only `__setattr__`, but there is no materialization for it to name. The golden model-override tests move with it -- they assert the published leaf (`config_leaf`) and the projection, which is where an override lands, instead of the record field it used to be replayed onto. Touching the gateway is what makes CI run its Rust lints on this stack, and the test helpers in `cache_aware.rs` fail one the current clippy added (`needless_borrows_for_generic_args`): the two borrows are dropped here, so the PR that activates the job is the one that leaves it green. The gateway's per-worker copy moves to `replace_resolved`: it keyed off the old flag name, and with the replay gone a plain `setattr` of `port` / `base_gpu_id` / `dp_size` is what the read-only record refuses. Both channels are covered: the existing gateway test keeps the older-wheel fallback, and a second one hands it a record that carries `replace_resolved` and asserts the three per-worker values reach the child as a declaration while the parent keeps what the operator passed. --- python/sglang/benchmark/one_batch.py | 2 +- python/sglang/srt/arg_groups/arg_utils.py | 8 +- python/sglang/srt/arg_groups/overrides.py | 39 ++--- .../attention/minimax_sparse_backend.py | 9 +- python/sglang/srt/managers/scheduler.py | 3 +- .../srt/mem_cache/kv_cache_configurator.py | 4 +- python/sglang/srt/ray/engine.py | 8 +- python/sglang/srt/server_args.py | 53 +++--- .../python/src/sglang_router/launch_server.py | 23 +-- .../python/tests/test_startup_sequence.py | 57 +++++++ sgl-model-gateway/src/policies/cache_aware.rs | 4 +- .../test_scheduler_internal_state_env_vars.py | 5 +- ...est_scheduler_internal_state_world_size.py | 42 ++--- .../test_resolution_declarations.py | 71 +++++--- .../test_resolution_is_reproducible.py | 4 +- test/registered/unit/test_model_overrides.py | 155 ++++++++++++------ .../unit/test_runtime_context_config_bags.py | 39 ++--- .../unit/test_runtime_context_override.py | 2 +- ...test_supplied_instance_exposure_ratchet.py | 6 +- 19 files changed, 329 insertions(+), 205 deletions(-) diff --git a/python/sglang/benchmark/one_batch.py b/python/sglang/benchmark/one_batch.py index 55df78a61b98..f72faea387cf 100644 --- a/python/sglang/benchmark/one_batch.py +++ b/python/sglang/benchmark/one_batch.py @@ -1005,7 +1005,7 @@ def main(server_args, bench_args): # The decode phase has to capture the batch sizes this run benchmarks, and # the per-phase convenience knob loses to an explicit --cuda-graph-config # JSON (resolution applies that last), so the size is merged into that JSON. - if getattr(server_args, "_declarations_materialized", False): + if getattr(server_args, "_resolution_finished", False): # A record the caller already resolved: nothing will parse a raw dict # again, so the declaration has to be the finished typed config. merged = resolution_result(server_args, "cuda_graph_config") diff --git a/python/sglang/srt/arg_groups/arg_utils.py b/python/sglang/srt/arg_groups/arg_utils.py index bc209aa271e8..17e4235af8cb 100644 --- a/python/sglang/srt/arg_groups/arg_utils.py +++ b/python/sglang/srt/arg_groups/arg_utils.py @@ -75,10 +75,10 @@ class Arg: # When True, this field is skipped by add_cli_args_from_dataclass. # Use for fields that have no CLI surface (e.g. injected via Python only). no_cli: bool = False - # When True, this field may be written by config resolution (model - # overrides and post-process passes): it is part of the whitelist accepted - # by the declaration stash, and its resolved value materializes onto the - # field at the end of __post_init__. + # When True, config resolution (model overrides and post-process passes) + # may decide this field: the declaration stash accepts the name, and + # `resolution_result` and the config bags answer with the decision. The + # field keeps what the operator passed. resolvable: bool = False diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 556479e54850..dc7e826db501 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -206,10 +206,9 @@ def register_post_process(fn: Callable[..., dict]) -> Callable[..., dict]: def _declaration_overlay(server_args: Any) -> Dict[str, Any]: """What the declarations say so far, last writer wins. - Passes declare without touching the fields until - ``materialize_declarations``, so a mid-resolution reader needs this to see - them; handlers and hooks write as they declare, and for those the overlay - repeats what the field already holds.""" + Passes declare without touching the fields, so a mid-resolution reader + needs this to see them; handlers and hooks write as they declare, and for + those the overlay repeats what the field already holds.""" overlay: Dict[str, Any] = {} for _source, declared in getattr(server_args, "_resolved_overrides", None) or (): overlay.update(declared) @@ -222,9 +221,9 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None: Evaluates the pass on the resolving state (a read-only view with the accumulated declarations overlaid from the stash) and appends its declaration to the stash. During ``__post_init__`` the fields stay - untouched — ``materialize_declarations`` applies the whole stash once at - the end of resolution; a pass invoked after materialization (a post-init - slot) writes through immediately. + untouched: the stash is what the config bags are projected from. A pass + invoked after resolution finished (a post-init slot) writes through + immediately, because there is no later projection to pick it up. """ declared = fn(ResolvedView(server_args, overlay=_declaration_overlay(server_args))) if not isinstance(declared, dict): @@ -244,7 +243,7 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None: stash = server_args._resolved_overrides = [] stash.append(entry) validate_declarations(server_args, [entry]) - if getattr(server_args, "_declarations_materialized", False): + if getattr(server_args, "_resolution_finished", False): _apply_fields(server_args, declared) @@ -384,18 +383,6 @@ def declare_direct_writes( return result -def materialize_declarations(server_args: Any) -> None: - """Apply the accumulated declarations onto ``server_args`` once, at the - end of ``__post_init__`` (gate order: last writer wins). After this the - fields carry the resolved configuration — every post-init reader, in any - process, reads them directly; ``resolved_view`` remains an internal - helper for mid-resolution code only.""" - for _source, declared in getattr(server_args, "_resolved_overrides", None) or (): - for field, value in declared.items(): - setattr(server_args, field, value) - server_args._declarations_materialized = True - - def resolution_result(server_args: Any, field: str, default: Any = None) -> Any: """What resolution decided for ``field``: the declaration if there is one, otherwise what the caller supplied. @@ -453,10 +440,14 @@ def _plain(value: Any) -> Any: def resolved_view(server_args: Any) -> ResolvedView: - """Read-only view of the resolving configuration for mid-resolution code - that is not a pass (``__post_init__`` handlers and hooks). Internal to - the resolution pipeline: after ``materialize_declarations`` runs, the - fields themselves carry the resolved values — read them directly.""" + """Read-only view of the resolving configuration: the declarations + overlaid on the fields, snapshotted per call. + + For mid-resolution code that is not a pass (``__post_init__`` handlers and + hooks), and for the record's own members that must answer with what + resolution decided -- a declaration-only resolver (a model-specific + override, a registry entry) never writes the field, so a field read there + answers with the raw input.""" return ResolvedView(server_args, overlay=_declaration_overlay(server_args)) diff --git a/python/sglang/srt/layers/attention/minimax_sparse_backend.py b/python/sglang/srt/layers/attention/minimax_sparse_backend.py index 49e617e39232..55808c74f214 100644 --- a/python/sglang/srt/layers/attention/minimax_sparse_backend.py +++ b/python/sglang/srt/layers/attention/minimax_sparse_backend.py @@ -247,11 +247,10 @@ def __init__(self, runner: ModelRunner): Phase, check_cuda_graph_backend, ) + from sglang.srt.runtime_context import get_spec - _sa = getattr(runner, "server_args", None) - self.speculative_num_draft_tokens = getattr( - _sa, "speculative_num_draft_tokens", None - ) + spec = get_spec() + self.speculative_num_draft_tokens = spec.speculative_num_draft_tokens _decode_cuda_graph = not check_cuda_graph_backend( Phase.DECODE, Backend.DISABLED ) @@ -264,7 +263,7 @@ def __init__(self, runner: ModelRunner): if ( self.use_msa and _decode_cuda_graph - and getattr(_sa, "speculative_algorithm", None) is not None + and spec.speculative_algorithm is not None ): raise NotImplementedError( "MiniMax-M3 MSA attention does not support speculative decoding under " diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index fd3a27bf5b45..bc50b2fcaf01 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -39,7 +39,6 @@ get_observability, get_parallel, get_schedule, - get_server_args, get_serving, get_spec, ) @@ -4412,7 +4411,7 @@ def get_internal_state(self, recv_req: GetInternalStateReq): # Resolved config (pristine server_args + post-publish overrides) so a # readback reflects values changed via /set_internal_state, not startup. ret = get_context().resolved_server_args_dict() - ret["world_size"] = compute_world_size(get_server_args()) + ret["world_size"] = compute_world_size(get_parallel().config) ret["last_gen_throughput"] = self.metrics_reporter.last_gen_throughput draft_graph_memory_usage = ( None if self.draft_worker is None else self.draft_worker.graph_memory_usage diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 75821be4df43..16f50edc8b34 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -1380,7 +1380,7 @@ def _build_hybrid_mla_swa_kv_pool( """ full_pool_class = DSATokenToKVPool if is_dsa_model else MLATokenToKVPool common = { - "page_size": self.server_args.page_size, + "page_size": get_schedule().page_size, "device": self.device, "enable_memory_saver": False, } @@ -1401,7 +1401,7 @@ def _build_hybrid_mla_swa_kv_pool( return SWAKVPool( size=full_max_total_num_tokens, size_swa=swa_max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, head_num=0, head_dim=0, diff --git a/python/sglang/srt/ray/engine.py b/python/sglang/srt/ray/engine.py index 7be4b194aba0..0b6505e35d44 100644 --- a/python/sglang/srt/ray/engine.py +++ b/python/sglang/srt/ray/engine.py @@ -35,7 +35,7 @@ from sglang.srt.runtime_context import ( get_parallel, ) -from sglang.srt.server_args import PortArgs, ServerArgs +from sglang.srt.server_args import PortArgs, ServerArgs, compute_world_size logger = logging.getLogger(__name__) @@ -108,14 +108,10 @@ def get_node_ip(): def _compute_world_size() -> int: """Compute world_size (total number of scheduler actors/GPUs needed). - Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size. Reads the published parallel leaves: the driver is sizing the actors that will hold the process groups, so there is nothing live to ask. """ - parallel = get_parallel().config - if parallel.enable_dp_attention: - return parallel.tp_size * parallel.pp_size - return parallel.dp_size * parallel.tp_size * parallel.pp_size + return compute_world_size(get_parallel().config) def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]: diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a90733df68a8..295b52c567b7 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -3709,7 +3709,7 @@ def resolve_once(self) -> None: arrived by pickle and brought its declarations along, so the child has nothing left to derive and projects what the parent decided. """ - if getattr(self, "_declarations_materialized", False): + if getattr(self, "_resolution_finished", False): return if getattr(self, "_resolution_failed", False): raise RuntimeError( @@ -3728,7 +3728,7 @@ def resolve_once(self) -> None: # Set here too, because the dummy/absent-model path returns before the # materialization that normally sets it: the gate is about whether the # handlers ran, not how far they got. - self._declarations_materialized = True + self._resolution_finished = True def resolved_dict(self) -> Dict[str, Any]: """This configuration as a plain dict of resolved field values. @@ -3770,7 +3770,7 @@ def replace_resolved(self, source: str, **changes: Any) -> ServerArgs: copy's deep structure in-process mutates the parent's too. """ replacement = dataclasses.replace(self, **changes) - if not getattr(self, "_declarations_materialized", False): + if not getattr(self, "_resolution_finished", False): # Not resolved yet: the copy goes through the gate itself. return replacement @@ -3780,7 +3780,7 @@ def replace_resolved(self, source: str, **changes: Any) -> ServerArgs: # (the read-only guard refuses the write). field_names = {field.name for field in dataclasses.fields(self)} for name, value in vars(self).items(): - if name in field_names or name == "_declarations_materialized": + if name in field_names or name == "_resolution_finished": continue if isinstance(value, (dict, list, set)): value = copy.copy(value) @@ -3791,7 +3791,7 @@ def replace_resolved(self, source: str, **changes: Any) -> ServerArgs: object.__setattr__(replacement, "_resolved_overrides", stash) if changes: stash.append((source, dict(changes))) - object.__setattr__(replacement, "_declarations_materialized", True) + object.__setattr__(replacement, "_resolution_finished", True) return replacement def _declare(self, source: str, **fields: Any) -> None: @@ -4022,13 +4022,7 @@ def _run_resolution_pipeline(self): # time; last declarations of the resolution, mirroring that order. self._handle_model_capability_adjustments() - # End of resolution: apply the accumulated declarations onto the - # fields once (gate order). From here on server_args carries the - # resolved configuration — post-init readers, in any process, read - # the fields directly. - from sglang.srt.arg_groups.overrides import materialize_declarations - - materialize_declarations(self) + self._resolution_finished = True def _handle_return_hidden_states_mode(self): cfg = resolving_view(self) @@ -9809,7 +9803,7 @@ def __setattr__(self, name, value): # get_context().override(source, ...); a value one runner or worker # owns travels as a constructor argument to it. if ( - getattr(self, "_declarations_materialized", False) + getattr(self, "_resolution_finished", False) and not getattr(self, "_internal_write", False) and name not in _CACHE_SLOTS and (not name.startswith("_") or name in _underscore_field_names()) @@ -9832,7 +9826,13 @@ def _resolved_attention_backends(self): return attention_backends_of(resolved_view(self)) def get_attention_backends(self): - return attention_backends_of(self) + """The (prefill, decode) pair resolution decided. + + Reads through the declaration stash, not the fields: the model-specific + overrides declare into the stash without writing the fields, so a field + read answers with what the operator typed. + """ + return attention_backends_of(resolved_view(self)) def use_mla_backend(self): from sglang.srt.configs.model_config import AttentionArch @@ -9884,7 +9884,7 @@ def max_speculative_num_draft_tokens(self) -> Optional[int]: # state needs steps + 1 draft-token slots. Revisit this if topk>1 # is supported. result = max(candidate_steps) + 1 - if getattr(self, "_declarations_materialized", False): + if getattr(self, "_resolution_finished", False): object.__setattr__(self, "_max_speculative_num_draft_tokens", result) return result @@ -9913,7 +9913,7 @@ def mamba_cache_chunk_size(self) -> int: assert ( max(chunk_size, page_size) % min(chunk_size, page_size) == 0 ), f"For SSM models, either chunk_size or page_size must be divisible by the other, got {chunk_size=}, {page_size=}" - if not getattr(self, "_declarations_materialized", False): + if not getattr(self, "_resolution_finished", False): return max(chunk_size, page_size) self._mamba_cache_chunk_size = max(chunk_size, page_size) return self._mamba_cache_chunk_size @@ -10574,8 +10574,8 @@ def describe_kv_events_publisher(self) -> Optional[dict]: "endpoint_host": host, "endpoint_port_base": port, "topic": cfg.topic, - "block_size": self.kv_event_block_size, - "dp_size": self.dp_size, + "block_size": resolved.kv_event_block_size, + "dp_size": resolved.dp_size, } def should_report_expert_balancedness(self) -> bool: @@ -10593,12 +10593,19 @@ def should_export_expert_balancedness_to_prometheus(self) -> bool: return cfg.expert_balancedness_report_mode in ("prometheus", "both") -def compute_world_size(server_args: ServerArgs) -> int: - """Return the total GPU count across all data-parallel replicas.""" +def compute_world_size(config) -> int: + """Return the total GPU count across all data-parallel replicas. + + Takes the resolved topology -- the published `parallel` bag, or a view over + the declarations. `enable_dp_attention` and `dp_size` are both resolution's + answers (`_handle_dwdp` fills the pair, DeepSeek MLA context parallelism + turns DP attention on), so a raw-record read would size the world from what + the operator typed. + """ return ( - (1 if server_args.enable_dp_attention else server_args.dp_size) - * server_args.tp_size - * server_args.pp_size + (1 if config.enable_dp_attention else config.dp_size) + * config.tp_size + * config.pp_size ) diff --git a/sgl-model-gateway/bindings/python/src/sglang_router/launch_server.py b/sgl-model-gateway/bindings/python/src/sglang_router/launch_server.py index adc6036ae9fa..6743f949a009 100644 --- a/sgl-model-gateway/bindings/python/src/sglang_router/launch_server.py +++ b/sgl-model-gateway/bindings/python/src/sglang_router/launch_server.py @@ -87,22 +87,25 @@ def launch_server_process( server_args: ServerArgs, worker_port: int, dp_id: int ) -> mp.Process: """Launch a single server process with the given args and port.""" - # This binding is installed against a released sglang, so it cannot call - # into helpers newer than that wheel. Copy first, then write through the - # sanctioned channel if the record is resolved (a resolved record refuses - # plain assignment), else assign. - worker_args = copy.deepcopy(server_args) changes = { "port": worker_port, "base_gpu_id": dp_id * server_args.tp_size, "dp_size": 1, } - late = getattr(worker_args, "_late_resolution", None) - if late is not None and getattr(worker_args, "_declarations_materialized", False): - late("sglang_router.launch_server_process", **changes) + # Three channels, newest first. A wheel that has the read-only record but not + # `replace_resolved` still has `_late_resolution`, and plain assignment there + # raises; only a wheel with neither accepts `setattr`. + replace_resolved = getattr(server_args, "replace_resolved", None) + if replace_resolved is not None: + worker_args = replace_resolved("sglang_router.launch_server_process", **changes) else: - for field, value in changes.items(): - setattr(worker_args, field, value) + worker_args = copy.deepcopy(server_args) + late = getattr(worker_args, "_late_resolution", None) + if late is not None: + late("sglang_router.launch_server_process", **changes) + else: + for field, value in changes.items(): + setattr(worker_args, field, value) server_args = worker_args proc = mp.Process(target=run_server, args=(server_args, dp_id)) diff --git a/sgl-model-gateway/bindings/python/tests/test_startup_sequence.py b/sgl-model-gateway/bindings/python/tests/test_startup_sequence.py index 8ae727d0257a..e659a4d8d0b1 100644 --- a/sgl-model-gateway/bindings/python/tests/test_startup_sequence.py +++ b/sgl-model-gateway/bindings/python/tests/test_startup_sequence.py @@ -856,6 +856,63 @@ def fake_killpg(pid, sig): assert (p1.pid, _sig.SIGTERM) in calls and (p2.pid, _sig.SIGTERM) in calls assert (p2.pid, _sig.SIGKILL) in calls + +def test_launch_server_process_declares_on_a_resolved_record(monkeypatch): + """A record that carries its resolution takes the declaration channel. + + The stub above has no `replace_resolved`, so it exercises the older wheels' + path. A current `ServerArgs` refuses plain assignment once resolution has + finished; the per-worker values reach the child as a declaration on a copy, + and the parent keeps what the operator passed. + """ + _install_sglang_stubs(monkeypatch) + import importlib + + ls = importlib.import_module("sglang_router.launch_server") + + created = {} + + class FakeProcess: + def __init__(self, target, args): + created["target"] = target + created["args"] = args + self.pid = 4243 + + def start(self): + created["started"] = True + + monkeypatch.setattr(ls.mp, "Process", FakeProcess) + + calls = [] + + class ResolvedServerArgs: + def __init__(self, **fields): + self.port = fields.get("port", 30000) + self.base_gpu_id = fields.get("base_gpu_id", 0) + self.dp_size = fields.get("dp_size", 4) + self.tp_size = fields.get("tp_size", 2) + + def replace_resolved(self, source, **changes): + calls.append((source, dict(changes))) + fields = dict(vars(self)) + fields.update(changes) + return ResolvedServerArgs(**fields) + + parent = ResolvedServerArgs() + proc = ls.launch_server_process(parent, worker_port=31002, dp_id=3) + + assert created.get("started") is True + assert proc.pid == 4243 + assert len(calls) == 1 + source, changes = calls[0] + assert source == "sglang_router.launch_server_process" + assert changes == {"port": 31002, "base_gpu_id": 6, "dp_size": 1} + + worker = created["args"][0] + assert (worker.port, worker.base_gpu_id, worker.dp_size) == (31002, 6, 1) + # The parent is untouched: the copy is what carries the worker's identity. + assert (parent.port, parent.base_gpu_id, parent.dp_size) == (30000, 0, 4) + def test_validation_error_handling(self): """Test error handling when validation fails.""" args = RouterArgs( diff --git a/sgl-model-gateway/src/policies/cache_aware.rs b/sgl-model-gateway/src/policies/cache_aware.rs index c1d7ee3dfd92..d0f5daa6b1e3 100644 --- a/sgl-model-gateway/src/policies/cache_aware.rs +++ b/sgl-model-gateway/src/policies/cache_aware.rs @@ -679,7 +679,7 @@ mod tests { let mut workers: Vec> = Vec::new(); for j in 0..num_workers { workers.push(Arc::new( - BasicWorkerBuilder::new(&format!("http://w{}:8000", j)) + BasicWorkerBuilder::new(format!("http://w{}:8000", j)) .worker_type(WorkerType::Regular) .build(), )); @@ -738,7 +738,7 @@ mod tests { let mut workers: Vec> = Vec::new(); for j in 0..num_workers { workers.push(Arc::new( - BasicWorkerBuilder::new(&format!("http://w{}:8000", j)) + BasicWorkerBuilder::new(format!("http://w{}:8000", j)) .worker_type(WorkerType::Regular) .build(), )); diff --git a/test/registered/unit/managers/test_scheduler_internal_state_env_vars.py b/test/registered/unit/managers/test_scheduler_internal_state_env_vars.py index 48606f9ba524..1d6ff3276397 100644 --- a/test/registered/unit/managers/test_scheduler_internal_state_env_vars.py +++ b/test/registered/unit/managers/test_scheduler_internal_state_env_vars.py @@ -48,10 +48,11 @@ def _get_internal_state(self) -> dict: ), patch( "sglang.srt.managers.scheduler.get_exec", return_value=SimpleNamespace(moe=SimpleNamespace(elastic_ep_backend=None)), - ), patch( - "sglang.srt.managers.scheduler.get_server_args", return_value=None ), patch( "sglang.srt.managers.scheduler.compute_world_size", return_value=1 + ), patch( + "sglang.srt.managers.scheduler.get_parallel", + return_value=SimpleNamespace(config=SimpleNamespace()), ): output = scheduler.get_internal_state(recv_req=GetInternalStateReq()) diff --git a/test/registered/unit/managers/test_scheduler_internal_state_world_size.py b/test/registered/unit/managers/test_scheduler_internal_state_world_size.py index e4c66f4785f9..a5141e29dc21 100644 --- a/test/registered/unit/managers/test_scheduler_internal_state_world_size.py +++ b/test/registered/unit/managers/test_scheduler_internal_state_world_size.py @@ -9,16 +9,16 @@ from sglang.srt.managers.io_struct import GetInternalStateReq from sglang.srt.managers.scheduler import Scheduler -from sglang.srt.server_args import ServerArgs, compute_world_size +from sglang.srt.server_args import compute_world_size register_cpu_ci(est_time=5, suite="base-a-test-cpu") -def _make_server_args( +def _make_parallel_config( *, tp_size: int, pp_size: int, dp_size: int, enable_dp_attention: bool -) -> ServerArgs: - return ServerArgs( - model_path="dummy", +) -> SimpleNamespace: + """The four `parallel` leaves the world size is computed from.""" + return SimpleNamespace( tp_size=tp_size, pp_size=pp_size, dp_size=dp_size, @@ -29,39 +29,39 @@ def _make_server_args( class TestComputeWorldSize(unittest.TestCase): def test_a_single_gpu_server_holds_one_gpu(self): """The default shape has to come out as one, or every consumer is off by a factor.""" - server_args = _make_server_args( + config = _make_parallel_config( tp_size=1, pp_size=1, dp_size=1, enable_dp_attention=False ) - self.assertEqual(compute_world_size(server_args), 1) + self.assertEqual(compute_world_size(config), 1) def test_tensor_and_pipeline_stages_multiply(self): """Each (pp_rank, tp_rank) pair is its own scheduler process on its own gpu.""" - server_args = _make_server_args( + config = _make_parallel_config( tp_size=2, pp_size=3, dp_size=1, enable_dp_attention=False ) - self.assertEqual(compute_world_size(server_args), 6) + self.assertEqual(compute_world_size(config), 6) def test_plain_data_parallel_replicas_each_hold_their_own_gpus(self): """Without dp attention every replica launches a full tensor-parallel group of its own.""" - server_args = _make_server_args( + config = _make_parallel_config( tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False ) - self.assertEqual(compute_world_size(server_args), 4) + self.assertEqual(compute_world_size(config), 4) def test_data_parallel_attention_shares_the_tensor_parallel_gpus(self): """With dp attention the dp ranks live inside the tensor-parallel world, not beside it.""" - server_args = _make_server_args( + config = _make_parallel_config( tp_size=4, pp_size=1, dp_size=2, enable_dp_attention=True ) - self.assertEqual(compute_world_size(server_args), 4) + self.assertEqual(compute_world_size(config), 4) class TestSchedulerInternalStateWorldSize(unittest.TestCase): - def _get_internal_state(self, server_args: ServerArgs) -> dict: + def _get_internal_state(self, config: SimpleNamespace) -> dict: scheduler = Scheduler.__new__(Scheduler) scheduler.metrics_reporter = SimpleNamespace( last_gen_throughput=1.0, @@ -94,8 +94,8 @@ def _get_internal_state(self, server_args: ServerArgs) -> dict: "sglang.srt.managers.scheduler.get_exec", return_value=SimpleNamespace(moe=SimpleNamespace(elastic_ep_backend=None)), ), patch( - "sglang.srt.managers.scheduler.get_server_args", - return_value=server_args, + "sglang.srt.managers.scheduler.get_parallel", + return_value=SimpleNamespace(config=config), ): output = scheduler.get_internal_state(recv_req=GetInternalStateReq()) @@ -103,24 +103,24 @@ def _get_internal_state(self, server_args: ServerArgs) -> dict: def test_the_internal_state_reports_the_whole_server(self): """A consumer sizing an external fleet reads the gpus the server occupies, not the declared sizes.""" - server_args = _make_server_args( + config = _make_parallel_config( tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False ) - internal_state = self._get_internal_state(server_args) + internal_state = self._get_internal_state(config) self.assertEqual(internal_state["world_size"], 4) def test_the_reported_size_is_not_one_replica_of_a_data_parallel_server(self): """Each plain dp replica has its own process group, so no scheduler can report the whole server from it.""" - server_args = _make_server_args( + config = _make_parallel_config( tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False ) - internal_state = self._get_internal_state(server_args) + internal_state = self._get_internal_state(config) self.assertNotEqual( - internal_state["world_size"], server_args.tp_size * server_args.pp_size + internal_state["world_size"], config.tp_size * config.pp_size ) diff --git a/test/registered/unit/server_args/test_resolution_declarations.py b/test/registered/unit/server_args/test_resolution_declarations.py index ee37b06cd4ec..5ce2c57e9f3d 100644 --- a/test/registered/unit/server_args/test_resolution_declarations.py +++ b/test/registered/unit/server_args/test_resolution_declarations.py @@ -233,6 +233,11 @@ def _bare_assignments(): return sorted(found) +def shape_key(shape): + """A shape rendered short enough for a failure message.""" + return ",".join(f"{k}={v}" for k, v in sorted(shape.items())) or "defaults" + + def _stash_overlay(server_args): """What the declarations say, last writer wins -- the projection's input.""" overlay = {} @@ -345,35 +350,43 @@ def test_the_stash_accounts_for_every_change_resolution_made(self): + "\n ".join(unexplained), ) - def test_the_projection_input_is_the_resolved_configuration(self): - """What the bags are built from equals what the record ends up holding. + def test_a_declaration_only_resolver_leaves_the_field_alone(self): + """The direction of travel: resolution decides, the record does not move. - The projection reads `raw input + declarations` rather than the - fields, so that it keeps working when the declarations stop - materializing. While they still do, the two have to agree leaf for - leaf -- a difference means the projection would publish something the - record does not say, which is the failure this whole transition is - meant to avoid. + A resolver that only declares -- a model-specific override, a registry + entry -- writes nothing onto the record. The projection carries its + answer and the field still holds what the caller passed. """ from sglang.srt.arg_groups.arg_utils import namespace_of from sglang.srt.arg_groups.overrides import resolution_result - differences = [] + found = [] for shape in _SHAPES: server_args = self._resolve(shape) + raw = getattr(server_args, "_raw_input", None) or {} for field in namespace_of(type(server_args)): - projected = resolution_result(server_args, field) + if field not in raw: + continue + decided = resolution_result(server_args, field) on_record = getattr(server_args, field) - if projected != on_record: - differences.append( - f"{shape} -> {field}: projection={projected!r} " - f"record={on_record!r}" - ) - self.assertEqual( - differences, + if decided == on_record: + continue + # It moved away from the record's value, so the record must + # still hold exactly what the caller passed. + self.assertEqual( + on_record, + raw[field], + f"{shape} -> {field}: the record holds {on_record!r}, which " + f"is neither the raw input {raw[field]!r} nor what " + f"resolution decided ({decided!r})", + ) + found.append((shape_key(shape), field)) + self.assertNotEqual( + found, [], - "the projection and the record disagree about a config leaf:\n " - + "\n ".join(differences), + "no field is resolved by declaration alone any more, so this check " + "no longer covers anything -- either the shapes stopped reaching " + "one or the declarations are writing the fields again", ) def test_the_whole_object_readback_carries_only_fields(self): @@ -661,7 +674,13 @@ def test_the_launcher_finishes_resolving_before_it_publishes(self): f"written:\n " + "\n ".join(too_late), ) - def test_the_stash_agrees_with_the_fields_it_declared(self): + def test_the_stash_is_what_the_projection_answers(self): + """The last declaration for a field is what `resolution_result` returns. + + The stash is the resolution result, and every bag is projected from + it, so a disagreement here means the projection is reading something + other than the declarations. + """ mismatches = [] for shape in _SHAPES: server_args = self._resolve(shape) @@ -669,16 +688,18 @@ def test_the_stash_agrees_with_the_fields_it_declared(self): for field, declared in overlay.items(): if field not in _RESOLVED_FIELDS: continue - actual = getattr(server_args, field) - if actual != declared: + answered = resolution_result(server_args, field) + if answered != declared: mismatches.append( - f"{shape} -> {field}: field={actual!r} stash={declared!r}" + f"{shape} -> {field}: projection={answered!r} " + f"stash={declared!r}" ) self.assertEqual( mismatches, [], - "a declared field and its stash entry disagree, so something " - "assigned the field behind the declaration:\n " + "\n ".join(mismatches), + "the projection disagrees with the last declaration for a field, so " + "the bags would publish something no resolver decided:\n " + + "\n ".join(mismatches), ) def test_no_immediate_writer_overrides_a_deferred_one(self): diff --git a/test/registered/unit/server_args/test_resolution_is_reproducible.py b/test/registered/unit/server_args/test_resolution_is_reproducible.py index dfeccd4b2333..0c0acabdfead 100644 --- a/test/registered/unit/server_args/test_resolution_is_reproducible.py +++ b/test/registered/unit/server_args/test_resolution_is_reproducible.py @@ -775,7 +775,7 @@ def test_a_bare_replace_would_resolve_a_second_time(self): parent = self._resolved() bare = dataclasses.replace(parent, dist_init_addr="1.2.3.4:5000") self.assertFalse( - getattr(bare, "_declarations_materialized", False), + getattr(bare, "_resolution_finished", False), "a bare replace carried the flag; then this test proves nothing", ) bare.resolve_once() @@ -792,7 +792,7 @@ def test_a_bare_replace_would_resolve_a_second_time(self): def test_replace_resolved_keeps_the_parents_resolution(self): parent = self._resolved() copy_ = parent.replace_resolved("ray.test", dist_init_addr="1.2.3.4:5000") - self.assertTrue(getattr(copy_, "_declarations_materialized", False)) + self.assertTrue(getattr(copy_, "_resolution_finished", False)) drifted = { field.name: (getattr(parent, field.name), getattr(copy_, field.name)) for field in dataclasses.fields(parent) diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index fbb40f2b4444..b02515befcf8 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -20,6 +20,7 @@ from sglang.srt.arg_groups.overrides import ( collect_model_override_declarations, register_model_override, + resolution_result, validate_declarations, ) from sglang.srt.configs.minicpm import MiniCPMHybridConfig @@ -294,9 +295,28 @@ def test_dummy_fixture_publishes_the_object_it_resolved(self): class TestGoldenModelOverrides(_IsolatedPublish): """Per-arch golden diff for migrated families: the declarative path must - reproduce the legacy imperative writes byte-identically on the - materialized server_args fields; the publish round-trip returns the same - object.""" + reproduce the legacy imperative writes byte-identically in the resolution + result; the publish round-trip returns the same object. + + `_resolved` is how the assertions read it. A model-specific override only + declares -- it does not write the field -- so the record keeps what the + caller passed and the projection carries the override. + """ + + def _resolved(self, server_args, field): + from sglang.srt.arg_groups.overrides import resolution_result + + return resolution_result(server_args, field) + + def _leaf(self, field): + """The published value of `field`, whichever bag owns it. + + The publish round-trip is checked on the bags: the record the process + publishes is the raw input, and the leaf is what every reader reads. + """ + from sglang.srt.runtime_context import get_context + + return get_context().config_leaf(field) _MINI_CONFIG = { "hidden_size": 64, @@ -571,40 +591,42 @@ def _publish(self, server_args): def test_mistral_large3_forces_bfloat16(self): sa = self._construct("MistralLarge3ForCausalLM", "mistral") - self.assertEqual(sa.dtype, "bfloat16") # materialized at end of resolution + self.assertEqual( + self._resolved(sa, "dtype"), "bfloat16" + ) # materialized at end of resolution self.assertIn( ("MODEL_OVERRIDES['MistralLarge3ForCausalLM']", {"dtype": "bfloat16"}), sa._resolved_overrides, ) - self.assertEqual(self._publish(sa).dtype, "bfloat16") + self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "bfloat16") def test_user_requested_dtype_is_still_overridden(self): # Legacy fidelity: the arch branch overwrote dtype unconditionally, - # so the declaration must too. The pristine request survives on - # provenance; the materialized field carries the override. + # so the declaration must too. The request survives on the record; the + # projection carries the override. sa = self._construct("MistralLarge3ForCausalLM", "mistral", dtype="float16") - self.assertEqual(sa.dtype, "bfloat16") # materialized - self.assertEqual(self._publish(sa).dtype, "bfloat16") + self.assertEqual(self._resolved(sa, "dtype"), "bfloat16") + self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "bfloat16") def test_control_arch_keeps_pristine_dtype(self): sa = self._construct("LlamaForCausalLM", "llama") - self.assertEqual(sa.dtype, "auto") + self.assertEqual(self._resolved(sa, "dtype"), "auto") declared = {f for _s, d in sa._resolved_overrides for f in d} self.assertNotIn("dtype", declared) # no arch declaration for Llama - # publish still materializes the whitelisted leaf with the pristine + # publish still projects the whitelisted leaf with the pristine # value: readers only ever read flags. - self.assertEqual(self._publish(sa).dtype, "auto") + self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "auto") def test_minimax_m2_enables_tf32_matmul(self): sa = self._construct("MiniMaxM2ForCausalLM", "llama") - self.assertTrue(sa.enable_tf32_matmul) # materialized + self.assertTrue(self._resolved(sa, "enable_tf32_matmul")) self.assertIn( ("_minimax_m2_overrides", {"enable_tf32_matmul": True}), sa._resolved_overrides, ) flags = self._publish(sa) - self.assertTrue(flags.enable_tf32_matmul) - self.assertFalse(flags.enable_multi_layer_eagle) # pristine materialize + self.assertTrue(self._leaf("enable_tf32_matmul")) + self.assertFalse(self._leaf("enable_multi_layer_eagle")) # the pristine value def test_minimax_m2_sm10x_nvfp4_uses_routed_trtllm(self): """MiniMax-M2 NVFP4 auto must avoid the unsupported plain TRT-LLM path.""" @@ -622,10 +644,14 @@ def test_minimax_m2_sm10x_nvfp4_uses_routed_trtllm(self): "MiniMaxM2ForCausalLM", "llama", quantization="modelopt_fp4" ) - self.assertEqual(explicit.moe_runner_backend, "flashinfer_cutlass") - self.assertEqual(non_nvfp4.moe_runner_backend, "auto") - self.assertEqual(nvfp4.moe_runner_backend, "flashinfer_trtllm_routed") - self.assertTrue(nvfp4.disable_shared_experts_fusion) + self.assertEqual( + self._resolved(explicit, "moe_runner_backend"), "flashinfer_cutlass" + ) + self.assertEqual(self._resolved(non_nvfp4, "moe_runner_backend"), "auto") + self.assertEqual( + self._resolved(nvfp4, "moe_runner_backend"), "flashinfer_trtllm_routed" + ) + self.assertTrue(self._resolved(nvfp4, "disable_shared_experts_fusion")) self.assertIn( ( "_minimax_m2_overrides", @@ -649,7 +675,7 @@ def test_minimax_m2_sm10x_nvfp4_uses_routed_trtllm(self): non_sm10x = self._construct( "MiniMaxM2ForCausalLM", "llama", quantization="modelopt_fp4" ) - self.assertEqual(non_sm10x.moe_runner_backend, "auto") + self.assertEqual(self._resolved(non_sm10x, "moe_runner_backend"), "auto") self._publish(nvfp4) self.assertEqual(get_exec().moe.moe_runner_backend, "flashinfer_trtllm_routed") @@ -1000,25 +1026,25 @@ def test_step3p_hierarchical_cache_golden(self): enable_hierarchical_cache=True, ) # materialized at the end of resolution - self.assertEqual(sa.swa_full_tokens_ratio, 1.0) - self.assertTrue(sa.disable_hybrid_swa_memory) + self.assertEqual(self._resolved(sa, "swa_full_tokens_ratio"), 1.0) + self.assertTrue(self._resolved(sa, "disable_hybrid_swa_memory")) flags = self._publish(sa) - self.assertEqual(flags.swa_full_tokens_ratio, 1.0) - self.assertTrue(flags.disable_hybrid_swa_memory) + self.assertEqual(self._leaf("swa_full_tokens_ratio"), 1.0) + self.assertTrue(self._leaf("disable_hybrid_swa_memory")) def test_gemma2_disables_hybrid_swa_memory(self): sa = self._construct("Gemma2ForCausalLM", "llama") - self.assertTrue(sa.disable_hybrid_swa_memory) # materialized + self.assertTrue(self._resolved(sa, "disable_hybrid_swa_memory")) # materialized self.assertIn( ("_gemma2_gemma3_overrides", {"disable_hybrid_swa_memory": True}), sa._resolved_overrides, ) - self.assertTrue(self._publish(sa).disable_hybrid_swa_memory) + self.assertTrue((self._publish(sa), self._leaf("disable_hybrid_swa_memory"))[1]) def test_olmo2_disables_hybrid_swa_memory(self): sa = self._construct("Olmo2ForCausalLM", "llama") - self.assertTrue(sa.disable_hybrid_swa_memory) # materialized - self.assertTrue(self._publish(sa).disable_hybrid_swa_memory) + self.assertTrue(self._resolved(sa, "disable_hybrid_swa_memory")) # materialized + self.assertTrue((self._publish(sa), self._leaf("disable_hybrid_swa_memory"))[1]) def test_exaone_conditional_on_sliding_window_pattern(self): # With the pattern the branch also asserts an explicit backend. @@ -1028,8 +1054,8 @@ def test_exaone_conditional_on_sliding_window_pattern(self): config_extra={"sliding_window_pattern": "LLLG"}, attention_backend="fa3", ) - self.assertTrue(sa.disable_hybrid_swa_memory) # materialized - self.assertTrue(self._publish(sa).disable_hybrid_swa_memory) + self.assertTrue(self._resolved(sa, "disable_hybrid_swa_memory")) # materialized + self.assertTrue((self._publish(sa), self._leaf("disable_hybrid_swa_memory"))[1]) def test_exaone_without_pattern_declares_nothing(self): from sglang.srt.arg_groups.overrides import _exaone_overrides @@ -1051,13 +1077,13 @@ def test_gpt_oss_mxfp4_forces_bfloat16(self): "llama", config_extra={"quantization_config": {"quant_method": "mxfp4"}}, ) - self.assertEqual(sa.dtype, "bfloat16") # materialized - self.assertEqual(self._publish(sa).dtype, "bfloat16") + self.assertEqual(self._resolved(sa, "dtype"), "bfloat16") + self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "bfloat16") def test_gpt_oss_without_mxfp4_keeps_pristine_dtype(self): sa = self._construct("GptOssForCausalLM", "llama") - self.assertEqual(sa.dtype, "auto") - self.assertEqual(self._publish(sa).dtype, "auto") + self.assertEqual(self._resolved(sa, "dtype"), "auto") + self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "auto") def test_gpt_oss_xpu_dtype_validation_reads_pristine(self): from sglang.srt.arg_groups.overrides import _gpt_oss_overrides @@ -1077,28 +1103,36 @@ def test_sampling_backend_default_pass(self): sa = self._construct("LlamaForCausalLM", "llama") expected = "flashinfer" if is_flashinfer_available() else "pytorch" - self.assertEqual(sa.sampling_backend, expected) # materialized + self.assertEqual( + self._resolved(sa, "sampling_backend"), expected + ) # materialized self.assertIn( ("_sampling_backend_default", {"sampling_backend": expected}), sa._resolved_overrides, ) - self.assertEqual(self._publish(sa).sampling_backend, expected) + self.assertEqual( + (self._publish(sa), self._leaf("sampling_backend"))[1], expected + ) def test_sampling_backend_user_choice_survives(self): sa = self._construct("LlamaForCausalLM", "llama", sampling_backend="pytorch") - self.assertEqual(sa.sampling_backend, "pytorch") + self.assertEqual(self._resolved(sa, "sampling_backend"), "pytorch") # the pass declared nothing; publish materializes the pristine choice - self.assertEqual(self._publish(sa).sampling_backend, "pytorch") + self.assertEqual( + (self._publish(sa), self._leaf("sampling_backend"))[1], "pytorch" + ) def test_deterministic_inference_forces_pytorch_sampling(self): sa = self._construct( "LlamaForCausalLM", "llama", enable_deterministic_inference=True ) - # two pass writers chain: default fill, then the deterministic force — - # last writer wins; materialization lands the end state on the fields. - self.assertEqual(sa.sampling_backend, "pytorch") + # two pass writers chain: default fill, then the deterministic force -- + # last writer wins. The end state lives in the stash, which is what the + # projection reads and the bags are built from; the field still holds + # what the caller passed. + self.assertEqual(resolution_result(sa, "sampling_backend"), "pytorch") flags = self._publish(sa) - self.assertEqual(flags.sampling_backend, "pytorch") + self.assertEqual(self._leaf("sampling_backend"), "pytorch") # the deterministic attention fill declared a compatible backend and # the compatibility default-fill then had nothing to do deterministic_fills = [ @@ -1107,8 +1141,10 @@ def test_deterministic_inference_forces_pytorch_sampling(self): if source == "_deterministic_attention_backend" ] self.assertEqual(len(deterministic_fills), 1) - self.assertEqual(sa.attention_backend, deterministic_fills[0]) - self.assertEqual(flags.attention_backend, deterministic_fills[0]) + self.assertEqual( + resolution_result(sa, "attention_backend"), deterministic_fills[0] + ) + self.assertEqual(self._leaf("attention_backend"), deterministic_fills[0]) def test_deterministic_incompatible_backend_raises(self): from sglang.srt.arg_groups.overrides import ( @@ -1148,13 +1184,17 @@ def test_dllm_forces_flashinfer_with_cuda_graph(self): disable_radix_cache=True, attention_backend="triton", ) - self.assertEqual(sa.attention_backend, "flashinfer") # materialized + self.assertEqual( + self._resolved(sa, "attention_backend"), "flashinfer" + ) # materialized self.assertIn( ("_dllm_attention_backend", {"attention_backend": "flashinfer"}), sa._resolved_overrides, ) # the deterministic fill lands on the attention_backend field - self.assertEqual(self._publish(sa).attention_backend, "flashinfer") + self.assertEqual( + (self._publish(sa), self._leaf("attention_backend"))[1], "flashinfer" + ) def test_attention_backend_leaf_materializes_end_state(self): # The default-fill pass declares the platform-selected backend; the @@ -1167,8 +1207,12 @@ def test_attention_backend_leaf_materializes_end_state(self): if "attention_backend" in d ] self.assertTrue(declared_values) # default fill declared - self.assertEqual(sa.attention_backend, declared_values[-1]) # materialized - self.assertEqual(self._publish(sa).attention_backend, declared_values[-1]) + self.assertEqual( + self._resolved(sa, "attention_backend"), declared_values[-1] + ) # materialized + self.assertEqual( + (self._publish(sa), self._leaf("attention_backend"))[1], declared_values[-1] + ) def test_post_materialize_pass_writes_through(self): from sglang.srt.arg_groups.overrides import run_post_process_pass @@ -1177,7 +1221,7 @@ def test_post_materialize_pass_writes_through(self): # legacy runner-side adjustments) declares AND writes through, so # field readers and the publish see the same end state. sa = self._construct("LlamaForCausalLM", "llama") - resolved_before = sa.attention_backend + resolved_before = self._resolved(sa, "attention_backend") def _force_triton(view): if view.attention_backend != "triton": @@ -1186,13 +1230,18 @@ def _force_triton(view): run_post_process_pass(sa, _force_triton) if resolved_before != "triton": - self.assertEqual(sa.attention_backend, "triton") - self.assertEqual(self._publish(sa).attention_backend, sa.attention_backend) + self.assertEqual(self._resolved(sa, "attention_backend"), "triton") + self.assertEqual( + (self._publish(sa), self._leaf("attention_backend"))[1], + self._resolved(sa, "attention_backend"), + ) def test_attention_backend_user_choice_declares_nothing_extra(self): sa = self._construct("LlamaForCausalLM", "llama", attention_backend="triton") - self.assertEqual(sa.attention_backend, "triton") - self.assertEqual(self._publish(sa).attention_backend, "triton") + self.assertEqual(self._resolved(sa, "attention_backend"), "triton") + self.assertEqual( + (self._publish(sa), self._leaf("attention_backend"))[1], "triton" + ) def test_compatibility_passes_at_callable_level(self): from sglang.srt.arg_groups.overrides import ( diff --git a/test/registered/unit/test_runtime_context_config_bags.py b/test/registered/unit/test_runtime_context_config_bags.py index c05a23cb37e2..73a0226de75e 100644 --- a/test/registered/unit/test_runtime_context_config_bags.py +++ b/test/registered/unit/test_runtime_context_config_bags.py @@ -11,6 +11,7 @@ from sglang.srt import runtime_context as rc from sglang.srt.arg_groups.arg_utils import NS, A +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -96,16 +97,15 @@ def test_the_bags_carry_what_resolution_produced(self): none-flags) the first resolution may have written -- so the assertion is "bag == what resolution produces", not "bag == the instance publish copied from". Reproducibility (`test_resolution_is_reproducible`) - licenses the sibling as a stand-in for the pipeline's output. The - raw-differs guard keeps the comparison meaningful: every sampled leaf - must have moved off its dataclass default, so each equality compares a - value resolution demonstrably wrote. Supplied construction inputs - (`model_path`, `device`, `random_seed`) and leaves resolution leaves - alone never enter the sample -- projection coverage for those lives in - `test_passthrough_leaves_project_into_their_namespaces`. Step 12 keeps - records at the user's raw input; then the sibling goes raw and this - assertion starts failing for every sampled leaf, which is the signal - the bags became the only home of the effective value. + licenses the sibling as a stand-in for the pipeline's output. + + The reference's resolved values are read through `resolution_result`, + because a record holds the user's raw input: the decision lives in the + declarations, and the bags are where a process reads it. The + raw-differs guard keeps the comparison meaningful -- every sampled leaf + must have moved off its dataclass default -- and the last assertion is + the other half of that invariant: the record still answers the raw + input for a leaf resolution decided. """ import dataclasses @@ -124,12 +124,13 @@ def test_the_bags_carry_what_resolution_produced(self): # The raw-differs guard: a sampled leaf that still sits on its # default (or has none to differ from) proves nothing. self.assertIsNot(defaults[leaf], dataclasses.MISSING) - self.assertNotEqual(getattr(reference, leaf), defaults[leaf]) - self.assertEqual(accessor(), getattr(reference, leaf)) - # And the record agrees today, which is what step 12 changes: when this - # assertion starts failing for a resolution-written leaf, the flip - # landed and the bag is the only place the effective value lives. - self.assertEqual(rc.get_schedule().page_size, sa.page_size) + resolved = resolution_result(reference, leaf) + self.assertNotEqual(resolved, defaults[leaf]) + self.assertEqual(accessor(), resolved) + # The record is the raw input, so the field still reads as the default + # for a leaf the bag now answers for. + self.assertEqual(sa.page_size, defaults["page_size"]) + self.assertNotEqual(rc.get_schedule().page_size, sa.page_size) def test_passthrough_leaves_project_into_their_namespaces(self): """Thin projection smoke over leaves resolution does not move. @@ -141,7 +142,7 @@ def test_passthrough_leaves_project_into_their_namespaces(self): sa = self._publish() sampled = ( (lambda: rc.get_serving().host, "host"), - (lambda: rc.get_memory().hicache_ratio, "hicache_ratio"), + (lambda: rc.get_memory().hicache_write_policy, "hicache_write_policy"), (lambda: rc.get_exec().moe.moe_runner_backend, "moe_runner_backend"), (lambda: rc.get_model().model_path, "model_path"), ) @@ -230,8 +231,8 @@ def test_read_only_by_bare_assignment(self): rc.get_memory().hicache_ratio = 9.0 def test_scoped_override_restores(self): - sa = self._publish() - original = sa.hicache_ratio + self._publish() + original = rc.get_memory().hicache_ratio with rc.get_memory().override(hicache_ratio=original + 1.0): self.assertEqual(rc.get_memory().hicache_ratio, original + 1.0) self.assertEqual(rc.get_memory().hicache_ratio, original) diff --git a/test/registered/unit/test_runtime_context_override.py b/test/registered/unit/test_runtime_context_override.py index 623125a00ef2..fecea1ddb879 100644 --- a/test/registered/unit/test_runtime_context_override.py +++ b/test/registered/unit/test_runtime_context_override.py @@ -110,7 +110,7 @@ def test_bare_server_args_write_raises_after_resolution(self): # server_args is read-only after resolution: resolved config changes go # to the bags, a per-runner config to a derived variant. sa = ServerArgs(model_path="dummy") - object.__setattr__(sa, "_declarations_materialized", True) + object.__setattr__(sa, "_resolution_finished", True) with self.assertRaises(AttributeError): sa.page_size = 999 diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index 95107549f366..e8f91a5ada1d 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -535,9 +535,9 @@ def _declarative_override_fields(self) -> set: ``MODEL_OVERRIDES`` maps arch -> {field: value}, and the ``@register_model_override``(-``_predicate``) providers return (or - build by subscript) {field: value} dicts; ``materialize_declarations`` - applies them all via setattr, so no assignment scan sees these writes - and a llama-only matrix never triggers them. Keys must be + build by subscript) {field: value} dicts, which go straight into the + declaration stash, so no assignment scan sees these writes and a + llama-only matrix never triggers them. Keys must be string literals; anything else fails loudly. """ tree = ast.parse( From 8bfb9d9cec44af599d1ed94bd230401a03e0c5db Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 13:06:49 +0000 Subject: [PATCH 24/26] config: ServerArgs holds the raw input; the declarations are the resolution `declare_resolution` no longer writes the field it declares. The record is what the caller passed, the stash is what resolution decided, and the config bags are projected from the stash -- which is what this whole series was for. Two things fall out. The pipeline becomes reproducible over a copy: resolving a bare `dataclasses.replace` now runs over the same input the parent got instead of over the parent's output, so the DP-attention halving and the conservativeness scaling apply once rather than twice. Only `random_seed` differs, because it is generated. `replace_resolved` keeps its reason -- it carries the parent's declarations and its `model_config`, so the copy answers without re-deriving anything -- and the test that used to assert the drift now asserts its absence. And a field read inside resolution becomes a real bug rather than a latent one, so `test_resolution_reads_the_declarations` pins it at zero over the two scopes it can derive exactly: every `arg_groups` function that takes a config, and every `ServerArgs` handler the dispatcher reaches. Handing a record to a helper that loads a decided field is the shape neither an attribute scan nor a `getattr` scan can see -- the field is spelled in the helper and the record at the call site -- so the guard derives those helpers and pins the four spellings that reach one: a bare name, an attribute, `get_server_args()` called inline, and a local bound to either. It reads a decided leaf straight off such a local too, and off a private attribute (`self._server_args.device`) or a record a script builds for itself (`ServerArgs.from_cli_args`), which is why the scan now covers `scripts/`, `examples/` and the gateway binding alongside the package. One reader is pinned with its reason: the gateway falls back to the field on a released wheel that has no `resolved_dict`. It also reads a decided leaf straight off such a local (`alias.`, `getattr(alias, "")`) -- the spelling where the leaf never appears at a call site. The spellings are pinned by a fixture rather than by a comment: the scan runs over a sample module holding each of them plus the legal forms next to them, so a shape cannot quietly stop being covered. Verified by comparing the whole resolution result -- all 476 fields, 20 launch shapes -- against the previous commit: zero differences. Four readers outside `srt/` follow, each of which was reading a field resolution fills in: the named-stream factory (`device` -- `torch.get_device_module(None)` lands on CUDA whatever the host is), the gateway's worker count (`dp_size`, which `--dwdp-size` fills, so a DWDP launch started one worker), the speculative benchmark (`mem_fraction_static` reached the child command as the string "None"), and the two checkpoint exporters (`model_path`, which ModelScope resolution rewrites to the downloaded directory). --- .../skills/sglang-runtime-context/SKILL.md | 7 +- examples/runtime/engine/save_remote_state.py | 3 +- examples/runtime/engine/save_sharded_state.py | 3 +- .../token_in_token_out_vlm_engine.py | 8 +- python/sglang/srt/arg_groups/overrides.py | 72 +-- python/sglang/srt/entrypoints/engine.py | 4 +- python/sglang/srt/entrypoints/http_server.py | 4 +- python/sglang/srt/layers/moe/utils.py | 4 +- .../sglang/srt/parser/template_detection.py | 11 +- python/sglang/srt/runtime_context.py | 35 +- python/sglang/srt/server_args.py | 29 +- scripts/playground/bench_speculative.py | 4 + .../python/src/sglang_router/launch_server.py | 6 +- .../cpu/test_server_args_backend.py | 13 +- .../unit/parser/test_template_manager.py | 64 +- .../test_model_config_reads_resolved_input.py | 4 +- .../test_resolution_declarations.py | 83 ++- .../test_resolution_is_reproducible.py | 34 +- .../test_resolution_reads_the_declarations.py | 612 ++++++++++++++++++ ..._server_args_no_instance_mutation_entry.py | 6 +- 20 files changed, 836 insertions(+), 170 deletions(-) create mode 100644 test/registered/unit/server_args/test_resolution_reads_the_declarations.py diff --git a/.claude/skills/sglang-runtime-context/SKILL.md b/.claude/skills/sglang-runtime-context/SKILL.md index 77e93c66f771..bfdae19ce0ca 100644 --- a/.claude/skills/sglang-runtime-context/SKILL.md +++ b/.claude/skills/sglang-runtime-context/SKILL.md @@ -99,8 +99,8 @@ resolved configuration lives in the namespace bags.** **Why a bag override cannot stand in for late resolution or per-runner construction.** The bags are projected at -publish *from the instance's fields*, so anything the runtime must read has to be on -the instance before publish — an override afterwards puts instance and bags back out +publish *from the declarations over the instance's raw fields*, so anything the +runtime must read has to be declared before publish — an override afterwards puts instance and bags back out of agreement, and whole-object readers (`ModelConfig.from_server_args`, `build_load_config`, `MMEncoder`'s own `self.server_args.X`) never see it. And bags do not cross a process boundary: a child publishes from the object it receives and @@ -114,7 +114,8 @@ bag to override at all. - **Per-runner values** — there is no per-runner `ServerArgs` any more. The draft-worker config copy is gone: every worker (`TpModelWorker`, the draft workers in `speculative/`) is handed the *same* instance the process published, - so `self.server_args.X` and the bag leaf agree **at publish** — a + so a bag leaf is the decision and `self.server_args.X` is the operator's + input — a post-publish `override` moves only the bag, which is exactly why a field that is process-wide config (`attention_backend`, `skip_tokenizer_init`, `kv_cache_dtype`) reads from the bags like any other, and why a residual diff --git a/examples/runtime/engine/save_remote_state.py b/examples/runtime/engine/save_remote_state.py index 84f43b604363..d96810334047 100644 --- a/examples/runtime/engine/save_remote_state.py +++ b/examples/runtime/engine/save_remote_state.py @@ -23,6 +23,7 @@ from pathlib import Path from sglang import Engine, ServerArgs +from sglang.srt.arg_groups.overrides import resolution_result parser = ArgumentParser() ServerArgs.add_cli_args(parser) @@ -44,7 +45,7 @@ def main(args): engine_args = ServerArgs.from_cli_args(args) engine_args.resolve_once() - model_path = engine_args.model_path + model_path = resolution_result(engine_args, "model_path") if not Path(model_path).is_dir(): raise ValueError("model path must be a local directory") # Create LLM instance from arguments diff --git a/examples/runtime/engine/save_sharded_state.py b/examples/runtime/engine/save_sharded_state.py index a27d5ba7062c..4994e6615350 100644 --- a/examples/runtime/engine/save_sharded_state.py +++ b/examples/runtime/engine/save_sharded_state.py @@ -28,6 +28,7 @@ from pathlib import Path from sglang import Engine, ServerArgs +from sglang.srt.arg_groups.overrides import resolution_result parser = ArgumentParser() ServerArgs.add_cli_args(parser) @@ -49,7 +50,7 @@ def main(args): engine_args = ServerArgs.from_cli_args(args) engine_args.resolve_once() - model_path = engine_args.model_path + model_path = resolution_result(engine_args, "model_path") if not Path(model_path).is_dir(): raise ValueError("model path must be a local directory") # Create LLM instance from arguments diff --git a/examples/runtime/token_in_token_out/token_in_token_out_vlm_engine.py b/examples/runtime/token_in_token_out/token_in_token_out_vlm_engine.py index 0853c8bdb255..b5db0dfe4f29 100644 --- a/examples/runtime/token_in_token_out/token_in_token_out_vlm_engine.py +++ b/examples/runtime/token_in_token_out/token_in_token_out_vlm_engine.py @@ -5,6 +5,7 @@ from sglang import Engine from sglang.lang.chat_template import get_chat_template_by_model_path +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.configs.model_config import ModelConfig from sglang.srt.server_args import ServerArgs from sglang.test.test_utils import DEFAULT_IMAGE_URL @@ -36,12 +37,13 @@ def get_input_ids( def token_in_out_example( server_args: ServerArgs, ): + cfg = resolving_view(server_args) input_ids, image_data = get_input_ids( server_args, ModelConfig( - server_args.model_path, - trust_remote_code=server_args.trust_remote_code, - model_override_args=server_args.json_model_override_args, + cfg.model_path, + trust_remote_code=cfg.trust_remote_code, + model_override_args=cfg.json_model_override_args, ), ) backend = Engine(server_args=server_args) diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index dc7e826db501..645c7854e762 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -14,9 +14,13 @@ """Declarative model-override registry. Model-identity adjustments to the server configuration are DECLARED here and -materialized onto ``server_args`` at the end of ``__post_init__`` (gate -order, last writer wins) — model code never mutates ``ServerArgs`` fields -imperatively. +appended to the record's declaration stash (gate order, last writer wins). +Nothing here writes back onto ``ServerArgs``: the record holds the user's raw +input, and a decision is read through ``resolution_result`` or the published +config bags — model code never mutates ``ServerArgs`` fields imperatively. The +one channel that still leaves a field changed is ``declare_direct_writes``, +which does not perform the write: it captures one an out-of-tree plugin already +made, and undoing it would surprise the plugin's own reads. Two declaration forms, keyed on ``hf_config.architectures[0]``: @@ -206,9 +210,8 @@ def register_post_process(fn: Callable[..., dict]) -> Callable[..., dict]: def _declaration_overlay(server_args: Any) -> Dict[str, Any]: """What the declarations say so far, last writer wins. - Passes declare without touching the fields, so a mid-resolution reader - needs this to see them; handlers and hooks write as they declare, and for - those the overlay repeats what the field already holds.""" + Nothing writes the fields, so a mid-resolution reader needs this to see a + decision at all; the fields keep what the caller supplied.""" overlay: Dict[str, Any] = {} for _source, declared in getattr(server_args, "_resolved_overrides", None) or (): overlay.update(declared) @@ -259,29 +262,17 @@ def _apply_fields(server_args: Any, fields: Dict[str, Any]) -> None: def declare_resolution(server_args: Any, source: str, **fields: Any) -> None: - """Record a resolution write in the declaration stash, and apply it now. - - The stash is what the projection reads, so a resolver that only assigns - the field leaves that write invisible to it. The immediate write keeps the - resolver's successors seeing the value where they read the field directly. - - What it does change is which writer wins. A declaration is appended and - replayed last, so a resolver that declares a field a *deferred* writer (a - post-process pass, a registry entry) also decides now beats it, where its - bare assignment used to be overwritten by that writer's declaration. A - resolver that gates on such a field has to read the resolving view rather - than the raw field, or it decides from a value that is already stale. - - For resolvers inside ``__post_init__``: the handlers on ``ServerArgs`` - (through ``self._declare``) and the ``arg_groups`` hooks and hardware - defaults they call. Resolution that has to wait for the launcher stage - goes through ``declare_late_resolution`` instead. - - Names arrive as keyword arguments, which accept anything; a misspelled one - would otherwise become a new attribute that nothing ever reads, so it is - rejected here. This is not the model-override whitelist: that one limits - which fields a *registry entry* may reach, while a resolver writing the - field it owns is the pipeline resolving by construction. + """Record a resolution write in the declaration stash. + + The stash *is* the resolution result: the bags are projected from it, + `resolution_result` answers from it, and no field is written. A resolver + reading a field another resolver may have decided must read `resolving_view` + (or `ServerArgs._resolved()`), which + `test_resolution_reads_the_declarations` pins. + + For resolvers inside ``__post_init__``; launcher-stage resolution goes + through ``declare_late_resolution``. A name that is not a field is rejected + here rather than becoming an attribute nothing reads. """ if dataclasses.is_dataclass(type(server_args)): unknown = sorted(set(fields) - field_names(type(server_args))) @@ -292,8 +283,6 @@ def declare_resolution(server_args: Any, source: str, **fields: Any) -> None: stash = [] object.__setattr__(server_args, "_resolved_overrides", stash) stash.append((source, dict(fields))) - for name, value in fields.items(): - setattr(server_args, name, value) def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> None: @@ -303,10 +292,10 @@ def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> Non normalization and the auto-parser detection need the launcher's validation stage (and, for the parsers, a tokenizer / chat-template load). They still belong to the resolution pipeline — they decide what the process will run - with — so they write the fields in place, before anything publishes the - object. Writing in place is the point: every holder of that instance (the - HTTP server, the multi-tokenizer workers it serializes for, the schedulers - it forks) must see the resolved value. + with — so their decision goes to the stash like any other, and the record + keeps what the caller passed. Every holder of that instance reads the + decision the same way the rest of the pipeline does: the bags it publishes, + or ``resolution_result``, both of which survive the pickle to a child. Refuses to touch the published instance: after publish the bags exist and a field write would desync them, which is what ``get_context().override`` is @@ -333,7 +322,6 @@ def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> Non stash = [] object.__setattr__(server_args, "_resolved_overrides", stash) stash.append((source, dict(fields))) - _apply_fields(server_args, fields) def declare_direct_writes( @@ -347,7 +335,10 @@ def declare_direct_writes( Out-of-tree platform plugins are handed the record and set fields on it. Their implementations live outside this tree, so they cannot be converted by editing the resolver; and the raw snapshot is taken before the pipeline - starts, so a plugin's default is neither declared nor raw. + starts, so a plugin's default is neither declared nor raw. The write itself + stays: this captures it into the stash so the projection and the bags carry + it, but reverting the field would break the plugin's own reads of what it + just set. It is the only field a record still carries from resolution. Rebinding is what the diff sees, and rebinding is all it needs to see: a plugin that mutates a value in place reaches the projection anyway, because @@ -388,7 +379,7 @@ def resolution_result(server_args: Any, field: str, default: Any = None) -> Any: otherwise what the caller supplied. This is what the config projection reads. Reading the field instead would - work only for as long as declarations materialize onto the record -- and + work whatever the caller passed onto the record -- and the point of declaring is that they will not, so the projection must not depend on it. A config that never ran the pipeline (a mock, a partial fixture) carries no raw snapshot; its fields are all it has. @@ -409,9 +400,8 @@ def resolution_projection(server_args: Any) -> Dict[str, Any]: The whole-object shape of ``resolution_result``, for the exits that hand out the entire configuration (``/server_info``, the gRPC and engine readbacks). - They used ``dataclasses.asdict``, which reads the fields -- correct only for - as long as declarations materialize onto the record, and the point of - declaring is that they will not. Field values only: the private resolution + They used ``dataclasses.asdict``, which reads the fields -- the operator's + input, not what resolution decided. Field values only: the private resolution bookkeeping and the ``model_config`` memo that a ``vars()`` dump carried into the readback are not configuration. """ diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 670831be1d17..c11ee781ba3b 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -251,7 +251,7 @@ def __init__(self, **kwargs): kwargs["log_level"] = "error" server_args = self.server_args_class(**kwargs) self.server_args = server_args - logger.info(f"{server_args=}") + logger.info(f"server_args={server_args.resolved_dict()}") # Rust Server is not supported with the offline Engine API if envs.SGLANG_RUST_SERVER.get(): @@ -1076,7 +1076,7 @@ def _launch_subprocesses( # Allocate ports for inter-process communications if port_args is None: port_args = PortArgs.init_new(server_args) - logger.info(f"{server_args=}") + logger.info(f"server_args={server_args.resolved_dict()}") # Start the engine info bootstrap server if per-rank info is needed. engine_info_bootstrap_server = None diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index d56aa71e39b5..a3f0b068e2db 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -803,8 +803,8 @@ async def get_server_info(): async def server_info(): """The startup configuration, plus live scheduler state. - The `ServerArgs` fields here are the record: what the launcher was given, - with resolution written back into it. Fields the control plane changes + The values here are the resolution result: what the launcher was given, + with every decision resolution made applied over it. Fields the control plane changes after publication -- the model a weight update swapped in, its load format, an operator-set weight version -- are reported by `/model_info`, and the HiCache mirror by `GET /hicache/storage-backend`. diff --git a/python/sglang/srt/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py index 11f8a4b156f1..2a844cd6790b 100644 --- a/python/sglang/srt/layers/moe/utils.py +++ b/python/sglang/srt/layers/moe/utils.py @@ -315,8 +315,8 @@ def initialize_moe_config(): """Seed the MoE runtime flags from the published configuration. Reads the bags: `moe_a2a_backend` and its siblings are resolution's - answers, and the record carries them only while declarations materialize - onto it. Called once per process after publish + answers, and the record carries the operator's input. Called once per + process after publish (scheduler init, the benchmark work functions). """ exec_moe = get_exec().moe diff --git a/python/sglang/srt/parser/template_detection.py b/python/sglang/srt/parser/template_detection.py index 13d5465a7eb2..766243bfd07a 100644 --- a/python/sglang/srt/parser/template_detection.py +++ b/python/sglang/srt/parser/template_detection.py @@ -702,12 +702,13 @@ def _architecture_auto_parsers(server_args, needs: Tuple[str, ...]) -> Dict[str, def resolve_auto_parsers(server_args) -> None: """Resolve ``--reasoning-parser=auto`` / ``--tool-call-parser=auto`` from the - chat template, in place, before anything publishes ``server_args``. + chat template, before anything publishes ``server_args``. - Performs a lightweight tokenizer load, so it runs once in engine init. In - place because everyone who holds this instance must see the resolved value: - the schedulers it forks, the HTTP server, and the tokenizer workers it is - serialized for. + Performs a lightweight tokenizer load, so it runs once in engine init. The + decision goes to this instance's declaration stash, so every holder of it + carries it -- the schedulers it forks, the HTTP server, the tokenizer + workers it is serialized for -- and each publishes bags projected from it. + The fields stay what the operator passed. """ cfg = resolving_view(server_args) needs = tuple( diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index c94b32c6ff5a..554ba1be1325 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -24,9 +24,10 @@ hop (``get_parallel().config.tp_size``), which reads the published ``parallel`` bag: bare is the live group, ``config`` is what was configured. -``get_server_args()`` returns the process-wide ``ServerArgs``. This is the pristine / resolved-at-startup **read-only** record kept -for debug and reproduction; business code reads resolved config from the -namespace bags below, not from this object. The context owns the storage: +``get_server_args()`` returns the process-wide ``ServerArgs``. This is the +user's raw input, kept **read-only** for debug and reproduction; what +resolution decided lives in the declarations (``resolution_result``) and, for +business code, in the namespace bags below -- never on this object's fields. The context owns the storage: publishing goes through ``RuntimeContext.set_server_args`` (the legacy ``set_global_server_args_for_scheduler`` / ``get_global_server_args`` are thin shims over this slot). @@ -440,8 +441,8 @@ class DpFlags(_FlagGroupBase): class Flags(_FlagGroupBase): """Root of the runtime-flags tier. - Resolved configuration lives on ``server_args`` fields (materialized at - the end of ``__post_init__``) — this tier only carries genuine runtime + Resolved configuration lives in the config bags below (projected from the + declarations at publish) — this tier only carries genuine runtime state whose value is not a function of the configuration alone, grouped by lifecycle (``capture``) or subsystem (``moe`` / ``dp``). """ @@ -700,8 +701,7 @@ def _build_config_bags(server_args: Any) -> dict: """Snapshot the resolution result into the namespace bag tree, driven by the ``NS(...)`` metadata on the dataclass fields. Each leaf comes from ``resolution_result`` -- the declaration if resolution made one, else what - the caller supplied -- rather than from the field, which carries the same - value only while declarations still materialize. Returns + the caller supplied. Returns ``{top_level_name: _ConfigBag}``, arbitrarily nested (``exec.moe.eplb.…``). Only dataclass fields carry ``NS`` markers, so derived properties/methods are naturally excluded (they stay on the bag). A name used as both a leaf and a @@ -783,7 +783,13 @@ def get_stream(self, name: str) -> Any: if stream is None: import torch - device = self._server_args.device if self._server_args else "cuda" + from sglang.srt.arg_groups.overrides import resolution_result + + device = ( + resolution_result(self._server_args, "device") + if self._server_args + else "cuda" + ) stream = torch.get_device_module(device).Stream() self.resources.streams[name] = stream return stream @@ -819,8 +825,8 @@ def set_server_args(self, server_args: ServerArgs) -> None: Overwrite-allowed: a re-publish replaces the slot (test kits re-publish per test; production ordering discipline lives at the call-sites, e.g. the draft-worker guard in ``ModelRunner.__init__``). The published - object already carries the resolved configuration (declarations - materialize at the end of ``__post_init__``). + object is the raw input; the resolution it carries is its declaration + stash, which is what the bags are projected from. """ # Seed the capture tier for the new lifecycle (defaults for sentinel # and mock publishes, which carry no config). @@ -1078,10 +1084,11 @@ def install(self) -> ServerArgs: } if declared: declare_late_resolution(server_args, "override_server_args", **declared) - _apply_fields( - server_args, - {name: value for name, value in self._fields.items() if name[0] == "_"}, - ) + # This hook stands in for a launch: the caller's values are both what + # the operator passed and what resolution decided, so they go on the + # record as well as into the stash. Production late resolution declares + # only -- there the record stays the operator's input. + _apply_fields(server_args, self._fields) ctx.set_server_args(server_args) self._installed = True return server_args diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 295b52c567b7..93141bb782b9 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -3721,13 +3721,13 @@ def resolve_once(self) -> None: try: self._run_resolution_pipeline() except BaseException: - # The handlers that ran already wrote to the record, and they are - # not idempotent over their own output. + # The handlers that ran already declared, and they are not + # idempotent over their own output. object.__setattr__(self, "_resolution_failed", True) raise # Set here too, because the dummy/absent-model path returns before the - # materialization that normally sets it: the gate is about whether the - # handlers ran, not how far they got. + # end of the pipeline that normally sets it: the gate is about whether + # the handlers ran, not how far they got. self._resolution_finished = True def resolved_dict(self) -> Dict[str, Any]: @@ -3735,9 +3735,8 @@ def resolved_dict(self) -> Dict[str, Any]: What the whole-object readbacks report (`/server_info` and its gRPC and in-process twins). `dataclasses.asdict(self)` reads the fields, which - carry resolution's result only while declarations materialize onto the - record; this reads the declarations, so it keeps answering with what - resolution decided once they stop. Nested dataclass fields are expanded + carry the raw input; this reads the declarations, so it answers with what + resolution decided. Nested dataclass fields are expanded the way `asdict` expands them; the private resolution bookkeeping and the `model_config` memo are not fields and do not appear. """ @@ -3750,12 +3749,11 @@ def replace_resolved(self, source: str, **changes: Any) -> ServerArgs: `dataclasses.replace` builds a new instance, so the copy carries none of what makes a record resolved: no raw snapshot, no declarations, no - materialization. The next publish therefore finds an unmaterialized - record and runs the pipeline over values it already decided -- DP - attention halves `chunked_prefill_size` a second time (8192 -> 4096 -> - 2048) and the schedule conservativeness is scaled again (0.3 -> 0.09). - The Ray paths replace `dist_init_addr` on a resolved record, which is - how they hit it. + finished flag. The next publish therefore resolves it again, which + drops every decision the stash held -- the late ones (the auto-detected + parsers) and the direct ones alike -- and re-runs the device probes in + whatever process opened the copy. The Ray paths replace + `dist_init_addr` on a resolved record, which is how they reach this. The change is appended to the stash rather than left on the field: the projection reads the raw snapshot plus the declarations, so a field the @@ -9797,9 +9795,8 @@ def _late_resolution(self, source: str, **fields) -> None: declare_late_resolution(self, source, **fields) def __setattr__(self, name, value): - # After materialization the fields are the resolved startup - # configuration -- the pristine, READ-ONLY record that the config bags - # were projected from. Resolved config changes go to the bags via + # Once resolution has finished the record is the READ-ONLY raw input + # the config bags were projected from. Resolved config changes go to the bags via # get_context().override(source, ...); a value one runner or worker # owns travels as a constructor argument to it. if ( diff --git a/scripts/playground/bench_speculative.py b/scripts/playground/bench_speculative.py index 54830a9dbc2c..6638a179618f 100644 --- a/scripts/playground/bench_speculative.py +++ b/scripts/playground/bench_speculative.py @@ -22,6 +22,7 @@ from sglang.bench_serving import benchmark, set_global_args from sglang.benchmark.datasets import DatasetRow from sglang.benchmark.datasets.mmmu import sample_mmmu_requests +from sglang.srt.arg_groups.overrides import resolution_projection from sglang.srt.server_args import ServerArgs from sglang.test.test_utils import ( DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, @@ -135,6 +136,9 @@ def send_one_batch(base_url, num_prompts, batch_size, processor, is_multimodal): def main(args, server_args): + from sglang.srt.arg_groups.overrides import resolution_projection + + server_args = SimpleNamespace(**resolution_projection(server_args)) base_url = "http://127.0.0.1:20000" configs = [] diff --git a/sgl-model-gateway/bindings/python/src/sglang_router/launch_server.py b/sgl-model-gateway/bindings/python/src/sglang_router/launch_server.py index 6743f949a009..ae7b680d3bb1 100644 --- a/sgl-model-gateway/bindings/python/src/sglang_router/launch_server.py +++ b/sgl-model-gateway/bindings/python/src/sglang_router/launch_server.py @@ -191,7 +191,11 @@ def main(): server_args.resolve_once() router_args = RouterArgs.from_cli_args(args, use_router_prefix=True) - # Find available ports for workers + # Find available ports for workers. The count is the operator's requested + # replica count, which is the raw field on purpose: `--dwdp-size` makes + # resolution declare a `dp_size` that describes one multi-rank server's + # internal topology, and spawning that many single-rank children would ask + # for dp_size^2 GPUs. worker_ports = find_available_ports( args.router_dp_worker_base_port, server_args.dp_size ) diff --git a/test/registered/cpu/test_server_args_backend.py b/test/registered/cpu/test_server_args_backend.py index 0f6cb7dc9f7f..423b11da36d8 100644 --- a/test/registered/cpu/test_server_args_backend.py +++ b/test/registered/cpu/test_server_args_backend.py @@ -4,6 +4,7 @@ import unittest from unittest.mock import patch +from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci @@ -25,8 +26,10 @@ def test_arm_cpu_defaults_to_torch_native(self, _mock_is_arm64): ServerArgs._handle_cpu_backends(server_args) - self.assertEqual(server_args.attention_backend, "torch_native") - self.assertEqual(server_args.sampling_backend, "pytorch") + self.assertEqual( + resolution_result(server_args, "attention_backend"), "torch_native" + ) + self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch") @patch("sglang.srt.server_args.is_host_cpu_arm64", return_value=False) def test_x86_cpu_defaults_to_intel_amx(self, _mock_is_arm64): @@ -34,8 +37,10 @@ def test_x86_cpu_defaults_to_intel_amx(self, _mock_is_arm64): ServerArgs._handle_cpu_backends(server_args) - self.assertEqual(server_args.attention_backend, "intel_amx") - self.assertEqual(server_args.sampling_backend, "pytorch") + self.assertEqual( + resolution_result(server_args, "attention_backend"), "intel_amx" + ) + self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch") class TestServerArgsIBDeviceValidation(unittest.TestCase): diff --git a/test/registered/unit/parser/test_template_manager.py b/test/registered/unit/parser/test_template_manager.py index af5e0ea0c02a..497077137d90 100644 --- a/test/registered/unit/parser/test_template_manager.py +++ b/test/registered/unit/parser/test_template_manager.py @@ -765,6 +765,18 @@ def test_minicpm5_not_misclassified_as_qwen(self): self.assertEqual(result, "minicpm5") +def _declared(server_args, field): + """What late resolution decided for `field` on this record. + + `resolve_auto_parsers` declares; the field keeps what the operator passed, + so the decision is read through the resolution result -- the same surface + the config bags are projected from. + """ + from sglang.srt.arg_groups.overrides import resolution_result + + return resolution_result(server_args, field) + + class TestResolveAutoParsers(unittest.TestCase): """Tests for resolve_auto_parsers().""" @@ -790,8 +802,8 @@ def test_resolves_both_parsers_with_tokenizer_template(self): with _patch_hf_transformers_utils(Mock(return_value=tokenizer)): resolve_auto_parsers(args) - self.assertEqual(args.reasoning_parser, "qwen3") - self.assertEqual(args.tool_call_parser, "qwen") + self.assertEqual(_declared(args, "reasoning_parser"), "qwen3") + self.assertEqual(_declared(args, "tool_call_parser"), "qwen") def test_resolves_reasoning_parser_only(self): args = self._make_server_args(reasoning_parser="auto", tool_call_parser=None) @@ -800,8 +812,8 @@ def test_resolves_reasoning_parser_only(self): with _patch_hf_transformers_utils(Mock(return_value=tokenizer)): resolve_auto_parsers(args) - self.assertEqual(args.reasoning_parser, "qwen3") - self.assertIsNone(args.tool_call_parser) + self.assertEqual(_declared(args, "reasoning_parser"), "qwen3") + self.assertIsNone(_declared(args, "tool_call_parser")) def test_resolves_tool_call_parser_only(self): args = self._make_server_args(reasoning_parser="qwen3", tool_call_parser="auto") @@ -810,14 +822,14 @@ def test_resolves_tool_call_parser_only(self): with _patch_hf_transformers_utils(Mock(return_value=tokenizer)): resolve_auto_parsers(args) - self.assertEqual(args.reasoning_parser, "qwen3") - self.assertEqual(args.tool_call_parser, "qwen") + self.assertEqual(_declared(args, "reasoning_parser"), "qwen3") + self.assertEqual(_declared(args, "tool_call_parser"), "qwen") def test_neither_auto_is_noop(self): args = self._make_server_args(reasoning_parser="qwen3", tool_call_parser="qwen") resolve_auto_parsers(args) - self.assertEqual(args.reasoning_parser, "qwen3") - self.assertEqual(args.tool_call_parser, "qwen") + self.assertEqual(_declared(args, "reasoning_parser"), "qwen3") + self.assertEqual(_declared(args, "tool_call_parser"), "qwen") def test_nonexistent_model_disables_both_parsers(self): args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") @@ -832,8 +844,8 @@ def test_nonexistent_model_disables_both_parsers(self): ): resolve_auto_parsers(args) - self.assertIsNone(args.reasoning_parser) - self.assertIsNone(args.tool_call_parser) + self.assertIsNone(_declared(args, "reasoning_parser")) + self.assertIsNone(_declared(args, "tool_call_parser")) def test_none_chat_template_disables_both_parsers(self): args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") @@ -842,8 +854,8 @@ def test_none_chat_template_disables_both_parsers(self): with _patch_hf_transformers_utils(Mock(return_value=tokenizer)): resolve_auto_parsers(args) - self.assertIsNone(args.reasoning_parser) - self.assertIsNone(args.tool_call_parser) + self.assertIsNone(_declared(args, "reasoning_parser")) + self.assertIsNone(_declared(args, "tool_call_parser")) def test_deepseek_v32_arch_without_chat_template_uses_custom_encoder(self): args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") @@ -855,8 +867,8 @@ def test_deepseek_v32_arch_without_chat_template_uses_custom_encoder(self): ): resolve_auto_parsers(args) - self.assertEqual(args.reasoning_parser, "deepseek-v3") - self.assertEqual(args.tool_call_parser, "deepseekv32") + self.assertEqual(_declared(args, "reasoning_parser"), "deepseek-v3") + self.assertEqual(_declared(args, "tool_call_parser"), "deepseekv32") def test_deepseek_v4_arch_without_chat_template_uses_custom_encoder(self): args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") @@ -868,8 +880,8 @@ def test_deepseek_v4_arch_without_chat_template_uses_custom_encoder(self): ): resolve_auto_parsers(args) - self.assertEqual(args.reasoning_parser, "deepseek-v4") - self.assertEqual(args.tool_call_parser, "deepseekv4") + self.assertEqual(_declared(args, "reasoning_parser"), "deepseek-v4") + self.assertEqual(_declared(args, "tool_call_parser"), "deepseekv4") def test_kimi_k3_arch_without_chat_template_uses_custom_encoder(self): args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") @@ -883,8 +895,8 @@ def test_kimi_k3_arch_without_chat_template_uses_custom_encoder(self): ): resolve_auto_parsers(args) - self.assertEqual(args.reasoning_parser, "kimi_k3") - self.assertEqual(args.tool_call_parser, "kimi_k3") + self.assertEqual(_declared(args, "reasoning_parser"), "kimi_k3") + self.assertEqual(_declared(args, "tool_call_parser"), "kimi_k3") def test_kimi_k3_model_type_without_architecture_uses_custom_encoder(self): args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") @@ -896,8 +908,8 @@ def test_kimi_k3_model_type_without_architecture_uses_custom_encoder(self): ): resolve_auto_parsers(args) - self.assertEqual(args.reasoning_parser, "kimi_k3") - self.assertEqual(args.tool_call_parser, "kimi_k3") + self.assertEqual(_declared(args, "reasoning_parser"), "kimi_k3") + self.assertEqual(_declared(args, "tool_call_parser"), "kimi_k3") def test_deepseek_arch_fallback_runs_when_tokenizer_load_fails(self): args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") @@ -909,8 +921,8 @@ def test_deepseek_arch_fallback_runs_when_tokenizer_load_fails(self): ): resolve_auto_parsers(args) - self.assertEqual(args.reasoning_parser, "deepseek-v3") - self.assertEqual(args.tool_call_parser, "deepseekv32") + self.assertEqual(_declared(args, "reasoning_parser"), "deepseek-v3") + self.assertEqual(_declared(args, "tool_call_parser"), "deepseekv32") def test_explicit_non_jinja_template_skips_architecture_fallback(self): args = self._make_server_args( @@ -926,8 +938,8 @@ def test_explicit_non_jinja_template_skips_architecture_fallback(self): resolve_auto_parsers(args) get_config.assert_not_called() - self.assertIsNone(args.reasoning_parser) - self.assertIsNone(args.tool_call_parser) + self.assertIsNone(_declared(args, "reasoning_parser")) + self.assertIsNone(_declared(args, "tool_call_parser")) def test_explicit_jinja_template_takes_precedence(self): tokenizer = _DummyTokenizer([], chat_template=None) @@ -947,8 +959,8 @@ def test_explicit_jinja_template_takes_precedence(self): with _patch_hf_transformers_utils(Mock(return_value=tokenizer)): resolve_auto_parsers(args) - self.assertEqual(args.reasoning_parser, "deepseek-v3") - self.assertEqual(args.tool_call_parser, "deepseekv32") + self.assertEqual(_declared(args, "reasoning_parser"), "deepseek-v3") + self.assertEqual(_declared(args, "tool_call_parser"), "deepseekv32") if __name__ == "__main__": diff --git a/test/registered/unit/server_args/test_model_config_reads_resolved_input.py b/test/registered/unit/server_args/test_model_config_reads_resolved_input.py index 8dabf06aa34c..be6b13f4bf21 100644 --- a/test/registered/unit/server_args/test_model_config_reads_resolved_input.py +++ b/test/registered/unit/server_args/test_model_config_reads_resolved_input.py @@ -26,8 +26,8 @@ _SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt" # Read for what the caller asked for: the constructor passes it through and -# never stores it, while resolution later overwrites the field with the value -# the architecture implies. Two quantities sharing one name. +# never stores it, while resolution declares the value the architecture implies. +# Two quantities sharing one name. _READ_BEFORE_RESOLUTION = frozenset({"is_embedding"}) # Declared after the first `get_model_config()`, so the cached configuration diff --git a/test/registered/unit/server_args/test_resolution_declarations.py b/test/registered/unit/server_args/test_resolution_declarations.py index 5ce2c57e9f3d..583f26701883 100644 --- a/test/registered/unit/server_args/test_resolution_declarations.py +++ b/test/registered/unit/server_args/test_resolution_declarations.py @@ -5,13 +5,11 @@ resolver declares now -- the record's handlers through `self._declare`, the hooks and hardware defaults through `declare_resolution` -- and that is pinned two ways: no bare assignment to a field survives anywhere a ServerArgs instance -is in reach, and after resolution every declared field agrees with what the -stash says. The second check is the one that keeps the transition honest -- -while a declaration still writes the field immediately, a stash entry and a -field can only disagree if something assigned the field behind the stash's -back. A third check runs the other way: every field resolution moved has to -be explained by the stash, which covers the spellings a source scan cannot -see. +is in reach, and after resolution `resolution_result` answers for every declared +field with what the stash holds. The second check is what the stash is measured +against: the two can disagree only if something wrote behind the stash's back. A +third check runs the other way -- every field resolution moved has to be +explained by the stash, which covers the spellings a source scan cannot see. """ import ast @@ -574,10 +572,9 @@ def test_late_resolution_reaches_the_projection(self): The parser detection and the LoRA normalization run at launcher stage -- they need a tokenizer, a chat template, an adapter directory -- and they - write through `declare_late_resolution`. If those writes only reached - the fields, the bags would describe the *unresolved* value: a server - launched with `--reasoning-parser auto` would advertise and apply - `auto` after detection had already replaced it. + declare through `declare_late_resolution`. The declaration is the only + home for what they decide: the record keeps `--reasoning-parser auto`, + and the bags a process publishes carry the detected parser. A real model path, not the dummy one: a dummy record never materializes, so its `resolve_once` re-runs and re-snapshots the raw input from @@ -599,15 +596,22 @@ def test_late_resolution_reaches_the_projection(self): ) publish(server_args, role="tokenizer") self.assertEqual(get_serving().reasoning_parser, "qwen3") - self.assertEqual(server_args.reasoning_parser, get_serving().reasoning_parser) + self.assertEqual( + server_args.reasoning_parser, + "auto", + "the record is the operator's input; late resolution declares, it " + "does not write back", + ) def test_validation_can_still_resolve_before_the_record_is_published(self): - """The LoRA checks normalize in place, so they must precede publish. + """The LoRA checks resolve, so they must precede publish. `check_server_args` is not read-only: it infers `enable_lora`, parses adapter paths and normalizes target modules through late resolution, which a published record refuses. The launcher order is what keeps this - legal, and this is the assertion that notices if it moves. + legal, and this is the assertion that notices if it moves. What those + declarations decide reaches the bags; the record keeps the raw form the + operator passed. """ from sglang.srt.runtime_context import get_lora, publish, reset_context @@ -621,9 +625,17 @@ def test_validation_can_still_resolve_before_the_record_is_published(self): self.addCleanup(reset_context) server_args.check_server_args() publish(server_args, role="tokenizer") - self.assertEqual(get_lora().enable_lora, server_args.enable_lora) self.assertEqual( - get_lora().lora_target_modules, server_args.lora_target_modules + get_lora().enable_lora, resolution_result(server_args, "enable_lora") + ) + self.assertEqual( + get_lora().lora_target_modules, + resolution_result(server_args, "lora_target_modules"), + ) + self.assertEqual( + server_args.lora_target_modules, + ["q_proj"], + "normalization is a declaration; the record keeps what was passed", ) def test_the_launcher_finishes_resolving_before_it_publishes(self): @@ -674,32 +686,37 @@ def test_the_launcher_finishes_resolving_before_it_publishes(self): f"written:\n " + "\n ".join(too_late), ) - def test_the_stash_is_what_the_projection_answers(self): - """The last declaration for a field is what `resolution_result` returns. + def test_an_undeclared_field_still_holds_the_raw_input(self): + """Nothing writes a field behind the stash's back. - The stash is the resolution result, and every bag is projected from - it, so a disagreement here means the projection is reading something - other than the declarations. + Comparing the stash against `resolution_result` would agree by + construction -- both are the same last-writer-wins walk over + `_resolved_overrides`, spelled forwards and backwards. The independent + source is the record's own `_raw_input` snapshot: a field with no + declaration has to still equal what the caller passed, because the only + sanctioned way to move one is to declare it. """ - mismatches = [] + moved = [] for shape in _SHAPES: server_args = self._resolve(shape) overlay = _stash_overlay(server_args) - for field, declared in overlay.items(): - if field not in _RESOLVED_FIELDS: + raw_input = getattr(server_args, "_raw_input", None) + self.assertTrue(raw_input, f"{shape}: the record kept no raw snapshot") + for field in dataclasses.fields(server_args): + name = field.name + if name in overlay or name not in raw_input: continue - answered = resolution_result(server_args, field) - if answered != declared: - mismatches.append( - f"{shape} -> {field}: projection={answered!r} " - f"stash={declared!r}" + current = getattr(server_args, name, None) + if current != raw_input[name]: + moved.append( + f"{shape} -> {name}: raw={raw_input[name]!r} " + f"field={current!r}" ) self.assertEqual( - mismatches, + moved, [], - "the projection disagrees with the last declaration for a field, so " - "the bags would publish something no resolver decided:\n " - + "\n ".join(mismatches), + "these fields moved without a declaration, so the bags publish one " + "value while the record shows another:\n " + "\n ".join(moved), ) def test_no_immediate_writer_overrides_a_deferred_one(self): diff --git a/test/registered/unit/server_args/test_resolution_is_reproducible.py b/test/registered/unit/server_args/test_resolution_is_reproducible.py index 0c0acabdfead..b64504c7b7ac 100644 --- a/test/registered/unit/server_args/test_resolution_is_reproducible.py +++ b/test/registered/unit/server_args/test_resolution_is_reproducible.py @@ -768,10 +768,15 @@ def _resolved(self): server_args.resolve_once() return server_args - def test_a_bare_replace_would_resolve_a_second_time(self): - """Why the helper exists. If this stops drifting, the pipeline became - idempotent and the helper's reason is gone -- read it again before - deleting either.""" + def test_a_bare_replace_resolves_again_and_lands_in_the_same_place(self): + """A bare copy resolves to the same place: the fields are the raw input. + + `dataclasses.replace` copies the fields, so a bare copy re-runs + resolution over the *same input* the parent got -- the DP-attention + halving and the conservativeness scaling apply once. `replace_resolved` + buys something else: it carries the parent's declarations and its + `model_config`, so the copy answers without resolving at all. + """ parent = self._resolved() bare = dataclasses.replace(parent, dist_init_addr="1.2.3.4:5000") self.assertFalse( @@ -779,14 +784,21 @@ def test_a_bare_replace_would_resolve_a_second_time(self): "a bare replace carried the flag; then this test proves nothing", ) bare.resolve_once() + drifted = { + field.name: ( + resolution_result(parent, field.name), + resolution_result(bare, field.name), + ) + for field in dataclasses.fields(parent) + if field.name not in ("dist_init_addr", "random_seed") + and repr(resolution_result(parent, field.name)) + != repr(resolution_result(bare, field.name)) + } self.assertEqual( - (bare.chunked_prefill_size, round(bare.schedule_conservativeness, 4)), - ( - parent.chunked_prefill_size // 2, - round(parent.schedule_conservativeness * 0.3, 4), - ), - "the second pass no longer drifts; this is the drift the copy " - "helper exists to avoid", + drifted, + {}, + "resolving a bare copy landed somewhere else, so the pipeline is " + "reading its own output again", ) def test_replace_resolved_keeps_the_parents_resolution(self): diff --git a/test/registered/unit/server_args/test_resolution_reads_the_declarations.py b/test/registered/unit/server_args/test_resolution_reads_the_declarations.py new file mode 100644 index 000000000000..3d940ae483e9 --- /dev/null +++ b/test/registered/unit/server_args/test_resolution_reads_the_declarations.py @@ -0,0 +1,612 @@ +"""Resolution reads its own decisions, not the record's fields. + +`declare_resolution` records a decision in the declaration stash and writes +nothing. The fields keep what the caller passed, so a resolver that reads a +field another resolver may have decided reads the raw input -- silently, and +only on the configurations where that other resolver fires. The whole pipeline +therefore reads through `resolving_view` (or `ServerArgs._resolved()`, which is +the same view spelled as the record's own member), and this pins that there is +nothing left reading a field directly. + +Subjects: every function in `arg_groups/` that takes a config, every +`ServerArgs` handler the dispatcher reaches, and every member of `ServerArgs` / +`PortArgs` -- the members are reached from the hooks and from business code, +which the handler walk cannot see, and a member that recomputes from a raw field +decides from what was typed. All three +are derived -- a new hook file, a new handler or a new member is covered the +moment it is written. Readers *outside* those +two -- the platform defaults, `ModelConfig`, the spec-algo hook -- are reached by +resolution too and have moved to the view as well, but enumerating them needs +the call-graph derivation `test_resolution_reads_no_bag` owns; this file pins +the two scopes it can derive exactly. +""" + +import ast +import dataclasses +import pathlib +import re + +import sglang +from sglang.srt.server_args import ServerArgs +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=8, suite="base-a-test-cpu") + +_SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt" +_FIELDS = frozenset(field.name for field in dataclasses.fields(ServerArgs)) + +# Names a config travels under. `args` is included because the platform hooks +# use it; a false positive would be a function taking an argparse Namespace and +# reading an attribute that happens to be a ServerArgs field name, which the +# allowlist below would then have to carry. +_HOLDER_NAMES = frozenset({"server_args", "sa", "args"}) + + +def _holders(fn): + names = { + arg.arg + for arg in list(fn.args.posonlyargs) + + list(fn.args.args) + + list(fn.args.kwonlyargs) + if arg.arg in _HOLDER_NAMES + } + for arg in ( + list(fn.args.posonlyargs) + list(fn.args.args) + list(fn.args.kwonlyargs) + ): + annotation = arg.annotation + text = ( + annotation.value + if isinstance(annotation, ast.Constant) + else ( + annotation.id + if isinstance(annotation, ast.Name) + else annotation.attr if isinstance(annotation, ast.Attribute) else None + ) + ) + if text == "ServerArgs": + names.add(arg.arg) + return names + + +def _field_reads(fn, holders): + for node in ast.walk(fn): + if ( + isinstance(node, ast.Attribute) + and node.attr in _FIELDS + and isinstance(node.value, ast.Name) + and node.value.id in holders + and isinstance(node.ctx, ast.Load) + ): + yield node.lineno, node.attr + + +def _resolution_handlers(): + """The `ServerArgs` methods the dispatcher reaches, transitively.""" + tree = ast.parse((_SRT / "server_args.py").read_text(encoding="utf-8-sig")) + cls = next( + node + for node in ast.walk(tree) + if isinstance(node, ast.ClassDef) and node.name == "ServerArgs" + ) + methods = { + node.name: node + for node in cls.body + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + } + assert "_run_resolution_pipeline" in methods, "the dispatcher was renamed" + seen, stack = set(), ["_run_resolution_pipeline"] + while stack: + name = stack.pop() + if name in seen or name not in methods: + continue + seen.add(name) + for node in ast.walk(methods[name]): + if ( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and isinstance(node.func.value, ast.Name) + and node.func.value.id == "self" + ): + stack.append(node.func.attr) + return {name: methods[name] for name in seen} + + +_DECLARERS = frozenset( + { + "_declare", + "declare_resolution", + "declare_late_resolution", + "declare_direct_writes", + } +) + + +def _declared_fields(): + """The fields resolution decides, read off every shape that reaches the stash. + + A keyword on a `declare_*` call is only one shape: the model-override and + post-process passes build a mapping instead (`MODEL_OVERRIDES` literals, + `overrides["dtype"] = ...`, a returned dict), and late resolution splats a + variable-keyed one. Deriving from keywords alone leaves nineteen fields + outside the subject set, `dtype` and `reasoning_parser` among them. + """ + fields = set() + # The declaration calls live wherever a resolver does; the mapping channels + # only exist where the override providers and post-process passes are. + keyword_sources = [_SRT / "server_args.py"] + for sub in ("arg_groups", "hardware_backend", "parser"): + keyword_sources += sorted((_SRT / sub).rglob("*.py")) + mapping_sources = {_SRT / "server_args.py", *(_SRT / "arg_groups").rglob("*.py")} + for path in keyword_sources: + tree = ast.parse(path.read_text(encoding="utf-8-sig")) + for node in ast.walk(tree): + # 1. `declare_resolution(sa, src, page_size=64)` and its siblings + if isinstance(node, ast.Call): + name = ( + node.func.id + if isinstance(node.func, ast.Name) + else getattr(node.func, "attr", None) + ) + if name in _DECLARERS: + for keyword in node.keywords: + if keyword.arg: + fields.add(keyword.arg) + elif isinstance(keyword.value, ast.Dict): + fields.update(_string_keys(keyword.value)) + # 2. every mapping literal in the files that declare through one: + # the MODEL_OVERRIDES tables, the dicts the override providers + # return, the ones the post-process passes build. Scanning + # unrelated files here would collect a plain kwarg dict + # (`tokenizer_config={"trust_remote_code": ...}`) and turn a + # passthrough read into a violation. + if isinstance(node, ast.Dict) and path in mapping_sources: + fields.update(_string_keys(node)) + # 3. `overrides["field"] = ...` + if ( + path in mapping_sources + and isinstance(node, ast.Assign) + and isinstance(node.targets[0], ast.Subscript) + and isinstance(node.targets[0].slice, ast.Constant) + and isinstance(node.targets[0].slice.value, str) + ): + fields.add(node.targets[0].slice.value) + return frozenset(fields & _FIELDS) + + +def _string_keys(node: ast.Dict) -> set: + return { + key.value + for key in node.keys + if isinstance(key, ast.Constant) and isinstance(key.value, str) + } + + +def _record_members(): + """Every member of `ServerArgs` / `PortArgs`, by class and name.""" + tree = ast.parse((_SRT / "server_args.py").read_text(encoding="utf-8-sig")) + members = {} + for node in ast.walk(tree): + if isinstance(node, ast.ClassDef) and node.name in ("ServerArgs", "PortArgs"): + for member in node.body: + if isinstance(member, (ast.FunctionDef, ast.AsyncFunctionDef)): + members[f"{node.name}.{member.name}"] = member + return members + + +def _config_reading_helpers(): + """Module functions that load a decided field off the config they are handed. + + A member that hands them `self`, or a call site that hands them a record, + reads the raw input through the callee -- the shape neither an attribute + scan nor a `getattr` scan can see, because the field name is spelled in the + helper and the record is spelled at the call site. + """ + decided = _declared_fields() + helpers = {} + sources = [_SRT / "server_args.py"] + sorted((_SRT / "arg_groups").rglob("*.py")) + for path in sources: + tree = ast.parse(path.read_text(encoding="utf-8-sig")) + for fn in ast.walk(tree): + if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)): + continue + params = { + arg.arg + for arg in list(fn.args.posonlyargs) + + list(fn.args.args) + + list(fn.args.kwonlyargs) + } - {"self", "cls"} + if not params: + continue + reads = { + node.attr + for node in ast.walk(fn) + if isinstance(node, ast.Attribute) + and isinstance(node.value, ast.Name) + and node.value.id in params + and isinstance(node.ctx, ast.Load) + and node.attr in decided + } + if reads: + helpers[fn.name] = sorted(reads) + return helpers + + +# The accessors that hand back the process-global record itself. A helper that +# is handed one of these reads the raw input exactly as a bare `self` would. +_RECORD_ACCESSORS = frozenset({"get_server_args", "global_server_args"}) + +# `self._server_args` is the same record under a private name; the scan has to +# see it or a reader inside the context object escapes every shape above. +_RECORD_ATTR = re.compile(r"^_*(server_args|sa)$") + + +def _record_arguments(node, aliases=frozenset()): + """The bare-record arguments of a call. + + Four spellings reach a helper with a record: the bare name (`self`, `sa`), + an attribute (`runner.server_args`), the process-global accessor called + inline (`get_server_args()`), and a local bound to either of the last two + earlier in the same function. + """ + out = [] + for arg in node.args: + if isinstance(arg, ast.Name) and arg.id in ("self", "server_args", "sa"): + out.append(arg.id) + elif isinstance(arg, ast.Attribute) and _RECORD_ATTR.match(arg.attr or ""): + out.append(ast.unparse(arg)) + elif ( + isinstance(arg, ast.Call) + and isinstance(arg.func, ast.Name) + and arg.func.id in _RECORD_ACCESSORS + ): + out.append(ast.unparse(arg)) + elif isinstance(arg, ast.Name) and arg.id in aliases: + out.append(arg.id) + return out + + +def _record_aliases(function): + """Locals bound to the record under a name of their own. + + `_sa = getattr(runner, "server_args", None)`, `cfg = get_server_args()` and + `engine_args = ServerArgs.from_cli_args(args)` all put the record behind a + name the argument scan does not recognise, so a later + `getattr(_sa, "")` reads what the operator typed. + """ + aliases = set() + for node in ast.walk(function): + if not isinstance(node, ast.Assign) or len(node.targets) != 1: + continue + target = node.targets[0] + if not isinstance(target, ast.Name): + continue + value = node.value + if isinstance(value, ast.Attribute): + if _RECORD_ATTR.match(value.attr or ""): + aliases.add(target.id) + continue + if not isinstance(value, ast.Call): + continue + func = value.func + if isinstance(func, ast.Name): + if func.id in _RECORD_ACCESSORS or func.id == "ServerArgs": + aliases.add(target.id) + elif ( + func.id == "getattr" + and len(value.args) >= 2 + and isinstance(value.args[1], ast.Constant) + and _RECORD_ATTR.match(str(value.args[1].value)) + ): + aliases.add(target.id) + elif ( + isinstance(func, ast.Attribute) + and func.attr in ("from_cli_args", "replace_resolved") + and isinstance(func.value, ast.Name) + and func.value.id == "ServerArgs" + ): + aliases.add(target.id) + return aliases + + +# The one reader for which the raw field is the right answer. The gateway sizes +# its worker pool from the operator's requested replica count; `--dwdp-size` +# makes resolution declare a `dp_size` describing one multi-rank server's +# internal topology, so reading the decision there would spawn dp_size +# single-rank children and ask for dp_size^2 GPUs. A new entry here needs that +# kind of reason next to it. +_NO_RESOLVED_SURFACE = frozenset( + { + ( + "sgl-model-gateway/bindings/python/src/sglang_router/launch_server.py", + "server_args.dp_size", + ), + } +) + + +def _is_record_base(node, aliases): + """Is this expression the record itself? + + A local bound to one, a parameter that carries one (`server_args`, `sa`, + `engine_args`), or an attribute holding one (`self._server_args`). + """ + if isinstance(node, ast.Name): + return node.id in aliases + if isinstance(node, ast.Attribute): + return bool(_RECORD_ATTR.match(node.attr or "")) + return False + + +def _record_handoff_offenders(rel, tree, helpers, decided, is_record=False): + """Every way a decided leaf is reached through a record in one module. + + Two shapes, both scanned under the record aliases the function binds: + handing the record to a helper that loads a decided field, and loading one + off the alias directly (`alias.` or `getattr(alias, "")`). The + second is what the MiniMax backend spelled, and an argument scan cannot see + it -- the leaf never appears at a call site. + """ + offenders, seen = [], set() + + def record(lineno, text): + if (lineno, text) in seen: + return + seen.add((lineno, text)) + offenders.append(f"{rel}:{lineno} {text}") + + scopes = [(tree, frozenset())] + [ + (fn, _record_aliases(fn)) + for fn in ast.walk(tree) + if isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)) + ] + for scope, aliases in scopes: + for node in ast.walk(scope): + if isinstance(node, ast.Call): + name = ( + node.func.id + if isinstance(node.func, ast.Name) + else getattr(node.func, "attr", None) + ) + if name in helpers: + for arg in _record_arguments(node, aliases): + if arg == "self" and not is_record: + continue + record( + node.lineno, + f"{name}({arg}) reads {', '.join(helpers[name])}", + ) + if ( + isinstance(node.func, ast.Name) + and node.func.id == "getattr" + and len(node.args) >= 2 + and isinstance(node.args[0], ast.Name) + and node.args[0].id in aliases + and isinstance(node.args[1], ast.Constant) + and node.args[1].value in decided + ): + record( + node.lineno, + f'getattr({node.args[0].id}, "{node.args[1].value}")', + ) + elif ( + isinstance(node, ast.Attribute) + and isinstance(node.ctx, ast.Load) + and node.attr in decided + and _is_record_base(node.value, aliases) + ): + record(node.lineno, f"{ast.unparse(node.value)}.{node.attr}") + return offenders + + +# Source the scanner must read the same way whether or not the tree happens to +# contain these shapes today. The first four are the spellings that reached +# production and were converted; the last two are the legal forms next to them, +# which have to stay quiet or the guard is unusable. +_SPELLINGS = """ +def hands_the_alias_to_a_helper(runner): + _sa = getattr(runner, "server_args", None) + return m3_fp8_attn_gemm_enabled(_sa) + + +def loads_a_leaf_off_the_alias(runner): + _sa = getattr(runner, "server_args", None) + return getattr(_sa, "speculative_num_draft_tokens", None) + + +def reads_a_leaf_through_the_alias(runner): + sa_local = runner.server_args + return sa_local.attention_backend + + +def hands_the_accessor_to_a_helper(): + return compute_world_size(get_server_args()) + + +def reads_the_view(runner): + cfg = resolving_view(runner.server_args) + return cfg.attention_backend + + +def reads_an_undecided_leaf(runner): + _sa = runner.server_args + return _sa.tp_size + + +def reads_a_leaf_off_a_private_attribute(self): + return self._server_args.attention_backend + + +def reads_a_leaf_off_a_constructed_record(cli): + engine_args = ServerArgs.from_cli_args(cli) + engine_args.resolve_once() + return engine_args.attention_backend +""" + + +class TestResolutionReadsTheDeclarations(CustomTestCase): + def test_no_hook_reads_a_field_off_the_record(self): + offenders = [] + files = sorted((_SRT / "arg_groups").glob("*.py")) + self.assertGreater(len(files), 5, "the hook scan found almost nothing") + for path in files: + rel = f"arg_groups/{path.name}" + tree = ast.parse(path.read_text(encoding="utf-8-sig")) + for fn in ast.walk(tree): + if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)): + continue + holders = _holders(fn) + if not holders: + continue + for lineno, field in _field_reads(fn, holders): + offenders.append(f"{rel}:{lineno} {fn.name} reads .{field}") + self.assertEqual( + offenders, + [], + "a resolution hook reads a field off the record; the field holds the " + "raw input, so this decides from what was typed rather than from " + "what resolution decided. Read `resolving_view(server_args)`:\n " + + "\n ".join(offenders), + ) + + def test_no_handler_reads_a_field_off_self(self): + handlers = _resolution_handlers() + self.assertGreater( + len(handlers), 50, f"only {len(handlers)} handlers were reached" + ) + offenders = [] + for name, fn in sorted(handlers.items()): + for lineno, field in _field_reads(fn, {"self"}): + offenders.append(f"server_args.py:{lineno} {name} reads self.{field}") + self.assertEqual( + offenders, + [], + "a resolution handler reads its own field; the field holds the raw " + "input. Bind `cfg = resolving_view(self)` and read that:\n " + + "\n ".join(offenders), + ) + + def test_no_member_recomputes_from_a_raw_field(self): + decided = _declared_fields() + self.assertGreater( + len(decided), 100, f"the declaration set derived only {len(decided)} fields" + ) + members = _record_members() + self.assertGreater(len(members), 100, f"only {len(members)} members were found") + offenders = [] + for name, fn in sorted(members.items()): + holders = _holders(fn) | {"self"} + for lineno, field in _field_reads(fn, holders): + if field in decided: + offenders.append(f"server_args.py:{lineno} {name} reads .{field}") + self.assertEqual( + offenders, + [], + "a record member recomputes from a field resolution decides; the " + "field holds the raw input, so the member answers for what was " + "typed. Bind `cfg = resolving_view(self)` and read that:\n " + + "\n ".join(offenders), + ) + + def test_no_reader_hands_the_record_to_a_config_helper(self): + helpers = _config_reading_helpers() + self.assertGreater( + len(helpers), 5, f"the helper derivation found only {len(helpers)}" + ) + decided = _declared_fields() + offenders = [] + # `scripts/`, `examples/` and the gateway binding are outside the + # package but hold records they resolve themselves, and every reader + # this scan found in them was reading a field resolution fills in. + _REPO = _SRT.parent.parent.parent + roots = ( + [_SRT] + + [_SRT.parent / d for d in ("benchmark", "lang")] + + [ + _REPO / d + for d in ( + "scripts", + "examples", + "sgl-model-gateway/bindings/python/src", + ) + ] + ) + for root in roots: + if not root.exists(): + continue + for path in sorted(root.rglob("*.py")): + # Package files keep their `srt/...` spelling (the skips below + # key on it); the repo-level roots are named from the repo. + try: + rel = path.relative_to(_SRT.parent).as_posix() + except ValueError: + rel = path.relative_to(_REPO).as_posix() + if rel.startswith(("srt/arg_groups/", "multimodal_gen/")): + continue + try: + tree = ast.parse(path.read_text(encoding="utf-8-sig")) + except SyntaxError: + continue + offenders += _record_handoff_offenders( + rel, tree, helpers, decided, is_record=rel == "srt/server_args.py" + ) + offenders = [ + line + for line in offenders + if (line.split(":", 1)[0], line.split(" ", 1)[1]) + not in _NO_RESOLVED_SURFACE + ] + self.assertEqual( + offenders, + [], + "a caller hands the record to a helper that loads a field " + "resolution decides; the helper then reads the raw input. Hand it " + "`resolving_view(record)` (or the published bag):\n " + + "\n ".join(offenders), + ) + + def test_the_scan_sees_every_spelling_that_reached_production(self): + """Every spelling that reached production, pinned next to the scanner. + + A shape the scan stops seeing is a silent hole, so each one is listed + here with the legal forms beside it and the flagged set compared + exactly. + + What it does not reach: a record that arrives as a *parameter* and was + resolved by the caller (`scripts/playground/bench_speculative.py` hands + `main(args, server_args)` one). Binding that would need the call graph, + and naming a parameter `server_args` is also how the resolution-time + readers spell a view. + """ + helpers = _config_reading_helpers() + decided = _declared_fields() + for name in ("m3_fp8_attn_gemm_enabled", "compute_world_size"): + self.assertIn(name, helpers, f"the helper derivation lost {name}") + for field in ("speculative_num_draft_tokens", "attention_backend"): + self.assertIn(field, decided, f"the declared set lost {field}") + + offenders = _record_handoff_offenders( + "sample.py", ast.parse(_SPELLINGS), helpers, decided + ) + flagged = {line.split(" ", 1)[1] for line in offenders} + self.assertEqual( + flagged, + { + "m3_fp8_attn_gemm_enabled(_sa)" + " reads " + ", ".join(helpers["m3_fp8_attn_gemm_enabled"]), + 'getattr(_sa, "speculative_num_draft_tokens")', + "sa_local.attention_backend", + "self._server_args.attention_backend", + "engine_args.attention_backend", + "compute_world_size(get_server_args())" + " reads " + ", ".join(helpers["compute_world_size"]), + }, + "the scan lost a spelling, or started flagging a legal one:\n " + + "\n ".join(sorted(flagged)), + ) + + +if __name__ == "__main__": + import unittest + + unittest.main() diff --git a/test/registered/unit/test_server_args_no_instance_mutation_entry.py b/test/registered/unit/test_server_args_no_instance_mutation_entry.py index 1afbf9fa1c22..d6b4e616bc00 100644 --- a/test/registered/unit/test_server_args_no_instance_mutation_entry.py +++ b/test/registered/unit/test_server_args_no_instance_mutation_entry.py @@ -6,9 +6,9 @@ write desyncs every namespace reader, and a copy invites publishing stale variants. Both are gone: post-publish changes go to the bags (``get_context().override``), a value one runner or worker owns travels as a -constructor argument, and late launcher-stage resolution writes in place -through ``arg_groups.overrides.declare_late_resolution``, which refuses the -published instance. +constructor argument, and late launcher-stage resolution declares through +``arg_groups.overrides.declare_late_resolution``, which writes no field and +refuses the published instance. The textual half of this guard matters because the resolution pipeline's own file is exempt from the mutation ratchet: a ``self.override(...)`` there — exactly From 19ad007856ef17d2cf4c39a8396bc6294ec20995 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sun, 23 Aug 2026 13:09:28 +0000 Subject: [PATCH 25/26] docs: the skill describes the record as raw input The runtime-context skill still said declarations materialize onto the fields at the end of `__post_init__`, and that a constructor may publish. Both changed: the record holds the raw input, resolution reads through `resolving_view` / `resolved_view`, `initialize_moe_config` takes no record, constructors assert instead of publishing, and a test that asserts what resolution decided reads `resolution_result` rather than the field. --- .../skills/sglang-runtime-context/SKILL.md | 72 +++++++++++++------ python/sglang/srt/server_args.py | 7 +- 2 files changed, 54 insertions(+), 25 deletions(-) diff --git a/.claude/skills/sglang-runtime-context/SKILL.md b/.claude/skills/sglang-runtime-context/SKILL.md index bfdae19ce0ca..c76a56399fde 100644 --- a/.claude/skills/sglang-runtime-context/SKILL.md +++ b/.claude/skills/sglang-runtime-context/SKILL.md @@ -22,12 +22,18 @@ flags/resources/forward tiers. ## Config: publish + namespace bags -**`ServerArgs` is a pristine seed. Business code never reads it for decisions — -resolved configuration lives in the namespace bags.** +**`ServerArgs` holds the raw input and nothing else. Resolution writes no field: +it declares, and the declarations are what the namespace bags are projected from. +Business code never reads the record for a decision — and after this cut, a field +read there answers with what the operator typed, not with what resolution +decided.** - Every publishing process entry calls `publish(server_args, role=...)` (`run_scheduler_process`, the Ray `SchedulerActor`, the DP controller, tokenizer, - detokenizer, encoder, weight-cache daemon, ...); the roles are enumerated once, + detokenizer, encoder, weight-cache daemon, the multi-tokenizer worker, the + spawned encoder TP/DP workers, the benchmark work functions, ...); constructors + do not publish — `ModelRunner`, `TokenizerManager` and `MMEncoder` call + `assert_published` and fail loudly if an entry forgot. The roles are enumerated once, as the keys of `ROLE_NAMESPACE_SETS` — there is no `launcher` role, the launch path publishes as `tokenizer`. The remaining non-publisher is `run_multi_detokenizer_router_process`: it *is* handed a `ServerArgs`, and uses @@ -82,13 +88,15 @@ resolved configuration lives in the namespace bags.** the target runner. - **Late launcher-stage resolution (pre-publish)**: a few rules cannot run inside `__post_init__` — LoRA normalization, and the auto-parser detection that needs a - tokenizer/chat-template load. They are resolution, not mutation, and they write - **in place** via `arg_groups.overrides.declare_late_resolution(server_args, - source, **fields)`, which refuses the published instance. In place is the point: - every holder of that object must see the resolved value — the HTTP server, the - multi-tokenizer workers it is serialized for, the schedulers it forks. Returning a - variant here is a bug: the launcher rebinds its local and everyone else keeps the - unresolved object. + tokenizer/chat-template load. They are resolution, not mutation, and they + **declare** via `arg_groups.overrides.declare_late_resolution(server_args, + source, **fields)`, which refuses the published instance. The declaration lands + in the stash on that very object, so every holder of it carries the decision — + the HTTP server, the multi-tokenizer workers it is serialized for, the + schedulers it forks — and each of them publishes bags projected from it. The + fields stay the operator's input; `resolution_result(sa, field)` and the bags + are what answer for the decision. Returning a variant here is a bug: the + launcher rebinds its local and everyone else keeps the unresolved object. - **A value another runner / worker owns is a constructor argument, not a config copy.** The draft worker's `context_length`, load format and attention backend travel as arguments to `TpModelWorker` / `ModelRunner` and live on the runner @@ -321,9 +329,9 @@ what sits beside it is residue, not a family — and not for one single reason: 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. + argument. Its `self.server_args` reads are on this list as a construction-path + convention, and the residual is real: they answer with the raw input, so a leaf + resolution decided and a post-publish `override` both pass them by. 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 @@ -358,14 +366,30 @@ if you do it, say so in the test. ### Mid-resolution reads (inside the pipeline only) -Resolution itself still runs in `__post_init__`: handlers and hooks read the -in-flight state through `resolved_view(server_args)` / `self._resolved()`, fields are -read-only during resolution, and declarations materialize once at the very end of -`__post_init__` (gate order, last writer wins) — *then* `publish` snapshots the -resolved values into the bags. `resolved_view` is pipeline-internal -(`server_args.py` / `arg_groups/`, plus helpers the pipeline itself invokes -mid-resolution, e.g. `adaptive_spec_params`); do not introduce new -out-of-pipeline call sites. +Resolution runs in `__post_init__` and **writes nothing onto the record**: a +handler declares (`self._declare` / `declare_resolution`), the declaration goes +into the stash, and the fields keep what the caller passed. So a mid-resolution +read of a field answers with the *raw input* — every reader in the pipeline goes +through a view instead: + +- `resolving_view(server_args)` / `self._resolved()` — the live view (walks the + stash per read). This is what handlers and hooks bind, conventionally as + `cfg = resolving_view(self)` at the top of the handler. +- `resolved_view(server_args)` — snapshots the overlay when built, which is what + a post-process pass wants: it reads the state at *its* slot. + +`test_resolution_reads_the_declarations` pins direct field reads at zero over the +two scopes it can derive exactly (every `arg_groups` function taking a config, +every `ServerArgs` handler the dispatcher reaches). Readers the pipeline calls +from elsewhere (`ModelConfig`, the platform defaults, the spec-algo hook) have +moved to the view as well — a field read there is the same bug, just one the +derivation cannot enumerate. + +One consequence worth knowing: because the fields are the raw input, resolving a +bare `dataclasses.replace` copy lands in the same place as the parent — the +pipeline reads only its own input. `replace_resolved` is the way to copy a +resolved record (it carries the declarations and the `model_config` memo, so the +copy does not re-resolve at all). ### Adding a model-specific config adjustment @@ -474,8 +498,12 @@ ONE thread — do not design for TBO threads that don't exist. them explicitly on the mock; `MagicMock(spec=...)` raises on attributes that only exist post-`__init__`, which is the fastest way to find a missed stub. - `reset_context()` in teardown when a test publishes outside a scoped override. -- `ServerArgs(model_path="dummy")` early-returns `__post_init__` (no materialization, no +- `ServerArgs(model_path="dummy")` early-returns the pipeline (few declarations, no strict guard) — fine for lightweight fixtures. +- **Asserting what resolution decided reads `resolution_result(sa, "field")`**, not + `sa.field`: the field is the raw input. Assert the field only when the point of + the case *is* that the record stayed pristine (the FA4 page-size and waterfill + cases do exactly that, and say so). - **Run changed test files per-file** (own process), the way CI does: a monolithic local pytest run lets a context published by an earlier file mask a missing-publish bug in a later one. diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 93141bb782b9..4a7820f28030 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -9786,9 +9786,10 @@ def _resolved(self): def _late_resolution(self, source: str, **fields) -> None: """Resolve fields at the launcher's validation stage (pre-publish). - See ``arg_groups.overrides.declare_late_resolution``: in place, because - every holder of this instance must see the resolved value, and refused - outright once the config is published. + See ``arg_groups.overrides.declare_late_resolution``: the decision goes + to this instance's declaration stash, so every holder of it carries the + decision and publishes bags that answer with it. Refused outright once + the config is published. """ from sglang.srt.arg_groups.overrides import declare_late_resolution From 8c33a7677d25890b51ed6f4a924806acb57d77ba Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Wed, 26 Aug 2026 11:55:10 +0000 Subject: [PATCH 26/26] ci: placeholder so the vehicle gets its own check suite