config: the runtime readers take the published bags (#36254)

This commit is contained in:
Cheng Wan
2026-08-26 05:08:25 -07:00
committed by GitHub
parent 5b7fc61306
commit 937af8538b
67 changed files with 796 additions and 552 deletions
+10 -10
View File
@@ -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
@@ -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
+69 -38
View File
@@ -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:
+3 -1
View File
@@ -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,
)
+28 -13
View File
@@ -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):
@@ -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(
+6 -4
View File
@@ -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
@@ -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,
},
}
@@ -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()
@@ -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:
@@ -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"):
@@ -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)
@@ -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}")
@@ -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,
@@ -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),
+2 -1
View File
@@ -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
@@ -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
+21 -22
View File
@@ -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()
+2 -1
View File
@@ -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)
+12 -3
View File
@@ -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.
+29 -22
View File
@@ -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
+1 -1
View File
@@ -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,
+1 -1
View File
@@ -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,
@@ -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
@@ -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()
@@ -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
@@ -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 "
+19 -15
View File
@@ -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
+1 -1
View File
@@ -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)
)
@@ -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
+34 -21
View File
@@ -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
)
@@ -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"
@@ -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"
@@ -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):
+9 -6
View File
@@ -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,
+3 -3
View File
@@ -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
@@ -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
@@ -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,
)
+1 -1
View File
@@ -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)(
@@ -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,
@@ -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(
@@ -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)
@@ -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)
@@ -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
@@ -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")
@@ -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,
)
@@ -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
@@ -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
@@ -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)
@@ -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,
)
+27 -33
View File
@@ -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
+5 -1
View File
@@ -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)
@@ -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
+11 -7
View File
@@ -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
@@ -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(
+36 -33
View File
@@ -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://<node0-ip>:<port> 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})"
@@ -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
+3 -1
View File
@@ -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
+19 -31
View File
@@ -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)
@@ -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(
@@ -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()
@@ -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")
@@ -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)
@@ -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 "
@@ -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()
+6 -3
View File
@@ -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 (
@@ -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: