diff --git a/.claude/skills/sglang-runtime-context/SKILL.md b/.claude/skills/sglang-runtime-context/SKILL.md index 0eea90569..77e93c66f 100644 --- a/.claude/skills/sglang-runtime-context/SKILL.md +++ b/.claude/skills/sglang-runtime-context/SKILL.md @@ -139,9 +139,9 @@ bag to override at all. `Engine`s can share one process, bags are last-publish-wins across them") is **retracted** — owner ruling (2026-08-15): a process holds at most one live config at a time (concurrent multi-Engine is unsupported; sequential rebuild - stays legal, unit tests rely on it). What still reads the instance in those - files is pinned pair by pair in the exposure ratchet, each with its own - disposition; none of it is a boundary to imitate. What + stays legal, unit tests rely on it). Nothing in those files reads the instance + any more -- the exposure ratchet's pin set is empty, so the next such read is a + new entry that has to argue for itself. What genuinely stays per-instance is what differs per *worker* within one engine: `base_gpu_id` travels as a constructor argument (`MMEncoder(gpu_id=...)`; `BaseMultimodalProcessor._fast_image_processor_device` is the shape to copy). @@ -149,18 +149,17 @@ bag to override at all. supplied-instance contract; don't rewrite the parameter reads unless the field is runtime-mutated (see the elastic-EP `ep_size` case in `eplb/expert_location.py`) — **or the field is one that resolution fills in - and the callee runs in a process that has published.** That second case is - pinned debt, not a style question: the record is destined to carry the - user's raw input, so `server_args.page_size` inside a runner-owned - constructor will read the raw pre-resolution value instead of the effective - one. Debt means a decision, not automatically a bag read: pick where the + and the callee runs in a process that has published.** That second case is a + decision, not a style question: the record carries the user's raw input, so a + resolution-filled field read off it inside a runner-owned constructor answers + with the pre-resolution value instead of the effective one. Debt means a decision, not automatically a bag read: pick where the value should come from — usually the `get_*()` bag, sometimes a runner stamp or a constructor argument (the per-mode attention pair and the encode-server `gpu_id` above are dispositions of exactly this debt). The per-instance boundaries above are **not** exempt from this unless-clause (the multi-Engine exemption is retracted); each one gets its own disposition. `test_supplied_instance_exposure_ratchet.py` - pins the remaining set — three spellings of the read: `server_args.field`, + pins that set (empty today) — three spellings of the read: `server_args.field`, literal-name `getattr(server_args, "field", default)`, and the parked form (`self.x = server_args` in a method that takes the parameter, read as `self.x.field` anywhere in the class) — and fails on a new one, so the @@ -411,7 +410,8 @@ probes, swappable ACTIVE values. Not for config mirrors (read the bag leaf inste - Groups are typed dataclasses on `Flags` (`capture` / `moe` / `dp`): typo-safe writes, transactional test-only `override(**kw)` context manager. -- `flags.moe` is materialized by `initialize_moe_config(server_args)` at scheduler init; +- `flags.moe` is materialized by `initialize_moe_config()` at scheduler init (it + reads `exec.moe` / `spec` / `model`, and takes no record); accessors (`get_moe_a2a_backend` etc.) are thin shims with lazy defaults. The speculative contexts (`speculative_moe_backend_context`) swap the ACTIVE leaves around draft forwards. - `flags.dp` is materialized by `initialize_dp_attention`; `is_dp_attention_enabled()` is a diff --git a/python/sglang/benchmark/offline_throughput.py b/python/sglang/benchmark/offline_throughput.py index dd8037b38..37acaad6f 100644 --- a/python/sglang/benchmark/offline_throughput.py +++ b/python/sglang/benchmark/offline_throughput.py @@ -30,6 +30,7 @@ from sglang.benchmark.datasets import DatasetRow, get_dataset 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 _create_ray_engine_backend(server_args: ServerArgs): 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 d6d12a499..55df78a61 100644 --- a/python/sglang/benchmark/one_batch.py +++ b/python/sglang/benchmark/one_batch.py @@ -64,6 +64,7 @@ import numpy as np import torch import torch.distributed as dist +from sglang.srt.arg_groups.overrides import resolution_result, resolving_view from sglang.srt.configs.model_config import ModelConfig from sglang.srt.distributed.parallel_state import ( destroy_distributed_environment, @@ -79,7 +80,11 @@ from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.managers.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw from sglang.srt.mem_cache.base_prefix_cache import EvictParams -from sglang.srt.model_executor.cuda_graph_config import Phase, cuda_graph_fully_disabled +from sglang.srt.model_executor.cuda_graph_config import ( + CudaGraphConfig, + Phase, + cuda_graph_fully_disabled, +) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.runtime_context import get_parallel, get_schedule, publish @@ -297,44 +302,45 @@ class BenchArgs: def load_model(server_args, port_args, gpu_id, tp_rank): + cfg = resolving_view(server_args) suppress_other_loggers() rank_print = print if tp_rank == 0 else lambda *args, **kwargs: None - moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size) + moe_ep_rank = tp_rank // (cfg.tp_size // cfg.ep_size) model_config = ModelConfig.from_server_args(server_args) attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = ( compute_dp_attention_world_info( - server_args.enable_dp_attention, + cfg.enable_dp_attention, tp_rank, - server_args.tp_size, - server_args.dp_size, - server_args.attn_cp_size, + cfg.tp_size, + cfg.dp_size, + cfg.attn_cp_size, ) ) ps = ParallelState( tp_rank=tp_rank, - tp_size=server_args.tp_size, + tp_size=cfg.tp_size, pp_rank=0, pp_size=1, dp_rank=None, - dp_size=server_args.dp_size, + dp_size=cfg.dp_size, attn_tp_rank=attn_tp_rank, attn_tp_size=attn_tp_size, attn_cp_rank=0, - attn_cp_size=server_args.attn_cp_size, - attn_dcp_rank=tp_rank % server_args.dcp_size, - attn_dcp_size=server_args.dcp_size, + attn_cp_size=cfg.attn_cp_size, + attn_dcp_rank=tp_rank % cfg.dcp_size, + attn_dcp_size=cfg.dcp_size, attn_dp_rank=attn_dp_rank, attn_dp_size=attn_dp_size, moe_ep_rank=moe_ep_rank, - moe_ep_size=server_args.ep_size, + moe_ep_size=cfg.ep_size, moe_dp_rank=None, - moe_dp_size=server_args.moe_dp_size, + moe_dp_size=cfg.moe_dp_size, gpu_id=gpu_id, ) runner_kwargs = dict( model_config=model_config, - mem_fraction_static=server_args.mem_fraction_static, + mem_fraction_static=cfg.mem_fraction_static, gpu_id=gpu_id, ps=ps, nccl_port=port_args.nccl_port, @@ -350,20 +356,20 @@ def load_model(server_args, port_args, gpu_id, tp_rank): model_runner = MlxModelRunnerStub(**runner_kwargs) else: model_runner = ModelRunner(**runner_kwargs) - if server_args.is_startup_weight_load_overlap: + if cfg.is_startup_weight_load_overlap: model_runner.start_startup_weight_load() model_runner.alloc_memory_pool() model_runner.init_attention_backends() model_runner.init_cuda_graphs() - if server_args.is_startup_weight_load_overlap: + if cfg.is_startup_weight_load_overlap: model_runner.finalize_startup_weight_load() rank_print(f"max_total_num_tokens={model_runner.max_total_num_tokens}") tokenizer = get_tokenizer( - server_args.tokenizer_path, - tokenizer_mode=server_args.tokenizer_mode, - trust_remote_code=server_args.trust_remote_code, + cfg.tokenizer_path, + tokenizer_mode=cfg.tokenizer_mode, + trust_remote_code=cfg.trust_remote_code, ) - if server_args.tp_size > 1: + if cfg.tp_size > 1: dist.barrier() if _use_mlx: @@ -584,19 +590,20 @@ class _MlxBenchRunner: """Wraps MlxModelRunner for the MLX benchmark path.""" def __init__(self, model_runner, server_args): + cfg = resolving_view(server_args) from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner # Radix cache requires the scheduler's allocator/trie; disable in # standalone bench mode where no scheduler is present. init_kwargs = dict( - model_path=server_args.model_path, - trust_remote_code=server_args.trust_remote_code, + model_path=cfg.model_path, + trust_remote_code=cfg.trust_remote_code, disable_radix_cache=True, - mem_fraction_static=server_args.mem_fraction_static, - quantization=server_args.quantization, + mem_fraction_static=cfg.mem_fraction_static, + quantization=cfg.quantization, ) - if server_args.max_total_tokens is not None: - init_kwargs["pool_size"] = server_args.max_total_tokens + if cfg.max_total_tokens is not None: + init_kwargs["pool_size"] = cfg.max_total_tokens self.mlx_runner = MlxModelRunner(**init_kwargs) self.mlx_runner.init_cache_pools(req_to_token_pool=None) self.fake_torch_runner = model_runner @@ -883,17 +890,18 @@ def latency_test( gpu_id, tp_rank, ): + cfg = resolving_view(server_args) # `main` runs this inline for tp_size == 1 and spawns it per rank otherwise; # a spawned child arrives with nothing published. publish(server_args, role="scheduler") - initialize_moe_config(server_args) - initialize_fp8_gemm_config(server_args) - initialize_fp4_gemm_config(server_args) + initialize_moe_config() + initialize_fp8_gemm_config() + initialize_fp4_gemm_config() - # Set CPU affinity if get_bool_env_var("SGLANG_SET_CPU_AFFINITY"): + parallel = get_parallel().config set_gpu_proc_affinity( - server_args.pp_size, server_args.tp_size, server_args.nnodes, tp_rank + parallel.pp_size, parallel.tp_size, parallel.nnodes, tp_rank ) # Configure the logger @@ -988,22 +996,45 @@ def latency_test( for result in result_list: fout.write(json.dumps(result) + "\n") - if server_args.tp_size > 1: + if cfg.tp_size > 1: destroy_model_parallel() destroy_distributed_environment() def main(server_args, bench_args): + # The decode phase has to capture the batch sizes this run benchmarks, and + # the per-phase convenience knob loses to an explicit --cuda-graph-config + # JSON (resolution applies that last), so the size is merged into that JSON. + if getattr(server_args, "_declarations_materialized", False): + # A record the caller already resolved: nothing will parse a raw dict + # again, so the declaration has to be the finished typed config. + merged = resolution_result(server_args, "cuda_graph_config") + merged = ( + merged.to_dict() + if isinstance(merged, CudaGraphConfig) + else dict(merged or {}) + ) + decode = dict(merged.get(Phase.DECODE) or {}) + decode["max_bs"] = max(bench_args.batch_size) + merged[Phase.DECODE] = decode + graph_config = CudaGraphConfig.from_dict(merged) + else: + explicit = server_args.cuda_graph_config + if isinstance(explicit, CudaGraphConfig): + explicit = explicit.to_dict() + graph_config = dict(explicit or {}) + decode = dict(graph_config.get(Phase.DECODE) or {}) + decode["max_bs"] = max(bench_args.batch_size) + graph_config[Phase.DECODE] = decode + server_args = server_args.replace_resolved( + "benchmark.one_batch", cuda_graph_config=graph_config + ) server_args.resolve_once() - - # The legacy cuda_graph_max_bs_decode field does not propagate; set the - # decode phase. - if server_args.cuda_graph_config is not None: - server_args.cuda_graph_config[Phase.DECODE].max_bs = max(bench_args.batch_size) + cfg = resolving_view(server_args) _set_envs_and_config(server_args) - if server_args.model_path: + if cfg.model_path: if bench_args.correctness_test: work_func = correctness_test else: diff --git a/python/sglang/benchmark/one_batch_server.py b/python/sglang/benchmark/one_batch_server.py index 76a4a0314..3ea5f80b9 100644 --- a/python/sglang/benchmark/one_batch_server.py +++ b/python/sglang/benchmark/one_batch_server.py @@ -33,6 +33,7 @@ from sglang.benchmark.datasets import get_dataset 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 ab5ac06f3..bb07cc5d0 100644 --- a/python/sglang/compile_deep_gemm.py +++ b/python/sglang/compile_deep_gemm.py @@ -17,13 +17,19 @@ import time 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 0e74efd52..e0e3a4c82 100644 --- a/python/sglang/lang/backend/runtime_endpoint.py +++ b/python/sglang/lang/backend/runtime_endpoint.py @@ -456,13 +456,15 @@ class Runtime: 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 0ee16c0ca..1ddbb4e68 100644 --- a/python/sglang/launch_server.py +++ b/python/sglang/launch_server.py @@ -5,6 +5,7 @@ import os import sys import warnings +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.plugins import load_plugins from sglang.srt.server_args import prepare_server_args from sglang.srt.utils import kill_process_tree @@ -18,10 +19,11 @@ def run_server(server_args): # The flags dispatched on below are decided by resolution (`--grpc-mode` # folds into `smg_grpc_mode`), and `prepare_server_args` returns raw input. server_args.resolve_once() + cfg = resolving_view(server_args) - if server_args.encoder_only: + if cfg.encoder_only: # For encoder disaggregation - if server_args.smg_grpc_mode or server_args.grpc_mode: + if cfg.smg_grpc_mode or cfg.grpc_mode: from sglang.srt.disaggregation.encoder.grpc_server import ( serve_grpc_encoder, ) @@ -31,7 +33,7 @@ def run_server(server_args): from sglang.srt.disaggregation.encoder.http_server import launch_server launch_server(server_args) - elif server_args.smg_grpc_mode: + elif cfg.smg_grpc_mode: # Legacy SMG gRPC server (--smg-grpc-mode, or the deprecated --grpc-mode # which __post_init__ folds into smg_grpc_mode). The native Rust gRPC # server is a separate path, enabled by --grpc-port, that starts @@ -39,7 +41,7 @@ def run_server(server_args): from sglang.srt.entrypoints.grpc_server import serve_grpc asyncio.run(serve_grpc(server_args)) - elif server_args.use_ray: + elif cfg.use_ray: # Ray mode: HTTP mode with Ray backend. try: from sglang.srt.ray.http_server import launch_server diff --git a/python/sglang/srt/configs/embedding_model_spec.py b/python/sglang/srt/configs/embedding_model_spec.py index c529e1999..6bcf5a326 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) + getattr(config, "prefill_only_disable_kv_cache", False) ), - "radix_cache_disabled": bool( - getattr(server_args, "disable_radix_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/disaggregation/encoder/grpc_server.py b/python/sglang/srt/disaggregation/encoder/grpc_server.py index 91024e859..58004e0eb 100644 --- a/python/sglang/srt/disaggregation/encoder/grpc_server.py +++ b/python/sglang/srt/disaggregation/encoder/grpc_server.py @@ -24,7 +24,12 @@ from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc from sglang.srt.disaggregation.encoder.server import MMEncoder, launch_encoder from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle from sglang.srt.managers.schedule_batch import Modality -from sglang.srt.runtime_context import get_disagg, publish +from sglang.srt.runtime_context import ( + get_disagg, + get_parallel, + get_serving, + publish, +) from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import random_uuid from sglang.srt.utils.network import NetworkAddress, get_zmq_socket @@ -212,11 +217,11 @@ async def serve_grpc_encoder(server_args: ServerArgs): dist_init_method = na.to_tcp() else: dist_init_method = NetworkAddress( - server_args.host or "127.0.0.1", port_args.nccl_port + get_serving().host or "127.0.0.1", port_args.nccl_port ).to_tcp() send_sockets: List[zmq.Socket] = [] - for rank in range(1, server_args.tp_size): + for rank in range(1, get_parallel().config.tp_size): schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}" send_sockets.append( get_zmq_socket(zmq_ctx, zmq.PUSH, schedule_path, bind=False) @@ -254,7 +259,9 @@ async def serve_grpc_encoder(server_args: ServerArgs): ) reflection.enable_server_reflection(SERVICE_NAMES, server) - listen_addr = NetworkAddress(server_args.host, server_args.port).to_host_port_str() + listen_addr = NetworkAddress( + get_serving().host, get_serving().port + ).to_host_port_str() server.add_insecure_port(listen_addr) await server.start() diff --git a/python/sglang/srt/disaggregation/encoder/http_server.py b/python/sglang/srt/disaggregation/encoder/http_server.py index e204048c6..4970e3d49 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 _register_encoder_url_with_bootstrap(server_args: ServerArgs): 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 6d8e0561a..bd6d3f06e 100644 --- a/python/sglang/srt/disaggregation/encoder/preprocessor.py +++ b/python/sglang/srt/disaggregation/encoder/preprocessor.py @@ -127,7 +127,7 @@ class EncoderPreprocessor: 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 @@ class EncoderPreprocessor: 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 @@ class EncoderPreprocessor: 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 @@ class EncoderPreprocessor: 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 @@ class EncoderPreprocessor: 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 50b75ccae..ace0bcd8d 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.multimodal_processor import get_mm_processor, import_pr 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 @@ class MMReceiverBase(ABC): # 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 @@ class MMReceiverBase(ABC): _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 433f0e9fc..a4d88cd9e 100644 --- a/python/sglang/srt/disaggregation/encoder/runtime.py +++ b/python/sglang/srt/disaggregation/encoder/runtime.py @@ -1572,10 +1572,10 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher: HTTP uses this entry point today. gRPC can reuse it later without importing HTTP application state or Uvicorn. """ - if get_parallel().config.dp_size <= 1 or server_args.tp_size != 1: + if get_parallel().config.dp_size <= 1 or get_parallel().config.tp_size != 1: raise ValueError( "Encoder DP mode requires --dp-size > 1 and --tp-size 1; got " - f"dp_size={get_parallel().config.dp_size}, tp_size={server_args.tp_size}." + f"dp_size={get_parallel().config.dp_size}, tp_size={get_parallel().config.tp_size}." ) dp_size = get_parallel().config.dp_size logger.info(f"Launching encoder in DP mode: dp_size={dp_size}") diff --git a/python/sglang/srt/disaggregation/encoder/server.py b/python/sglang/srt/disaggregation/encoder/server.py index cef8e884b..34b1386db 100644 --- a/python/sglang/srt/disaggregation/encoder/server.py +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -63,6 +63,7 @@ from sglang.srt.runtime_context import ( get_exec, get_mm, get_model, + get_parallel, publish, ) from sglang.srt.server_args import ServerArgs @@ -450,7 +451,7 @@ class MMEncoder: this instance's value, not a config change, so it travels as an argument.""" assert_published(server_args, role="encoder") - logger.info(f"init MMEncoder {rank}/{server_args.tp_size}") + logger.info(f"init MMEncoder {rank}/{get_parallel().config.tp_size}") self.server_args = server_args configure_media_url_security( get_mm().allowed_media_domains, @@ -470,7 +471,7 @@ class MMEncoder: self.load_config = LoadConfig( load_format=get_model().load_format, download_dir=server_args.download_dir, - model_loader_extra_config=server_args.model_loader_extra_config, + model_loader_extra_config=get_model().model_loader_extra_config, remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip, remote_instance_weight_loader_seed_instance_service_port=server_args.remote_instance_weight_loader_seed_instance_service_port, remote_instance_weight_loader_send_weights_group_ports=server_args.remote_instance_weight_loader_send_weights_group_ports, @@ -491,12 +492,14 @@ class MMEncoder: init_distributed_environment( backend=get_default_distributed_backend(self.device), - world_size=server_args.tp_size, + world_size=get_parallel().config.tp_size, rank=rank, distributed_init_method=dist_init_method, local_rank=rank, ) - initialize_model_parallel(tensor_model_parallel_size=server_args.tp_size) + initialize_model_parallel( + tensor_model_parallel_size=get_parallel().config.tp_size + ) initialize_dp_attention(server_args, self.model_config) self.model = load_model( @@ -554,7 +557,7 @@ class MMEncoder: ) self.mm_global_cache = EmbeddingCacheController( rank, - server_args.tp_size, + get_parallel().config.tp_size, embedding_store=embedding_store, hidden_dims=self._embedding_dims, tp_group=get_tp_group().cpu_group, @@ -1032,7 +1035,7 @@ class MMEncoder: ) def _broadcast_global_cache_mask(self, mask_tensor: torch.Tensor): - if self.server_args.tp_size > 1: + if get_parallel().config.tp_size > 1: torch.distributed.broadcast( mask_tensor, src=0, diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index fb0c3460c..f72fb25fc 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -46,7 +46,7 @@ from sglang.srt.disaggregation.utils import ( 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 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): ): 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 623b9aa63..f5a4c4387 100644 --- a/python/sglang/srt/distributed/bootstrap.py +++ b/python/sglang/srt/distributed/bootstrap.py @@ -31,6 +31,7 @@ from sglang.srt.platforms import current_platform from sglang.srt.runtime_context import ( get_exec, get_parallel, + get_serving, ) from sglang.srt.server_args import ServerArgs from sglang.srt.utils import ( @@ -187,7 +188,7 @@ def _resolve_dist_init_method(*, server_args: ServerArgs, dist_port: int) -> str dist_init_method = na.to_tcp() else: dist_init_method = NetworkAddress( - server_args.host or "127.0.0.1", dist_port + get_serving().host or "127.0.0.1", dist_port ).to_tcp() return dist_init_method diff --git a/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py b/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py index 74e82d73a..0e7c6a7a0 100644 --- a/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py +++ b/python/sglang/srt/distributed/device_communicators/triton_symm_mem_ag.py @@ -423,19 +423,19 @@ def recommended_max_tokens(include_prefill: bool, floor: int = 0) -> int: NCCL. Covers the spec-decode batch plus, if ``include_prefill``, a prefill chunk. Returns ``floor`` if server args are unavailable.""" try: - from sglang.srt.runtime_context import get_server_args + from sglang.srt.runtime_context import get_schedule, get_spec - sa = get_server_args() + def g(value) -> int: + return value if isinstance(value, int) and value > 0 else 0 - def g(name: str) -> int: - v = getattr(sa, name, 0) - return v if isinstance(v, int) and v > 0 else 0 - - tokens = g("max_running_requests") * max( - g("speculative_num_draft_tokens"), g("speculative_eagle_topk"), 1 + schedule, spec = get_schedule(), get_spec() + tokens = g(schedule.max_running_requests) * max( + g(spec.speculative_num_draft_tokens), g(spec.speculative_eagle_topk), 1 ) if include_prefill: - tokens = max(tokens, g("chunked_prefill_size"), g("max_prefill_tokens")) + tokens = max( + tokens, g(schedule.chunked_prefill_size), g(schedule.max_prefill_tokens) + ) return max(tokens, floor) except Exception: return floor diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index e4e625c60..670831be1 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -48,6 +48,7 @@ import torch 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_detection import resolve_auto_parsers from sglang.srt.parser.template_manager import TemplateManager from sglang.srt.plugins import load_plugins from sglang.srt.runtime_context import ( + get_disagg, get_exec, get_model, get_parallel, @@ -309,9 +311,9 @@ class Engine(EngineScoreMixin, EngineBase): 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 @@ class Engine(EngineScoreMixin, EngineBase): # Needs a tokenizer and a chat template, so it cannot live in the # pipeline; after the plugins, which may register the parser detected. - if ( - server_args.reasoning_parser == "auto" - or server_args.tool_call_parser == "auto" - ): + parsers = resolving_view(server_args) + if parsers.reasoning_parser == "auto" or parsers.tool_call_parser == "auto": resolve_auto_parsers(server_args) # This publish replaces whatever was published before it, so the @@ -1627,27 +1627,28 @@ class Engine(EngineScoreMixin, EngineBase): 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 _set_envs_and_config(server_args: ServerArgs): 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 _set_envs_and_config(server_args: ServerArgs): 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 9349d3f92..4adf1d737 100644 --- a/python/sglang/srt/entrypoints/grpc_bridge.py +++ b/python/sglang/srt/entrypoints/grpc_bridge.py @@ -14,6 +14,7 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional from pydantic import ValidationError +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.configs.embedding_model_spec import resolved_embedding_plan from sglang.srt.runtime_context import ( get_lora, @@ -416,7 +417,7 @@ class RuntimeHandle: if embedding_model_spec is not None: result["embedding"] = resolved_embedding_plan( embedding_model_spec, - server_args=self.server_args, + config=resolving_view(self.server_args), model_config=model_config, ) return json.dumps(result, default=str) diff --git a/python/sglang/srt/entrypoints/grpc_server.py b/python/sglang/srt/entrypoints/grpc_server.py index c3e22762c..0c2a8c559 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 serve_grpc(server_args, model_info=None): ) 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 serve_grpc(server_args, model_info=None): ) 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 13817d468..d56aa71e3 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -63,6 +63,7 @@ from fastapi.middleware.cors import CORSMiddleware 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 @@ -284,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() @@ -488,6 +489,7 @@ from sglang.srt.runtime_context import ( get_exec, get_lora, get_model, + get_observability, get_parallel, get_serving, publish, @@ -771,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 @@ -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 @@ def _setup_and_run_http_server( # 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 @@ 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 ), @@ -2707,10 +2713,11 @@ def _setup_and_run_http_server( 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 d7107a975..754950af2 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 93f77bd3c..3591fe79f 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/expert_distribution.py b/python/sglang/srt/eplb/expert_distribution.py index ec6133956..458ecfc75 100644 --- a/python/sglang/srt/eplb/expert_distribution.py +++ b/python/sglang/srt/eplb/expert_distribution.py @@ -764,7 +764,7 @@ class _UtilizationRateAccumulatorMixin(_Accumulator): single_pass_global_physical_count, num_gpu=self._expert_location_metadata.ep_size, ) - gpu_physical_count = gpu_physical_count.to(self._server_args.device) + gpu_physical_count = gpu_physical_count.to(get_device_namespace().device) torch.distributed.reduce( gpu_physical_count, dst=0, op=torch.distributed.ReduceOp.SUM ) @@ -898,9 +898,9 @@ class _StatAccumulator(_UtilizationRateAccumulatorMixin): # 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/hardware_backend/xpu/graph_runner/xpu_graph_runner.py b/python/sglang/srt/hardware_backend/xpu/graph_runner/xpu_graph_runner.py index b8c7fa592..bcc8b4e62 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 @@ import torch 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 @@ class XPUGraphRunner(DecodeCudaGraphRunner): 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 4c9c8ea46..83c22746e 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 @@ class DualChunkFlashAttentionBackend(AttentionBackend): 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/minimax_sparse_backend.py b/python/sglang/srt/layers/attention/minimax_sparse_backend.py index def19cbc6..49e617e39 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 @@ from typing import TYPE_CHECKING, Optional, Tuple 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 @@ class MiniMaxSparseAttnBackend(AttentionBackend): self.max_context_len = int(runner.model_config.context_len) # Per-forward cache for the native decode block table (rebuilt each forward). self._native_decode_bt: dict = {} - self.fp8_attn_gemm = m3_fp8_attn_gemm_enabled(runner.server_args) + self.fp8_attn_gemm = m3_fp8_attn_gemm_enabled( + resolving_view(runner.server_args) + ) if self.fp8_attn_gemm: assert self.kv_pool.main_pool.dtype == torch.float8_e4m3fn, ( "fp8 attn-GEMM mode requires an fp8_e4m3fn main KV pool, got " diff --git a/python/sglang/srt/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index c57e33e97..225fc4fdb 100644 --- a/python/sglang/srt/layers/cp/base.py +++ b/python/sglang/srt/layers/cp/base.py @@ -33,7 +33,6 @@ from sglang.srt.runtime_context import get_parallel if TYPE_CHECKING: from sglang.srt.model_executor.forward_batch_info import ForwardBatch - from sglang.srt.server_args import ServerArgs class ContextParallelStrategyKind(IntEnum): @@ -237,23 +236,27 @@ def _is_dsa_active() -> bool: _STRATEGY: Optional[ContextParallelStrategy] = None -def init_cp_strategy(server_args: ServerArgs) -> None: - """Bind the configured CP strategy for this process.""" - from sglang.srt.arg_groups.overrides import resolving_view +def init_cp_strategy( + *, enable_prefill_cp: bool, cp_size: int, cp_strategy: str +) -> None: + """Bind the CP strategy for this process. - cfg = resolving_view(server_args) + Takes the three values: resolution calls this from inside `__post_init__`, + where the bags do not exist yet, and `get_cp_strategy` calls it lazily in a + worker, which reads them off the published bags. Each caller reads from the + source it has. + """ global _STRATEGY - if not cfg.enable_prefill_cp: + if not enable_prefill_cp: _STRATEGY = None return - cp_size = cfg.attn_cp_size if cp_size <= 1: _STRATEGY = None return - kind = ContextParallelStrategyKind.from_string(cfg.cp_strategy) + kind = ContextParallelStrategyKind.from_string(cp_strategy) if kind == ContextParallelStrategyKind.ZIGZAG: from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy @@ -264,8 +267,7 @@ def init_cp_strategy(server_args: ServerArgs) -> None: _STRATEGY = InterleaveCPStrategy(cp_size=cp_size) else: raise ValueError( - f"Unsupported cp_strategy kind {kind} for " - f"cp_strategy={cfg.cp_strategy!r}" + f"Unsupported cp_strategy kind {kind} for cp_strategy={cp_strategy!r}" ) @@ -280,14 +282,16 @@ def get_cp_strategy() -> Optional[ContextParallelStrategy]: global _STRATEGY if _STRATEGY is None: - from sglang.srt.runtime_context import get_server_args - try: - server_args = get_server_args() + parallel = get_parallel().config except ValueError: return None - if server_args is not None and get_parallel().config.enable_prefill_cp: - init_cp_strategy(server_args) + if parallel.enable_prefill_cp: + init_cp_strategy( + enable_prefill_cp=True, + cp_size=parallel.attn_cp_size, + cp_strategy=parallel.cp_strategy, + ) return _STRATEGY diff --git a/python/sglang/srt/layers/cp/utils.py b/python/sglang/srt/layers/cp/utils.py index b397c4fed..be2ff1757 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/flashinfer_comm_fusion.py b/python/sglang/srt/layers/flashinfer_comm_fusion.py index cd399649c..8c1aa58a7 100644 --- a/python/sglang/srt/layers/flashinfer_comm_fusion.py +++ b/python/sglang/srt/layers/flashinfer_comm_fusion.py @@ -13,7 +13,7 @@ from sglang.srt.distributed import ( 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 7f4e51024..11f8a4b15 100644 --- a/python/sglang/srt/layers/moe/utils.py +++ b/python/sglang/srt/layers/moe/utils.py @@ -4,7 +4,6 @@ import logging 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.environ import envs from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) -from sglang.srt.runtime_context import get_exec, get_flags, get_forward, get_parallel +from sglang.srt.runtime_context import ( + get_exec, + get_flags, + get_forward, + get_model, + get_parallel, + get_spec, +) from sglang.srt.utils import is_cuda, is_npu _is_npu = is_npu() -if TYPE_CHECKING: - from sglang.srt.server_args import ServerArgs - from sglang.srt.runtime_context import get_server_args from sglang.srt.utils.common import log_info_on_rank0 @@ -308,37 +311,47 @@ def get_ascend_dispatcher_output_dtype(dispatcher): return DispatcherOutputDtype.BF16 -def initialize_moe_config(server_args: ServerArgs): +def initialize_moe_config(): + """Seed the MoE runtime flags from the published configuration. + + Reads the bags: `moe_a2a_backend` and its siblings are resolution's + answers, and the record carries them only while declarations materialize + onto it. Called once per process after publish + (scheduler init, the benchmark work functions). + """ + exec_moe = get_exec().moe + overlap = get_exec().overlap + spec = get_spec() moe = get_flags().moe - moe.a2a_backend = MoeA2ABackend(server_args.moe_a2a_backend) - moe.runner_backend = MoeRunnerBackend(server_args.moe_runner_backend) + moe.a2a_backend = MoeA2ABackend(exec_moe.moe_a2a_backend) + moe.runner_backend = MoeRunnerBackend(exec_moe.moe_runner_backend) moe.speculative_runner_backend = ( - MoeRunnerBackend(server_args.speculative_moe_runner_backend) - if server_args.speculative_moe_runner_backend is not None + MoeRunnerBackend(spec.speculative_moe_runner_backend) + if spec.speculative_moe_runner_backend is not None else moe.runner_backend ) moe.speculative_a2a_backend = ( - MoeA2ABackend(server_args.speculative_moe_a2a_backend) - if server_args.speculative_moe_a2a_backend is not None + MoeA2ABackend(spec.speculative_moe_a2a_backend) + if spec.speculative_moe_a2a_backend is not None else moe.a2a_backend ) - moe.deepep_mode = DeepEPMode(server_args.deepep_mode) - moe.deepep_config = server_args.deepep_config or "" - moe.tbo_enabled = server_args.enable_two_batch_overlap - moe.sbo_enabled = server_args.enable_single_batch_overlap + moe.deepep_mode = DeepEPMode(exec_moe.deepep_mode) + moe.deepep_config = exec_moe.deepep_config or "" + moe.tbo_enabled = overlap.enable_two_batch_overlap + moe.sbo_enabled = overlap.enable_single_batch_overlap if moe.sbo_enabled and is_cuda(): if torch.cuda.get_device_capability()[0] == 9: raise ValueError( "SBO (single batch overlap) is not supported on SM90 GPUs with latest sgl-deep-gemm wheel. Please try removing --enable-single-batch-overlap argument." ) - moe.tbo_token_distribution_threshold = server_args.tbo_token_distribution_threshold - moe.disable_fp4_allgather = server_args.disable_flashinfer_cutlass_moe_fp4_allgather - moe.quantization = server_args.quantization + moe.tbo_token_distribution_threshold = overlap.tbo_token_distribution_threshold + moe.disable_fp4_allgather = exec_moe.disable_flashinfer_cutlass_moe_fp4_allgather + moe.quantization = get_model().quantization # Seeded with the user's intent; each model's gate refines the ACTIVE # value for its own build (install_shared_experts_fusion_decision). - moe.disable_shared_experts_fusion = server_args.disable_shared_experts_fusion + moe.disable_shared_experts_fusion = exec_moe.disable_shared_experts_fusion moe.speculative_disable_shared_experts_fusion = ( - server_args.disable_shared_experts_fusion + exec_moe.disable_shared_experts_fusion ) diff --git a/python/sglang/srt/layers/quantization/fp4_utils.py b/python/sglang/srt/layers/quantization/fp4_utils.py index 1b5858d71..52588dbbf 100644 --- a/python/sglang/srt/layers/quantization/fp4_utils.py +++ b/python/sglang/srt/layers/quantization/fp4_utils.py @@ -2,10 +2,11 @@ from __future__ import annotations 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.common import ( ) 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 @@ class Fp4GemmRunnerBackend(Enum): 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 f71ba4784..267bf3b97 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -3,23 +3,10 @@ from __future__ import annotations 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 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( 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 @@ from sglang.srt.utils import ( is_xpu, offloader, ) +from sglang.srt.utils.common import torch_release from sglang.srt.utils.custom_op import register_custom_op logger = logging.getLogger(__name__) @@ -799,11 +793,11 @@ def _dispatch_auto_backend() -> Callable: return triton_w8a8_block_fp8_linear -def initialize_fp8_gemm_config(server_args: ServerArgs) -> None: +def initialize_fp8_gemm_config() -> None: """Initialize FP8 GEMM configuration.""" global FP8_GEMM_RUNNER_BACKEND - backend = server_args.fp8_gemm_runner_backend + backend = get_exec().kernel.fp8_gemm_runner_backend if backend == "auto" and is_sm120_supported(): backend = "cutlass" diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index 13d65d612..485e33cd9 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.io_struct import ( ) from sglang.srt.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread -from sglang.srt.runtime_context import 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 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): 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 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): 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/rust_server.py b/python/sglang/srt/managers/rust_server.py index 79c1a48aa..4251f1525 100644 --- a/python/sglang/srt/managers/rust_server.py +++ b/python/sglang/srt/managers/rust_server.py @@ -20,6 +20,7 @@ from typing import TYPE_CHECKING, Any, Dict, FrozenSet, List, Optional, Tuple 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.managers.utils import ( ) from sglang.srt.runtime_context import ( get_mm, + get_observability, + get_parallel, get_serving, ) from sglang.srt.utils.flatten import ( @@ -272,7 +275,7 @@ class NativeMmHost: ) 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 @@ class 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 @@ class RustServer: 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 @@ class RustServer: 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 24b2689c9..fd3a27bf5 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -903,11 +903,11 @@ class Scheduler( "moe_topk", ) if any(hasattr(config_to_check, attr) for attr in moe_topk_attrs): - initialize_moe_config(self.server_args) + initialize_moe_config() # Initialize GEMM-related configuration for FP8 and FP4 backends. - initialize_fp8_gemm_config(self.server_args) - initialize_fp4_gemm_config(self.server_args) + initialize_fp8_gemm_config() + initialize_fp4_gemm_config() initialize_bf16_gemm_config(self.server_args) # This must be called after initialize_moe_config diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index 0860a5e9f..41923444b 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -179,8 +179,8 @@ class TokenizerControlMixin: ) 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 ee97103c8..6ec332baa 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -3591,7 +3591,7 @@ def get_processor_wrapper(server_args): tokenizer_mode=get_serving().tokenizer_mode, trust_remote_code=get_model().trust_remote_code, revision=get_model().revision, - image_processor_backend=resolve_image_processor_backend(server_args), + image_processor_backend=resolve_image_processor_backend(get_mm()), tokenizer_backend=get_serving().tokenizer_backend, model_name=get_model().model_path, ) diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index c91fe88af..343571225 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -86,7 +86,7 @@ class HiRadixCache(RadixCache): 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 e29de3a01..2b38326b0 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 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) -def _get_allocator_type(server_args: ServerArgs) -> str: - return get_allocator_type(server_args) +def _get_allocator_type() -> str: + return get_allocator_type() def _evict_swa_for_device_alloc(cache: UnifiedRadixCache, required_size: int) -> None: @@ -126,7 +126,7 @@ def build_kv_host_pool( get_memory().hicache_size if host_size is None else host_size, page_size, get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), pool_label=pool_label, **kwargs, ) @@ -540,7 +540,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=swa_num_host_pages, slot_page_size=kvcache.swa_page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator entries.append( @@ -567,7 +567,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=num_host_pages, slot_page_size=page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) c4_indexer_host_pool = DeepSeekV4PagedHostPool( pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER), @@ -579,7 +579,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=num_host_pages, slot_page_size=page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) entries.extend( [ @@ -610,7 +610,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=swa_num_host_pages, swa_page_size=kvcache.swa_page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) c4_indexer_state_host_pool = DeepSeekV4StateHostPool( pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER_STATE), @@ -621,7 +621,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=swa_num_host_pages, swa_page_size=kvcache.swa_page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) entries.extend( [ @@ -653,7 +653,7 @@ def build_deepseek_v4_hicache_stack( num_host_pages=num_host_pages, slot_page_size=page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) # C128 state pool is intentionally not registered with hicache. # page_size=256 % 128 == 0, so state pool is not consumed on load. @@ -739,7 +739,7 @@ def build_hybrid_mamba_stack( mamba_pool, get_memory().hicache_ratio, mamba_host_size, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), layout=get_memory().hicache_mem_layout, ) entries = [ @@ -1038,7 +1038,7 @@ def build_full_draft_pools( host_to_device_ratio=host_pool_group.logical_size / pool.size, page_size=controller.page_size, layout=get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), pool_label="draft", ) draft_layer_mapping = {i: i for i in range(pool.layer_num)} @@ -1064,7 +1064,7 @@ def build_full_draft_pools( pool, draft_host_pool, get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) specs.append( SidecarPoolSpec( @@ -1111,7 +1111,7 @@ def build_swa_draft_pools( num_host_pages=target_swa_host_pool.num_host_pages, slot_page_size=draft_swa_pool.page_size, layout=target_swa_host_pool.layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ) else: host_pool = _build_mha_mla_host_pool( @@ -1119,7 +1119,7 @@ def build_swa_draft_pools( host_to_device_ratio=target_swa_host_pool.size / draft_swa_pool.size, page_size=target_swa_host_pool.page_size, layout=target_swa_host_pool.layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), pool_label="draft_swa", ) @@ -1513,7 +1513,7 @@ class _DsaStrategy(StackStrategy): full_kv_pool, kv_host_pool, get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ), prefetch_threshold=prefetch_threshold, model_name=model_name, @@ -1951,7 +1951,7 @@ def attach_hybrid_dsa_pool_to_hiradix_cache( kv, kv_host_pool, get_memory().hicache_mem_layout, - allocator_type=_get_allocator_type(server_args), + allocator_type=_get_allocator_type(), ), model_name=get_serving().served_model_name, storage_backend_extra_config=extra_config, diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index ed0215bff..75821be4d 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 @@ from typing import TYPE_CHECKING, Any, Optional import msgspec import torch +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.configs.hybrid_arch import ( hybrid_gdn_config, kimi_linear_config, @@ -1262,7 +1263,7 @@ class KVCacheConfigurator: sparse_layer_ids=sparse_layer_ids, disable_value_sparse_layer_ids=disable_value_sparse_layer_ids, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, ) @@ -1533,7 +1534,7 @@ class KVCacheConfigurator: # with the widening-dequant contract. index_dtype=( self.kv_cache_dtype - if m3_fp8_attn_gemm_enabled(self.server_args) + if m3_fp8_attn_gemm_enabled(resolving_view(self.server_args)) else self.model_dtype ), head_num=self.model_config.get_num_kv_heads( diff --git a/python/sglang/srt/mem_cache/pool_host/common.py b/python/sglang/srt/mem_cache/pool_host/common.py index 770135654..d781c2095 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 4f9bd1240..1a7429e84 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 @@ from sglang.srt.observability.metrics_collector import ( 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 @@ class UnifiedRadixCache(BasePrefixCache): 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 @@ class UnifiedRadixCache(BasePrefixCache): 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 @@ class UnifiedRadixCache(BasePrefixCache): # 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 @@ class UnifiedRadixCache(BasePrefixCache): 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 7a1b716bf..6d7843dc0 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -609,7 +609,7 @@ class CPUGraphRunner: 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/mindspore_runner.py b/python/sglang/srt/model_executor/mindspore_runner.py index 4cdcaed50..6d3a4d06a 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._c_expression import GroupOptions 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 f7947fb61..d4d6e3e54 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1202,7 +1202,6 @@ class ModelRunner: # Pre-expand RoPE cache before CUDA Graph capture reserve_rope_cache_for_long_sequences( self.model, - self.server_args, self.model_config, logger, ) diff --git a/python/sglang/srt/model_executor/model_runner_components/moe_ep_setup.py b/python/sglang/srt/model_executor/model_runner_components/moe_ep_setup.py index 27c9951cc..25cdea784 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.eplb.lplb_solver import ( ) from sglang.srt.layers.moe.hash_topk import HashTopK from sglang.srt.layers.moe.topk import TopK +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import get_bool_env_var, is_hip, log_info_on_rank0 if TYPE_CHECKING: @@ -54,7 +55,7 @@ def prepare_moe_topk( # Redundant experts therefore need to be included in the per-rank # expert count used for Waterfill's shared-expert slot remapping. num_physical_routed_experts = ( - num_routed_experts + server_args.ep_num_redundant_experts + num_routed_experts + get_exec().moe.ep_num_redundant_experts ) if isinstance(module, TopK): routed_scaling_factor = module.topk_config.routed_scaling_factor diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 1ebc26310..180c1795e 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 @@ class BaseRunner(ABC): self.model_runner = model_runner self.device = model_runner.device self.device_module = torch.get_device_module(self.device) - self.tp_size = model_runner.server_args.tp_size + self.tp_size = get_parallel().config.tp_size # elastic-EP scale-up rewrites dp_size on the published config self.dp_size = get_parallel().config.dp_size self.pp_size = get_parallel().config.pp_size diff --git a/python/sglang/srt/model_executor/runner_backend/utils.py b/python/sglang/srt/model_executor/runner_backend/utils.py index a07ec247a..63c98b45c 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/ray/data_parallel_controller.py b/python/sglang/srt/ray/data_parallel_controller.py index 4bbd15f68..89f86453b 100644 --- a/python/sglang/srt/ray/data_parallel_controller.py +++ b/python/sglang/srt/ray/data_parallel_controller.py @@ -88,7 +88,7 @@ class RayDataParallelController(DataParallelController): 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 @@ class RayDataParallelController(DataParallelController): 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 @@ class RayDataParallelController(DataParallelController): 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 @@ class RayDataParallelController(DataParallelController): 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 @@ class RayDataParallelController(DataParallelController): 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 @@ class RayDataParallelController(DataParallelController): 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 @@ class RayDataParallelController(DataParallelController): _, _, 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 3dbf3e57b..7be4b194a 100644 --- a/python/sglang/srt/ray/engine.py +++ b/python/sglang/srt/ray/engine.py @@ -105,18 +105,17 @@ def _get_bundle_node_ip(placement_group: PlacementGroup, bundle_idx: int) -> str ) -def _compute_world_size(server_args: ServerArgs) -> int: +def _compute_world_size() -> int: """Compute world_size (total number of scheduler actors/GPUs needed). Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size. + Reads the published parallel leaves: the driver is sizing the actors that + will hold the process groups, so there is nothing live to ask. """ - if get_parallel().config.enable_dp_attention: - return server_args.tp_size * get_parallel().config.pp_size - return ( - get_parallel().config.dp_size - * server_args.tp_size - * get_parallel().config.pp_size - ) + parallel = get_parallel().config + if parallel.enable_dp_attention: + return parallel.tp_size * parallel.pp_size + return parallel.dp_size * parallel.tp_size * parallel.pp_size def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]: @@ -274,16 +273,13 @@ class RayEngine(Engine): placement_group as create_placement_group, ) - if get_parallel().config.enable_dp_attention: - total_gpus = server_args.tp_size * get_parallel().config.pp_size + parallel = get_parallel().config + if parallel.enable_dp_attention: + total_gpus = parallel.tp_size * parallel.pp_size else: - total_gpus = ( - get_parallel().config.dp_size - * server_args.tp_size - * get_parallel().config.pp_size - ) + total_gpus = parallel.dp_size * parallel.tp_size * parallel.pp_size - nnodes = server_args.nnodes + nnodes = parallel.nnodes gpus_per_node = total_gpus // nnodes strategy = "STRICT_PACK" if nnodes == 1 else "SPREAD" @@ -300,8 +296,8 @@ class RayEngine(Engine): ray.get(pg.ready()) is_custom_pg = placement_group is not None - nnodes = server_args.nnodes - world_size = _compute_world_size(server_args) + nnodes = get_parallel().config.nnodes + world_size = _compute_world_size() if not is_custom_pg: engine_bundle, engine_ip = _find_engine_bundle(pg, nnodes) @@ -341,7 +337,7 @@ class RayEngine(Engine): _calculate_rank_ranges( nnodes, get_parallel().config.pp_size, - server_args.tp_size, + get_parallel().config.tp_size, node_rank=node_idx, ) ) @@ -377,9 +373,10 @@ class RayEngine(Engine): f"bundle_indices={bundle_indices}" ) + tp_size = get_parallel().config.tp_size for rank in range(world_size): - pp_rank = rank // server_args.tp_size - tp_rank = rank % server_args.tp_size + pp_rank = rank // tp_size + tp_rank = rank % tp_size bundle_idx = bundle_indices[rank] actor = _create_scheduler_actor( @@ -455,21 +452,18 @@ class RayEngine(Engine): 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/server_args.py b/python/sglang/srt/server_args.py index 17a4afa49..a90733df6 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -7100,7 +7100,11 @@ class ServerArgs: from sglang.srt.layers.cp.base import init_cp_strategy - init_cp_strategy(self) + init_cp_strategy( + enable_prefill_cp=bool(cfg.enable_prefill_cp), + cp_size=cfg.attn_cp_size, + cp_strategy=cfg.cp_strategy, + ) def _handle_dwdp(self): cfg = resolving_view(self) diff --git a/python/sglang/srt/speculative/draft_worker_common.py b/python/sglang/srt/speculative/draft_worker_common.py index c7c1c14fe..eb9ddb806 100644 --- a/python/sglang/srt/speculative/draft_worker_common.py +++ b/python/sglang/srt/speculative/draft_worker_common.py @@ -9,6 +9,7 @@ import torch from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode +from sglang.srt.runtime_context import attention_backends, get_spec from sglang.srt.server_args import DRAFT_ATTENTION_BACKEND_CHOICES, ServerArgs from sglang.srt.speculative.dflash_info import DFlashVerifyInput from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 @@ -28,12 +29,16 @@ class DraftWorkerBundle(msgspec.Struct, frozen=True): resolved_attention_backend: str -def _resolve_draft_attention_backend_fallback( - *, server_args: ServerArgs, algo_label: str -) -> str: - draft_backend = server_args.speculative_draft_attention_backend +def _resolve_draft_attention_backend_fallback(*, algo_label: str) -> str: + """The draft's attention backend, from the published leaves. + + `spec.speculative_draft_attention_backend` when the operator named one, + otherwise the process's prefill backend. Both are resolution's answers, so + they come from the bags. + """ + draft_backend = get_spec().speculative_draft_attention_backend if draft_backend is None: - draft_backend, _ = server_args.get_attention_backends() + draft_backend, _ = attention_backends() if draft_backend is None: return "triton" if torch.version.hip else "flashinfer" if draft_backend not in DRAFT_ATTENTION_BACKEND_CHOICES: @@ -65,9 +70,7 @@ def build_draft_tp_worker( # validated (e.g. a self-drafting architecture); it skips the generic # supported-backend fallback below. draft_backend = attention_backend_override or ( - _resolve_draft_attention_backend_fallback( - server_args=server_args, algo_label=algo_label - ) + _resolve_draft_attention_backend_fallback(algo_label=algo_label) ) from sglang.srt.layers.moe.utils import draft_model_build_scope diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index c7972f925..8f826768e 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -4623,11 +4623,15 @@ def cached_triton_kernel(key_fn=None): return decorator -def reserve_rope_cache_for_long_sequences( - model, server_args, model_config, logger=None -): - """Pre-expand RoPE cache for long sequences and speculative decoding.""" +def reserve_rope_cache_for_long_sequences(model, model_config, logger=None): + """Pre-expand RoPE cache for long sequences and speculative decoding. + + Runs inside `ModelRunner`, past publish, so the three config inputs come + from the bags: the context length and the two speculative counts are + resolution's answers. + """ from sglang.srt.environ import envs + from sglang.srt.runtime_context import get_model, get_spec SAFETY_FACTOR = envs.SGLANG_SPEC_EXPANSION_SAFETY_FACTOR.get() MARGIN = envs.SGLANG_ROPE_CACHE_SAFETY_MARGIN.get() @@ -4635,7 +4639,7 @@ def reserve_rope_cache_for_long_sequences( # 1) Estimate base context upper bound base_ctx = ( - getattr(server_args, "context_length", None) + get_model().context_length or getattr(model_config, "context_len", None) or getattr(model_config, "max_model_len", None) or getattr(model_config.hf_text_config, "max_position_embeddings", None) @@ -4643,8 +4647,8 @@ def reserve_rope_cache_for_long_sequences( ) # 2) Speculative decoding expansion - steps = int(getattr(server_args, "speculative_num_steps", 0) or 0) - draft = int(getattr(server_args, "speculative_num_draft_tokens", 0) or 0) + steps = int(get_spec().speculative_num_steps or 0) + draft = int(get_spec().speculative_num_draft_tokens or 0) reserve = base_ctx + steps * draft * SAFETY_FACTOR + MARGIN # 3) Align to reduce reallocation frequency diff --git a/python/sglang/srt/utils/hf_transformers/processor.py b/python/sglang/srt/utils/hf_transformers/processor.py index de50416db..c0c7a8411 100644 --- a/python/sglang/srt/utils/hf_transformers/processor.py +++ b/python/sglang/srt/utils/hf_transformers/processor.py @@ -54,11 +54,16 @@ from .tokenizer import ( _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/weight_cache/daemon.py b/python/sglang/srt/weight_cache/daemon.py index d14d79075..c61dd8394 100644 --- a/python/sglang/srt/weight_cache/daemon.py +++ b/python/sglang/srt/weight_cache/daemon.py @@ -47,6 +47,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple 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,28 @@ class WeightCacheDaemon: dist_init_method: Optional[str] = None, ): self.server_args = server_args - self.model_path = server_args.model_path + 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 +225,7 @@ class WeightCacheDaemon: 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 +283,7 @@ class WeightCacheDaemon: 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 +674,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 +708,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 +719,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 +741,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 +775,7 @@ def launch_weight_cache_daemons( ) logger.info( - f"All {num_daemons} weight cache daemons on node {server_args.node_rank} are ready " + f"All {num_daemons} weight cache daemons on node {cfg.node_rank} are ready " f"(pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, " f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, " f"dist_init_method={dist_init_method})" diff --git a/python/sglang/test/scripted_runtime/scheduler_hook.py b/python/sglang/test/scripted_runtime/scheduler_hook.py index 58253cfcb..51462b5d4 100644 --- a/python/sglang/test/scripted_runtime/scheduler_hook.py +++ b/python/sglang/test/scripted_runtime/scheduler_hook.py @@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Generator, List, Optional, Tuple import zmq +from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.environ import envs from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle from sglang.srt.utils.network import get_zmq_socket @@ -56,8 +57,8 @@ def _drive_engine_through_warmup(ctx: ScriptedContext) -> Generator: """Run the engine until the server warmup request has been received and fully processed, so scripts never observe foreign warmup traffic.""" scheduler = ctx.scheduler - server_args = scheduler.server_args - if server_args.skip_server_warmup: + cfg = resolving_view(scheduler.server_args) + if cfg.skip_server_warmup: logger.info("scripted_runtime: skip_server_warmup set, not driving warmup") return @@ -67,7 +68,7 @@ def _drive_engine_through_warmup(ctx: ScriptedContext) -> Generator: # is_fully_idle() can transiently report idle while a PP microbatch result # is still in flight, so require it to hold for two full microbatch # rotations after the warmup request was observed on the recv socket. - quiesce_iters = 2 * (server_args.pp_size + server_args.pp_async_batch_depth) + quiesce_iters = 2 * (cfg.pp_size + cfg.pp_async_batch_depth) proxy = ctx._tokenizer_recv_proxy deadline = start_time + WARMUP_DRIVE_TIMEOUT_S diff --git a/test/manual/ep/test_flashinfer_dispatcher.py b/test/manual/ep/test_flashinfer_dispatcher.py index cbb6bccdf..8bb540f66 100644 --- a/test/manual/ep/test_flashinfer_dispatcher.py +++ b/test/manual/ep/test_flashinfer_dispatcher.py @@ -10,6 +10,7 @@ from sglang.srt.distributed.parallel_state import ( 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 @@ class TestFlashinferDispatcher(CustomTestCase): server_args.moe_runner_backend = "flashinfer_cutlass" server_args.moe_a2a_backend = "flashinfer" set_global_server_args_for_scheduler(server_args) - initialize_moe_config(server_args) + publish(server_args, role="scheduler") + initialize_moe_config() init_distributed_environment( world_size=-1, # Auto-detect from environment diff --git a/test/registered/cp/test_cp_strategy_unit.py b/test/registered/cp/test_cp_strategy_unit.py index 4b8d6f8bd..53cc4fa2e 100644 --- a/test/registered/cp/test_cp_strategy_unit.py +++ b/test/registered/cp/test_cp_strategy_unit.py @@ -60,7 +60,7 @@ class _FakeCPGroup: 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 @@ class TestCPStrategyUnit(CustomTestCase): 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 @@ class TestCPStrategyUnit(CustomTestCase): 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 @@ class TestCPStrategyUnit(CustomTestCase): 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 @@ class TestPrefillCPBCGReplay(CustomTestCase): 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 @@ class TestPrefillCPBCGReplay(CustomTestCase): 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 @@ class TestCPZigzagStrategy(CustomTestCase): class TestCPInterleaveStrategy(CustomTestCase): def setUp(self): init_cp_strategy( - SimpleNamespace( - enable_prefill_cp=True, - cp_strategy="interleave", - attn_cp_size=4, - attention_backend="fa3", - ) + enable_prefill_cp=True, + cp_size=4, + cp_strategy="interleave", ) def tearDown(self): - init_cp_strategy(SimpleNamespace(enable_prefill_cp=False)) + init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="zigzag") def _metadata_for_rank(self, rank, *, cp_size, seq_lens, extend_seq_lens): strategy = InterleaveCPStrategy(cp_size=cp_size) diff --git a/test/registered/unit/configs/test_embedding_model_spec.py b/test/registered/unit/configs/test_embedding_model_spec.py index 6e1d5018e..bd8850c4b 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 @@ class TestEmbeddingModelSpec(unittest.TestCase): ) 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/layers/test_flashinfer_comm_fusion.py b/test/registered/unit/layers/test_flashinfer_comm_fusion.py index 9eaf5c8fc..5fdc935b8 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 @@ class TestFlashInferCommFusion(CustomTestCase): 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 @@ class TestFlashInferCommFusion(CustomTestCase): 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 @@ class TestFlashInferCommFusion(CustomTestCase): 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/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index c94a30d2a..9407f3b57 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -27,6 +27,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( 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 @@ -2356,9 +2357,13 @@ class TestGrpcServerArgs(CustomTestCase): 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", @@ -2373,7 +2378,7 @@ class TestGrpcServerArgs(CustomTestCase): tokenizer_manager=MagicMock(), template_manager=MagicMock(), scheduler_info={}, - grpc_port=resolution_result(server_args, "grpc_port"), + grpc_port=get_serving().grpc_port, ) self.assertEqual(handle, "handle") diff --git a/test/registered/unit/spec/test_draft_construction_isolation.py b/test/registered/unit/spec/test_draft_construction_isolation.py index 0469ed4e6..5609009d9 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 @@ from sglang.srt.layers.moe.utils import ( 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 @@ class TestFusionDecisionFlag(CustomTestCase): from sglang.srt.server_args import ServerArgs self._seed() - initialize_moe_config( - ServerArgs(model_path="dummy", disable_shared_experts_fusion=True) + publish( + ServerArgs(model_path="dummy", disable_shared_experts_fusion=True), + role="scheduler", ) + initialize_moe_config() moe = get_flags().moe self.assertTrue(moe.disable_shared_experts_fusion) self.assertTrue(moe.speculative_disable_shared_experts_fusion) diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index 6b840f970..b183374b6 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 @@ _SLOT_OWNERS = ("srt/runtime_context.py", "srt/server_args.py", "srt/arg_groups/ # 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 @@ _CONFIGURED_SIZE_CALL_SITES = { "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 @@ _CONFIGURED_SIZE_CALL_SITES = { ("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 @@ _CONFIGURED_SIZE_CALL_SITES = { "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_ray_driver_reads_the_bags.py b/test/registered/unit/test_ray_driver_reads_the_bags.py new file mode 100644 index 000000000..7532334e7 --- /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 11d4fd836..fbcc472f1 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -519,8 +519,6 @@ class TestMoeFlagsGroup(_IsolatedServerArgs): swap under the speculative contexts and restore on exit.""" def _init(self, **kw): - from types import SimpleNamespace - from sglang.srt.layers.moe.utils import initialize_moe_config defaults = dict( @@ -538,7 +536,12 @@ class TestMoeFlagsGroup(_IsolatedServerArgs): disable_shared_experts_fusion=False, ) defaults.update(kw) - initialize_moe_config(SimpleNamespace(**defaults)) + # The flags are seeded from the bags, so the test publishes a config + # carrying these values. + override = get_context().override_server_args(**defaults) + override.install() + self.addCleanup(override.restore) + initialize_moe_config() def test_lazy_defaults_before_initialize(self): from sglang.srt.layers.moe.utils import ( diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index 23e863670..95107549f 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -134,51 +134,12 @@ _ENV_MATRIX = (({}, {"SGLANG_IS_IN_CI": "true"}),) # are step-12 exposure like any other pair. _PASSED = frozenset({"model_path", "device", "random_seed"}) -# The reads that still take a value off the supplied instance. `initialize_moe_config` -# is handed the record until the replay goes away; the rest are pre-publish launcher -# reads. -_EXPOSED = { - ("disaggregation/encoder/server.py", "model_loader_extra_config"), - ("layers/moe/utils.py", "deepep_mode"), - ("layers/moe/utils.py", "disable_shared_experts_fusion"), - ("layers/moe/utils.py", "moe_a2a_backend"), - ("layers/moe/utils.py", "moe_runner_backend"), - ("layers/moe/utils.py", "quantization"), - ("layers/moe/utils.py", "speculative_moe_runner_backend"), - ("configs/embedding_model_spec.py", "chunked_prefill_size"), - ("configs/embedding_model_spec.py", "cuda_graph_config"), - ("configs/embedding_model_spec.py", "disable_radix_cache"), - ("configs/embedding_model_spec.py", "is_embedding"), - ("configs/embedding_model_spec.py", "prefill_only_disable_kv_cache"), - ("entrypoints/engine.py", "enable_symm_mem"), - ("entrypoints/engine.py", "reasoning_parser"), - ("entrypoints/engine.py", "tool_call_parser"), - ("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"), - ("layers/moe/utils.py", "deepep_mode"), - ("layers/moe/utils.py", "moe_a2a_backend"), - ("layers/moe/utils.py", "moe_runner_backend"), - ("layers/moe/utils.py", "quantization"), - ("layers/moe/utils.py", "speculative_moe_runner_backend"), - ("speculative/draft_worker_common.py", "speculative_draft_attention_backend"), - ("utils/common.py", "speculative_num_draft_tokens"), - ("utils/common.py", "speculative_num_steps"), - ("utils/hf_transformers/processor.py", "image_processor_backend"), - ("weight_cache/daemon.py", "attn_cp_size"), - ("weight_cache/daemon.py", "deepep_mode"), - ("weight_cache/daemon.py", "dp_size"), - ("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 @@ -191,28 +152,7 @@ _EXPOSED_CUDA_ONLY: frozenset = frozenset() # Axis two: (file, field) pairs where a supplied-instance read names a field that # some code overrides post-publish. Each needs an ordering judgment, not a blanket # conversion; the list exists so a new one is a decision made when it is written. -_OVERRIDDEN_AND_READ = { - ("entrypoints/engine.py", "reasoning_parser"), - ("entrypoints/engine.py", "tool_call_parser"), - ("weight_cache/daemon.py", "dp_size"), - ("weight_cache/daemon.py", "dtype"), - ("weight_cache/daemon.py", "ep_size"), - ("weight_cache/daemon.py", "load_format"), - ("weight_cache/daemon.py", "model_path"), - ("mem_cache/pool_host/common.py", "hicache_storage_backend"), - ("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"), - ("mem_cache/unified_radix_cache.py", "hicache_storage_backend"), - ("mem_cache/unified_radix_cache.py", "hicache_storage_backend_extra_config"), - ("mem_cache/unified_radix_cache.py", "hicache_storage_prefetch_policy"), - ("mem_cache/unified_radix_cache.py", "hicache_write_policy"), - ("utils/common.py", "speculative_num_draft_tokens"), - ("utils/common.py", "speculative_num_steps"), - ("weight_cache/daemon.py", "dp_size"), - ("weight_cache/daemon.py", "dtype"), - ("weight_cache/daemon.py", "ep_size"), - ("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: