diff --git a/.claude/skills/sglang-runtime-context/SKILL.md b/.claude/skills/sglang-runtime-context/SKILL.md index 0eea90569526..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 @@ -99,8 +107,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 +122,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 @@ -139,9 +148,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 +158,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 @@ -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 @@ -411,7 +435,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 @@ -473,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/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/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 38addce4b388..f72faea387cf 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,10 +80,14 @@ 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 +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 @@ -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: @@ -538,7 +544,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 +554,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(), ) @@ -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 @@ -681,6 +688,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,14 +890,18 @@ def latency_test( gpu_id, tp_rank, ): - initialize_moe_config(server_args) - initialize_fp8_gemm_config(server_args) - initialize_fp4_gemm_config(server_args) + 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() + 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 @@ -983,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, "_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") + 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/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/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 2f16545287ca..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]``: @@ -29,6 +33,7 @@ from __future__ import annotations +import copy import dataclasses import inspect import json @@ -149,6 +154,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 @@ -169,10 +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 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.""" + 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) @@ -185,9 +224,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): @@ -207,7 +246,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) @@ -223,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))) @@ -256,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: @@ -267,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 @@ -297,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( @@ -311,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 @@ -347,24 +374,12 @@ 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. 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. @@ -380,11 +395,49 @@ 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 -- 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. + """ + 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 - 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)) @@ -545,16 +598,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, @@ -575,7 +629,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( @@ -594,7 +648,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", @@ -606,7 +660,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 @@ -614,7 +668,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 @@ -623,7 +677,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( @@ -637,9 +691,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 ) @@ -652,8 +706,8 @@ 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) - if _dspark_verify_on_decode_backend(backend, q_len, server_args.kv_cache_dtype): + _, 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( "Kimi-K3 DSPARK on SM100/SM103: decode/verify attention backend " @@ -691,7 +745,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 {} @@ -721,6 +776,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] = {} @@ -731,39 +787,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" ) @@ -790,9 +846,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( @@ -800,7 +856,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." @@ -809,22 +865,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" ) @@ -834,8 +890,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 @@ -843,7 +900,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" @@ -853,13 +910,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" @@ -873,10 +931,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 @@ -888,16 +947,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 " @@ -915,7 +970,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 @@ -926,42 +981,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)." ) @@ -972,7 +1021,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 " @@ -980,9 +1029,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(): @@ -1000,9 +1048,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( @@ -1042,6 +1088,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(): @@ -1064,14 +1111,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 = ( @@ -1081,7 +1128,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" @@ -1119,9 +1166,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) @@ -1143,18 +1190,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" @@ -1162,8 +1210,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" @@ -1177,6 +1225,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(): @@ -1187,9 +1236,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" @@ -1219,7 +1268,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( @@ -1227,7 +1277,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(): @@ -1237,29 +1287,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() @@ -1270,7 +1320,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 {} @@ -1279,24 +1330,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 {} @@ -1307,13 +1361,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 @@ -1327,11 +1382,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: @@ -1341,9 +1396,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()) @@ -1371,6 +1426,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] = {} @@ -1379,14 +1435,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}. @@ -1411,6 +1467,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] = {} @@ -1421,7 +1478,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": @@ -1446,22 +1503,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}" @@ -1482,27 +1539,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" @@ -1519,7 +1576,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. @@ -1536,8 +1594,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 { @@ -1549,7 +1607,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 {} @@ -1557,11 +1616,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." ) @@ -1578,10 +1634,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 @@ -1591,8 +1648,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( @@ -1604,6 +1661,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) @@ -1612,7 +1670,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 @@ -1622,8 +1680,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( @@ -1638,13 +1696,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: @@ -1659,6 +1718,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(): @@ -1667,12 +1727,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" ) @@ -2134,7 +2194,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/configs/embedding_model_spec.py b/python/sglang/srt/configs/embedding_model_spec.py index c529e19993de..6bcf5a326cec 100644 --- a/python/sglang/srt/configs/embedding_model_spec.py +++ b/python/sglang/srt/configs/embedding_model_spec.py @@ -221,17 +221,17 @@ 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)`, which is where a decision lives. """ - 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 +239,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 +252,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/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/disaggregation/encoder/grpc_server.py b/python/sglang/srt/disaggregation/encoder/grpc_server.py index 163dc276d0b8..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 +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 @@ -201,6 +206,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() @@ -211,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) @@ -253,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 6d8e0561ad7b..bd6d3f06e847 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" @@ -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 50b75ccaea13..ace0bcd8da19 100644 --- a/python/sglang/srt/disaggregation/encoder/receiver.py +++ b/python/sglang/srt/disaggregation/encoder/receiver.py @@ -35,7 +35,14 @@ 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_model, + get_parallel, + get_serving, +) from sglang.srt.server_args import ServerArgs from sglang.srt.utils import ImageData from sglang.srt.utils.common import safe_pickle_loads @@ -1721,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() @@ -1836,9 +1843,9 @@ 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(server_args), + image_processor_backend=resolve_image_processor_backend(get_mm()), **extra_kwargs, ) @@ -2659,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 330a8b9330b0..a4d88cd9ea9a 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( @@ -1570,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 3d9596d81d58..34b1386db723 100644 --- a/python/sglang/srt/disaggregation/encoder/server.py +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -57,12 +57,14 @@ ) from sglang.srt.observability.metrics_collector import EncoderMetricsCollector from sglang.srt.runtime_context import ( - ensure_published, + assert_published, get_device, get_disagg, get_exec, get_mm, get_model, + get_parallel, + publish, ) from sglang.srt.server_args import ServerArgs from sglang.srt.utils import configure_media_url_security @@ -448,8 +450,8 @@ 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") - logger.info(f"init MMEncoder {rank}/{server_args.tp_size}") + assert_published(server_args, role="encoder") + 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, @@ -469,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, @@ -490,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( @@ -553,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, @@ -1031,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, @@ -2036,6 +2040,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/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/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/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/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/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/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/engine.py b/python/sglang/srt/entrypoints/engine.py index 118b05f852f8..c11ee781ba3b 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, @@ -249,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(): @@ -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 @@ -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 @@ -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, @@ -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_bridge.py b/python/sglang/srt/entrypoints/grpc_bridge.py index 56dfc1acfb1f..4adf1d737091 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 @@ -15,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, @@ -417,13 +417,13 @@ 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) 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/grpc_server.py b/python/sglang/srt/entrypoints/grpc_server.py index c3e22762c76f..0c2a8c559316 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 @@ -204,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( @@ -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/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index df6d09fc310d..a3f0b068e2db 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 @@ -232,6 +233,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}" @@ -282,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() @@ -486,8 +489,10 @@ async def custom_handler(request: Request): get_exec, get_lora, get_model, + get_observability, get_parallel, get_serving, + publish, ) elastic_ep_router.route_class = ORJSONRoute @@ -768,7 +773,7 @@ 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, + config=resolving_view(_global_state.tokenizer_manager.server_args), model_config=model_config, ) return result @@ -798,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`. @@ -811,10 +816,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, @@ -2140,7 +2144,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], @@ -2319,7 +2322,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, @@ -2540,7 +2542,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. @@ -2602,12 +2604,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 ), @@ -2621,10 +2624,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, @@ -2658,10 +2662,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, @@ -2689,12 +2694,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 ), @@ -2707,10 +2713,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", @@ -2748,12 +2755,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/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/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/eplb/eplb_manager.py b/python/sglang/srt/eplb/eplb_manager.py index 00738b40f5c5..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 ), @@ -193,7 +190,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 +200,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,14 +215,13 @@ 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(), ) 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: @@ -235,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/eplb/expert_distribution.py b/python/sglang/srt/eplb/expert_distribution.py index 693d69264141..458ecfc75d80 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, @@ -766,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 ) @@ -900,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/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/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/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/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/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/attention/minimax_sparse_backend.py b/python/sglang/srt/layers/attention/minimax_sparse_backend.py index def19cbc655b..55808c74f214 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 " @@ -244,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 ) @@ -261,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/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index 4191089ed6a8..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,20 +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.""" +def init_cp_strategy( + *, enable_prefill_cp: bool, cp_size: int, cp_strategy: str +) -> None: + """Bind the CP strategy for this process. + + 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 getattr(server_args, "enable_prefill_cp", False): + if not enable_prefill_cp: _STRATEGY = None return - cp_size = getattr(server_args, "attn_cp_size", 1) if cp_size <= 1: _STRATEGY = None return - kind = ContextParallelStrategyKind.from_string(server_args.cp_strategy) + kind = ContextParallelStrategyKind.from_string(cp_strategy) if kind == ContextParallelStrategyKind.ZIGZAG: from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy @@ -261,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={server_args.cp_strategy!r}" + f"Unsupported cp_strategy kind {kind} for cp_strategy={cp_strategy!r}" ) @@ -277,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/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/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/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/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/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py index 7f4e510241e0..2a844cd6790b 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 the operator's input. 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 1b5858d71fc7..52588dbbf726 100644 --- a/python/sglang/srt/layers/quantization/fp4_utils.py +++ b/python/sglang/srt/layers/quantization/fp4_utils.py @@ -2,10 +2,11 @@ import logging from enum import Enum -from typing import TYPE_CHECKING, Optional +from typing import Optional import torch +from sglang.srt.runtime_context import get_exec from sglang.srt.utils.common import ( get_device_capability, is_cuda, @@ -13,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__) @@ -142,11 +140,11 @@ 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 = 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/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/lora/lora_manager.py b/python/sglang/srt/lora/lora_manager.py index ba07718a05e3..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 @@ -1030,7 +1028,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/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 " 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/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/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 56b2200cabad..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, ) @@ -416,7 +415,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 @@ -903,15 +902,15 @@ 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 - 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 +1288,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 +1340,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 +1491,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 +2087,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 +2210,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, @@ -4415,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/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..41923444b8dd 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] @@ -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/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index c21fe1754c92..6ec332baa001 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 @@ -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 @@ -2040,7 +2038,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 @@ -3593,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/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 b1900ea65d21..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, ) @@ -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, @@ -542,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( @@ -569,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), @@ -581,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( [ @@ -612,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), @@ -623,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( [ @@ -655,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. @@ -741,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 = [ @@ -1040,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)} @@ -1066,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( @@ -1113,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( @@ -1121,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", ) @@ -1515,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, @@ -1953,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/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 068fda733023..16f50edc8b34 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, @@ -1152,7 +1153,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, @@ -1263,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, ) @@ -1356,7 +1356,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, @@ -1381,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, } @@ -1396,14 +1395,13 @@ 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, ), ) 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, @@ -1536,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( @@ -2199,10 +2197,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/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/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index be0a80da3e3c..6d7843dc01b4 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 ) @@ -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/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/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/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index b990b4914c49..d4d6e3e54f24 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, @@ -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. @@ -332,20 +331,16 @@ 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 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: @@ -433,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) @@ -594,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, ) @@ -705,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, ) @@ -730,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, @@ -1079,9 +1069,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() @@ -1094,9 +1082,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() @@ -1116,7 +1102,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, @@ -1127,7 +1112,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, ) @@ -1198,11 +1182,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, @@ -1217,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, ) @@ -1284,7 +1268,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, @@ -1474,13 +1457,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: @@ -1929,7 +1910,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..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 @@ -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, @@ -311,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/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/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/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/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..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,11 +70,10 @@ 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 - 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 +81,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..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 @@ -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/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/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/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/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/python/sglang/srt/parser/template_detection.py b/python/sglang/srt/parser/template_detection.py index d353aa1c60ca..766243bfd07a 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 "" @@ -701,24 +702,26 @@ 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( 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 +734,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/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..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__) @@ -105,18 +105,13 @@ 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 - ) + return compute_world_size(get_parallel().config) def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]: @@ -274,16 +269,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 +292,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 +333,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 +369,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 +448,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/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 20fd371209fa..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). @@ -969,11 +975,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 +991,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 @@ -1077,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 @@ -1374,29 +1382,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. +def assert_published(server_args, *, role: str) -> RuntimeContext: + """This record, under this role, is already published -- or fail loud. - 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. + 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. - 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. + 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/srt/server_args.py b/python/sglang/srt/server_args.py index fb5c69cdff8a..4a7820f28030 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 @@ -3708,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( @@ -3720,26 +3721,39 @@ 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. - self._declarations_materialized = True + # 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]: + """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 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. + """ + 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. `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 @@ -3754,7 +3768,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 @@ -3764,7 +3778,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) @@ -3775,7 +3789,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: @@ -3824,6 +3838,8 @@ def _run_resolution_pipeline(self): # _handle_model_specific_adjustments never runs. self._resolved_overrides = [] + cfg = resolving_view(self) + from sglang.srt.arg_groups.mega_moe_hook import handle_mega_moe handle_mega_moe(self) @@ -3835,7 +3851,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() @@ -3896,8 +3912,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) @@ -4005,21 +4020,16 @@ 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): - 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", @@ -4031,7 +4041,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, @@ -4068,8 +4079,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 " @@ -4095,7 +4106,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", @@ -4136,7 +4147,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() @@ -4152,15 +4163,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. @@ -4184,7 +4195,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." @@ -4205,13 +4216,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 @@ -4229,20 +4241,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 " @@ -4250,23 +4263,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 @@ -4275,7 +4287,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" ), ) @@ -4283,47 +4295,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." @@ -4337,7 +4350,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. " @@ -4346,42 +4359,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 " @@ -4391,19 +4407,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. " @@ -4435,7 +4451,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, " @@ -4451,7 +4467,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, @@ -4459,18 +4475,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( @@ -4480,41 +4496,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." @@ -4538,17 +4554,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(), @@ -4556,14 +4573,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 @@ -4573,22 +4590,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, @@ -4596,7 +4613,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, @@ -4612,6 +4629,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 @@ -4641,42 +4659,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", @@ -4687,8 +4706,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=( @@ -4708,21 +4728,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", @@ -4730,22 +4752,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 @@ -4756,10 +4779,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] @@ -4773,9 +4797,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] @@ -4786,6 +4811,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() @@ -4799,7 +4825,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." @@ -4807,16 +4833,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): @@ -4826,8 +4853,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 @@ -4838,7 +4865,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: @@ -4851,36 +4879,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(): @@ -4902,6 +4930,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 @@ -4910,7 +4939,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 @@ -4921,27 +4950,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 = [ ( @@ -4949,8 +4980,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(), @@ -4966,7 +4997,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 @@ -4974,50 +5005,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 @@ -5037,12 +5069,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", @@ -5063,11 +5095,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(): @@ -5076,7 +5109,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): @@ -5085,10 +5118,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" @@ -5103,15 +5137,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; " @@ -5126,22 +5161,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", @@ -5179,14 +5215,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, @@ -5196,59 +5233,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, @@ -5257,7 +5294,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, @@ -5266,7 +5303,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( @@ -5288,34 +5325,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 ) @@ -5327,31 +5364,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) @@ -5374,14 +5409,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." @@ -5392,41 +5427,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 @@ -5442,47 +5478,49 @@ 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) 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(): @@ -5506,9 +5544,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" ): @@ -5520,10 +5559,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] @@ -5552,19 +5592,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 @@ -5618,12 +5659,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, @@ -5633,7 +5675,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 @@ -5644,22 +5686,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, ) @@ -5744,7 +5786,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) ): @@ -5757,17 +5799,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 @@ -5781,15 +5823,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 " @@ -5801,8 +5843,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 " @@ -5814,17 +5856,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 " @@ -5834,7 +5876,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 @@ -5843,8 +5885,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: @@ -5869,7 +5911,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( @@ -5955,7 +5997,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 @@ -5978,14 +6020,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 @@ -5997,7 +6039,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 " @@ -6023,7 +6065,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). @@ -6252,6 +6294,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() @@ -6276,8 +6319,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 @@ -6313,6 +6356,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 @@ -6338,24 +6382,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 []) ): @@ -6398,7 +6442,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 @@ -6419,7 +6463,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 " @@ -6436,7 +6480,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 @@ -6467,7 +6511,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( @@ -6477,8 +6522,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() @@ -6486,7 +6532,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( @@ -6564,7 +6610,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." @@ -6572,24 +6619,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." ) @@ -6599,7 +6648,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 " @@ -6609,12 +6658,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." @@ -6636,22 +6685,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) @@ -6659,10 +6710,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 @@ -6678,7 +6729,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 @@ -6686,7 +6737,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 " @@ -6704,34 +6755,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 ( @@ -6759,7 +6810,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 " @@ -6775,9 +6826,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 @@ -6785,12 +6836,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 @@ -6803,14 +6854,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( @@ -6831,8 +6882,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", @@ -6843,23 +6894,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 " @@ -6869,17 +6920,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", @@ -6890,36 +6942,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: @@ -6946,7 +6998,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 @@ -6957,38 +7010,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 " @@ -7003,70 +7056,75 @@ 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 - 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): - 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 " @@ -7075,7 +7133,7 @@ def _handle_dwdp(self): self._declare( "_handle_dwdp", - dp_size=self.dwdp_size, + dp_size=cfg.dwdp_size, ) self._declare( "_handle_dwdp", @@ -7092,9 +7150,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, @@ -7112,8 +7170,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, " @@ -7123,6 +7181,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, @@ -7130,8 +7189,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 " @@ -7145,23 +7204,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 @@ -7170,14 +7229,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 @@ -7197,6 +7256,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, @@ -7215,7 +7275,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": @@ -7226,7 +7286,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", @@ -7290,12 +7350,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 @@ -7305,29 +7366,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() @@ -7359,6 +7420,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, @@ -7375,13 +7437,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", @@ -7392,17 +7454,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`). @@ -7414,11 +7476,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" @@ -7435,7 +7497,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", @@ -7445,7 +7507,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(), ( @@ -7455,12 +7517,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", @@ -7469,7 +7531,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)." ) @@ -7481,7 +7543,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", @@ -7491,7 +7553,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(), ( @@ -7502,17 +7564,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", @@ -7525,8 +7590,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", @@ -7537,23 +7602,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." ) @@ -7562,65 +7628,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." @@ -7628,16 +7694,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." ) @@ -7645,57 +7711,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 " @@ -7714,18 +7780,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 " @@ -7746,6 +7812,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 " @@ -7754,20 +7821,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, @@ -7789,26 +7856,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 " @@ -7821,13 +7889,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." @@ -7842,7 +7910,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 " @@ -7851,7 +7919,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." @@ -7867,8 +7935,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, ( @@ -7895,11 +7964,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 ), ) @@ -7910,13 +7980,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 @@ -7933,42 +8004,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 " @@ -7978,9 +8050,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 " @@ -7988,18 +8061,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." @@ -8014,13 +8087,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", @@ -8031,8 +8105,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", @@ -8043,19 +8117,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", @@ -8063,17 +8138,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, @@ -8085,15 +8161,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( @@ -8110,21 +8186,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", @@ -8133,34 +8210,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." @@ -8170,8 +8247,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." @@ -8181,7 +8258,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( @@ -8193,7 +8270,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(), @@ -8205,7 +8282,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 " @@ -8216,7 +8293,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 " @@ -8224,7 +8301,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." ) @@ -8243,13 +8320,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: @@ -8257,23 +8335,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( @@ -8293,23 +8369,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" @@ -8322,19 +8399,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 " @@ -8343,34 +8421,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", @@ -8469,23 +8547,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." ) @@ -8494,7 +8573,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." ) @@ -8516,11 +8595,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 " @@ -8541,7 +8621,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 " @@ -8551,14 +8631,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" @@ -8601,7 +8681,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 " @@ -8615,7 +8695,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." @@ -8626,7 +8706,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 " @@ -8634,8 +8714,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": @@ -8643,7 +8723,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." ) @@ -8654,8 +8734,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 " @@ -8683,19 +8763,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. " @@ -8708,7 +8789,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( @@ -8740,48 +8821,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: " @@ -8795,7 +8877,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." ) @@ -8809,8 +8892,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." ) @@ -8840,7 +8923,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] @@ -8887,7 +8970,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) @@ -8922,17 +9005,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)." @@ -8941,24 +9025,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. @@ -8972,14 +9056,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), " @@ -8987,7 +9071,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) " @@ -9003,12 +9087,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 @@ -9022,7 +9107,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", @@ -9059,10 +9144,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 " @@ -9074,28 +9159,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, @@ -9115,8 +9201,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" ) @@ -9124,18 +9210,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" ) @@ -9144,13 +9230,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." ) @@ -9159,7 +9245,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." ) @@ -9170,43 +9256,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 " @@ -9227,37 +9313,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." @@ -9265,25 +9349,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", @@ -9293,7 +9376,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(): @@ -9323,7 +9406,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) @@ -9623,15 +9707,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. @@ -9661,6 +9751,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) @@ -9673,12 +9764,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", @@ -9695,22 +9786,22 @@ 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 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 ( - 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()) @@ -9733,7 +9824,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 @@ -9742,17 +9839,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]: @@ -9763,26 +9861,28 @@ 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 # 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 @@ -9811,7 +9911,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 @@ -9821,15 +9921,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( @@ -9839,20 +9941,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 " @@ -9860,67 +9964,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. @@ -9933,41 +10037,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" + 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( - "--prompt-tokens-buckets", self.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" ) @@ -9983,23 +10085,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" ) @@ -10008,65 +10110,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}" @@ -10076,9 +10180,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) @@ -10116,7 +10220,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=[ @@ -10126,56 +10230,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): @@ -10184,20 +10288,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() @@ -10206,20 +10311,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", @@ -10286,12 +10391,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 @@ -10325,6 +10431,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): @@ -10334,7 +10441,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." ) @@ -10368,14 +10475,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 @@ -10431,8 +10539,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: @@ -10463,26 +10572,38 @@ 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: - 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") + 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: - """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 ) @@ -10642,6 +10763,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: @@ -10674,7 +10796,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}", @@ -10705,7 +10827,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/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/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/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/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/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index d1fd11a023fd..8f826768e8f3 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 @@ -4625,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() @@ -4637,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) @@ -4645,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/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/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/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/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/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/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/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 adc6036ae9fa..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 @@ -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)) @@ -188,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/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/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/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/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/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/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_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( 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/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/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/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/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() 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/managers/test_mm_process_config.py b/test/registered/unit/managers/test_mm_process_config.py index 04a9e82f187f..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): @@ -81,21 +85,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/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/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, ) 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/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 e21d6a462fe7..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 @@ -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): diff --git a/test/registered/unit/server_args/test_resolution_declarations.py b/test/registered/unit/server_args/test_resolution_declarations.py index 2925811613f4..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 @@ -233,6 +231,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 +348,69 @@ 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): + """`/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): @@ -535,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 @@ -560,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 @@ -582,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): @@ -635,24 +686,37 @@ 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): - mismatches = [] + def test_an_undeclared_field_still_holds_the_raw_input(self): + """Nothing writes a field behind the stash's back. + + 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. + """ + 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 - actual = getattr(server_args, field) - if actual != declared: - mismatches.append( - f"{shape} -> {field}: field={actual!r} 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, [], - "a declared field and its stash entry disagree, so something " - "assigned the field behind the declaration:\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): @@ -738,7 +802,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..b64504c7b7ac 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: @@ -761,31 +768,43 @@ 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( - getattr(bare, "_declarations_materialized", False), + getattr(bare, "_resolution_finished", False), "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): 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) @@ -847,9 +866,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_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/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 5c3e085b5951..fc2f4c4f5697 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, @@ -26,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 @@ -61,7 +63,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 +78,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 +116,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 +144,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 +162,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 +236,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 +247,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 +287,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 +304,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 +332,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 +345,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 +360,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 +375,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 +402,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 +431,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 +448,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 +466,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 +480,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 +517,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 +600,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 +633,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 +653,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 +717,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 +736,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 +925,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 +938,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 +973,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 +995,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 +1018,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 +1114,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 +1383,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 +1403,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 +1491,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 +1576,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 +1623,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 +1642,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 +1660,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 +1676,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 +1688,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 +1725,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 +1815,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 +1846,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 +1860,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 +1888,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 +1923,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 +1955,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 +1966,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 +2167,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 +2317,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) @@ -2237,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", @@ -2254,7 +2379,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=get_serving().grpc_port, ) self.assertEqual(handle, "handle") 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/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, ) 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_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index 6b840f97057a..b183374b63cf 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -53,6 +53,18 @@ # 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" + ), + ("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" @@ -68,6 +80,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" @@ -115,6 +134,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 " @@ -212,6 +251,27 @@ "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/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_model_overrides.py b/test/registered/unit/test_model_overrides.py index c7d748f6a623..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 @@ -280,17 +281,42 @@ 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): """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, @@ -565,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.""" @@ -616,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", @@ -643,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") @@ -994,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. @@ -1022,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 @@ -1045,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 @@ -1071,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 = [ @@ -1101,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 ( @@ -1142,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 @@ -1161,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 @@ -1171,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": @@ -1180,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_publish_precedes_bag_reads.py b/test/registered/unit/test_publish_precedes_bag_reads.py index 8429b0915da5..78ced6117755 100644 --- a/test/registered/unit/test_publish_precedes_bag_reads.py +++ b/test/registered/unit/test_publish_precedes_bag_reads.py @@ -78,15 +78,22 @@ "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", "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"), } ) 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..7532334e7294 --- /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. 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 +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() diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index b655102d14a1..fbcc472f1c2e 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", ) @@ -533,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( @@ -552,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_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 e7dfd144e4f4..fecea1ddb879 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() @@ -106,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_server_args_migration.py b/test/registered/unit/test_server_args_migration.py index cac5d7cc6340..67b2d9eb72e8 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,10 @@ def test_media_url_security_args(self): "32", ] ) + # The normalization is a declaration. 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: 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 diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index aa9af73ce56c..e8f91a5ada1d 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -134,80 +134,12 @@ # are step-12 exposure like any other pair. _PASSED = frozenset({"model_path", "device", "random_seed"}) -_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"), - ("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"), - ("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"), - ("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"), - ("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/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"), - ("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"), - ("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 @@ -220,35 +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 = { - ("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"), - ("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"), - ("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"), - ("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: @@ -631,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(