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 `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 **retracted** — owner ruling (2026-08-15): a process holds at most one live
config at a time (concurrent multi-Engine is unsupported; sequential rebuild 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 stays legal, unit tests rely on it). Nothing in those files reads the instance
files is pinned pair by pair in the exposure ratchet, each with its own any more -- the exposure ratchet's pin set is empty, so the next such read is a
disposition; none of it is a boundary to imitate. What new entry that has to argue for itself. What
genuinely stays per-instance is what differs per *worker* within one engine: genuinely stays per-instance is what differs per *worker* within one engine:
`base_gpu_id` travels as a constructor argument (`MMEncoder(gpu_id=...)`; `base_gpu_id` travels as a constructor argument (`MMEncoder(gpu_id=...)`;
`BaseMultimodalProcessor._fast_image_processor_device` is the shape to copy). `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 supplied-instance contract; don't rewrite the parameter reads unless the
field is runtime-mutated (see the elastic-EP `ep_size` case in 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 `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 and the callee runs in a process that has published.** That second case is a
pinned debt, not a style question: the record is destined to carry the decision, not a style question: the record carries the user's raw input, so a
user's raw input, so `server_args.page_size` inside a runner-owned resolution-filled field read off it inside a runner-owned constructor answers
constructor will read the raw pre-resolution value instead of the effective with the pre-resolution value instead of the effective one. Debt means a decision, not automatically a bag read: pick where the
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 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 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 `gpu_id` above are dispositions of exactly this debt). The per-instance
boundaries above are **not** exempt from this unless-clause (the multi-Engine boundaries above are **not** exempt from this unless-clause (the multi-Engine
exemption is retracted); each one gets its own disposition. exemption is retracted); each one gets its own disposition.
`test_supplied_instance_exposure_ratchet.py` `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 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 = 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 `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, - Groups are typed dataclasses on `Flags` (`capture` / `moe` / `dp`): typo-safe writes,
transactional test-only `override(**kw)` context manager. 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 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. 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 - `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.datasets.random import sample_random_requests
from sglang.benchmark.utils import get_tokenizer, set_ulimit from sglang.benchmark.utils import get_tokenizer, set_ulimit
from sglang.lang.backend.runtime_endpoint import Runtime 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.entrypoints.engine import Engine
from sglang.srt.server_args import ServerArgs 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 RayEngine requires a placement group, so we launch it inside a Ray actor
and return a lightweight proxy that forwards calls via ray.get(). and return a lightweight proxy that forwards calls via ray.get().
""" """
cfg = resolving_view(server_args)
import ray import ray
from ray.runtime_env import RuntimeEnv from ray.runtime_env import RuntimeEnv
from ray.util.placement_group import placement_group 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(): if not ray.is_initialized():
ray.init(runtime_env=RuntimeEnv(env_vars=env_vars)) 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") pg = placement_group([{"CPU": 1, "GPU": total_gpus}], strategy="STRICT_PACK")
ray.get(pg.ready()) ray.get(pg.ready())
@@ -398,7 +400,7 @@ def _create_ray_engine_backend(server_args: ServerArgs):
placement_group=pg, placement_group=pg,
placement_group_bundle_index=0, placement_group_bundle_index=0,
), ),
).remote(**dict(server_args._raw_input)) ).remote(**dict(cfg._raw_input))
class _Proxy: class _Proxy:
"""Forwards method calls to the remote RayEngine actor.""" """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 # A programmatic caller may hand over a freshly constructed record, and
# the backends below read the resolved paths and the raw snapshot. # 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 bench_args.backend == "engine":
if server_args.use_ray: if cfg.use_ray:
backend = _create_ray_engine_backend(server_args) backend = _create_ray_engine_backend(server_args)
else: else:
backend = Engine(server_args=server_args) backend = Engine(server_args=server_args)
if not backend: if not backend:
raise ValueError("Please provide valid engine arguments") raise ValueError("Please provide valid engine arguments")
elif bench_args.backend == "runtime": elif bench_args.backend == "runtime":
backend = Runtime(**dict(server_args._raw_input)) backend = Runtime(**dict(cfg._raw_input))
else: else:
raise ValueError('Please set backend to either "engine" or "runtime"') 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) tokenizer = get_tokenizer(tokenizer_id)
# Set global environments # Set global environments
+69 -38
View File
@@ -64,6 +64,7 @@ import numpy as np
import torch import torch
import torch.distributed as dist 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.configs.model_config import ModelConfig
from sglang.srt.distributed.parallel_state import ( from sglang.srt.distributed.parallel_state import (
destroy_distributed_environment, 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.schedule_batch import Req, ScheduleBatch
from sglang.srt.managers.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw 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.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.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_parallel, get_schedule, publish 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): def load_model(server_args, port_args, gpu_id, tp_rank):
cfg = resolving_view(server_args)
suppress_other_loggers() suppress_other_loggers()
rank_print = print if tp_rank == 0 else lambda *args, **kwargs: None 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) model_config = ModelConfig.from_server_args(server_args)
attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = ( attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = (
compute_dp_attention_world_info( compute_dp_attention_world_info(
server_args.enable_dp_attention, cfg.enable_dp_attention,
tp_rank, tp_rank,
server_args.tp_size, cfg.tp_size,
server_args.dp_size, cfg.dp_size,
server_args.attn_cp_size, cfg.attn_cp_size,
) )
) )
ps = ParallelState( ps = ParallelState(
tp_rank=tp_rank, tp_rank=tp_rank,
tp_size=server_args.tp_size, tp_size=cfg.tp_size,
pp_rank=0, pp_rank=0,
pp_size=1, pp_size=1,
dp_rank=None, dp_rank=None,
dp_size=server_args.dp_size, dp_size=cfg.dp_size,
attn_tp_rank=attn_tp_rank, attn_tp_rank=attn_tp_rank,
attn_tp_size=attn_tp_size, attn_tp_size=attn_tp_size,
attn_cp_rank=0, attn_cp_rank=0,
attn_cp_size=server_args.attn_cp_size, attn_cp_size=cfg.attn_cp_size,
attn_dcp_rank=tp_rank % server_args.dcp_size, attn_dcp_rank=tp_rank % cfg.dcp_size,
attn_dcp_size=server_args.dcp_size, attn_dcp_size=cfg.dcp_size,
attn_dp_rank=attn_dp_rank, attn_dp_rank=attn_dp_rank,
attn_dp_size=attn_dp_size, attn_dp_size=attn_dp_size,
moe_ep_rank=moe_ep_rank, 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_rank=None,
moe_dp_size=server_args.moe_dp_size, moe_dp_size=cfg.moe_dp_size,
gpu_id=gpu_id, gpu_id=gpu_id,
) )
runner_kwargs = dict( runner_kwargs = dict(
model_config=model_config, model_config=model_config,
mem_fraction_static=server_args.mem_fraction_static, mem_fraction_static=cfg.mem_fraction_static,
gpu_id=gpu_id, gpu_id=gpu_id,
ps=ps, ps=ps,
nccl_port=port_args.nccl_port, 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) model_runner = MlxModelRunnerStub(**runner_kwargs)
else: else:
model_runner = ModelRunner(**runner_kwargs) 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.start_startup_weight_load()
model_runner.alloc_memory_pool() model_runner.alloc_memory_pool()
model_runner.init_attention_backends() model_runner.init_attention_backends()
model_runner.init_cuda_graphs() 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() model_runner.finalize_startup_weight_load()
rank_print(f"max_total_num_tokens={model_runner.max_total_num_tokens}") rank_print(f"max_total_num_tokens={model_runner.max_total_num_tokens}")
tokenizer = get_tokenizer( tokenizer = get_tokenizer(
server_args.tokenizer_path, cfg.tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode, tokenizer_mode=cfg.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=cfg.trust_remote_code,
) )
if server_args.tp_size > 1: if cfg.tp_size > 1:
dist.barrier() dist.barrier()
if _use_mlx: if _use_mlx:
@@ -584,19 +590,20 @@ class _MlxBenchRunner:
"""Wraps MlxModelRunner for the MLX benchmark path.""" """Wraps MlxModelRunner for the MLX benchmark path."""
def __init__(self, model_runner, server_args): def __init__(self, model_runner, server_args):
cfg = resolving_view(server_args)
from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner
# Radix cache requires the scheduler's allocator/trie; disable in # Radix cache requires the scheduler's allocator/trie; disable in
# standalone bench mode where no scheduler is present. # standalone bench mode where no scheduler is present.
init_kwargs = dict( init_kwargs = dict(
model_path=server_args.model_path, model_path=cfg.model_path,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=cfg.trust_remote_code,
disable_radix_cache=True, disable_radix_cache=True,
mem_fraction_static=server_args.mem_fraction_static, mem_fraction_static=cfg.mem_fraction_static,
quantization=server_args.quantization, quantization=cfg.quantization,
) )
if server_args.max_total_tokens is not None: if cfg.max_total_tokens is not None:
init_kwargs["pool_size"] = server_args.max_total_tokens init_kwargs["pool_size"] = cfg.max_total_tokens
self.mlx_runner = MlxModelRunner(**init_kwargs) self.mlx_runner = MlxModelRunner(**init_kwargs)
self.mlx_runner.init_cache_pools(req_to_token_pool=None) self.mlx_runner.init_cache_pools(req_to_token_pool=None)
self.fake_torch_runner = model_runner self.fake_torch_runner = model_runner
@@ -883,17 +890,18 @@ def latency_test(
gpu_id, gpu_id,
tp_rank, tp_rank,
): ):
cfg = resolving_view(server_args)
# `main` runs this inline for tp_size == 1 and spawns it per rank otherwise; # `main` runs this inline for tp_size == 1 and spawns it per rank otherwise;
# a spawned child arrives with nothing published. # a spawned child arrives with nothing published.
publish(server_args, role="scheduler") publish(server_args, role="scheduler")
initialize_moe_config(server_args) initialize_moe_config()
initialize_fp8_gemm_config(server_args) initialize_fp8_gemm_config()
initialize_fp4_gemm_config(server_args) initialize_fp4_gemm_config()
# Set CPU affinity
if get_bool_env_var("SGLANG_SET_CPU_AFFINITY"): if get_bool_env_var("SGLANG_SET_CPU_AFFINITY"):
parallel = get_parallel().config
set_gpu_proc_affinity( 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 # Configure the logger
@@ -988,22 +996,45 @@ def latency_test(
for result in result_list: for result in result_list:
fout.write(json.dumps(result) + "\n") fout.write(json.dumps(result) + "\n")
if server_args.tp_size > 1: if cfg.tp_size > 1:
destroy_model_parallel() destroy_model_parallel()
destroy_distributed_environment() destroy_distributed_environment()
def main(server_args, bench_args): 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() server_args.resolve_once()
cfg = resolving_view(server_args)
# 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)
_set_envs_and_config(server_args) _set_envs_and_config(server_args)
if server_args.model_path: if cfg.model_path:
if bench_args.correctness_test: if bench_args.correctness_test:
work_func = correctness_test work_func = correctness_test
else: 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.endpoint import acquire_endpoint
from sglang.benchmark.utils import get_processor, get_tokenizer from sglang.benchmark.utils import get_processor, get_tokenizer
from sglang.profiler import run_profile 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.disaggregation.utils import FAKE_BOOTSTRAP_HOST
from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -1234,6 +1235,7 @@ def run_benchmark_internal(
def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs): def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
cfg = resolving_view(server_args)
results, server_info = run_benchmark_internal(server_args, bench_args) results, server_info = run_benchmark_internal(server_args, bench_args)
# Save results as pydantic models in the JSON format # 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( save_results_as_pydantic_models(
results, results,
pydantic_result_filename=bench_args.pydantic_result_filename, 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, server_args=bench_args.server_args_for_metrics,
) )
+28 -13
View File
@@ -17,13 +17,19 @@ import time
import requests import requests
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST
from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.entrypoints.warmup import warmup from sglang.srt.entrypoints.warmup import warmup
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import GenerateReqInput from sglang.srt.managers.io_struct import GenerateReqInput
from sglang.srt.managers.tokenizer_manager import TokenizerManager 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.server_args import ServerArgs
from sglang.srt.utils import kill_process_tree from sglang.srt.utils import kill_process_tree
@@ -59,8 +65,7 @@ async def warm_up_compile(
disaggregation_mode: str, tokenizer_manager: TokenizerManager disaggregation_mode: str, tokenizer_manager: TokenizerManager
): ):
print("\nGenerate warm up request for compiling DeepGEMM...\n") print("\nGenerate warm up request for compiling DeepGEMM...\n")
server_args = tokenizer_manager.server_args dp_size = get_parallel().config.dp_size
dp_size = server_args.dp_size
base_ids = [0, 1, 2, 3] base_ids = [0, 1, 2, 3]
sampling_params = { sampling_params = {
"temperature": 0.0, "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_host = [FAKE_BOOTSTRAP_HOST] * dp_size
generate_req_input.bootstrap_room = [ 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: else:
input_ids = ( 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 # Keeps the device probe out of the fork below, for a caller that reaches
# this without resolving first. # this without resolving first.
server_args.resolve_once() server_args.resolve_once()
cfg = resolving_view(server_args)
proc = multiprocessing.Process(target=launch_server_internal, args=(server_args,)) proc = multiprocessing.Process(target=launch_server_internal, args=(server_args,))
proc.start() proc.start()
@@ -125,7 +132,7 @@ def launch_server_process_and_send_one_request(
if response.status_code == 200: if response.status_code == 200:
# Rank-0 node send a request to sync with other node and then return. # Rank-0 node send a request to sync with other node and then return.
if server_args.node_rank == 0: if server_args.node_rank == 0:
dp_size = server_args.dp_size dp_size = cfg.dp_size
base_ids = [0, 1, 2, 3] base_ids = [0, 1, 2, 3]
payload = { payload = {
"sampling_params": { "sampling_params": {
@@ -133,11 +140,11 @@ def launch_server_process_and_send_one_request(
"temperature": 0, "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["input_ids"] = [list(base_ids) for _ in range(dp_size)]
payload["bootstrap_host"] = [FAKE_BOOTSTRAP_HOST] * dp_size payload["bootstrap_host"] = [FAKE_BOOTSTRAP_HOST] * dp_size
payload["bootstrap_room"] = [ 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) for i in range(dp_size)
] ]
else: 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 """The config this script serves with: no cuda graph, no torch compile, and a
watchdog that outlives the compilation.""" watchdog that outlives the compilation."""
args.enable_torch_compile = False 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. # Watchdog timeout follows compile_args.timeout because compilation takes long.
args.watchdog_timeout = compile_args.timeout args.watchdog_timeout = compile_args.timeout
args.warmups = "compile-deep-gemm" 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...") 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): def run_compile(server_args: ServerArgs, compile_args: CompileArgs):
@@ -456,13 +456,15 @@ class Runtime:
self.endpoint.cache_prefix(prefix) self.endpoint.cache_prefix(prefix)
def get_tokenizer(self): def get_tokenizer(self):
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.utils.hf_transformers_utils import get_tokenizer from sglang.srt.utils.hf_transformers_utils import get_tokenizer
cfg = resolving_view(self.server_args)
return get_tokenizer( return get_tokenizer(
self.server_args.tokenizer_path or self.server_args.model_path, cfg.tokenizer_path or cfg.model_path,
tokenizer_mode=self.server_args.tokenizer_mode, tokenizer_mode=cfg.tokenizer_mode,
trust_remote_code=self.server_args.trust_remote_code, trust_remote_code=cfg.trust_remote_code,
revision=self.server_args.revision, revision=cfg.revision,
) )
async def async_generate( async def async_generate(
+6 -4
View File
@@ -5,6 +5,7 @@ import os
import sys import sys
import warnings import warnings
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.plugins import load_plugins from sglang.srt.plugins import load_plugins
from sglang.srt.server_args import prepare_server_args from sglang.srt.server_args import prepare_server_args
from sglang.srt.utils import kill_process_tree 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` # The flags dispatched on below are decided by resolution (`--grpc-mode`
# folds into `smg_grpc_mode`), and `prepare_server_args` returns raw input. # folds into `smg_grpc_mode`), and `prepare_server_args` returns raw input.
server_args.resolve_once() server_args.resolve_once()
cfg = resolving_view(server_args)
if server_args.encoder_only: if cfg.encoder_only:
# For encoder disaggregation # 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 ( from sglang.srt.disaggregation.encoder.grpc_server import (
serve_grpc_encoder, serve_grpc_encoder,
) )
@@ -31,7 +33,7 @@ def run_server(server_args):
from sglang.srt.disaggregation.encoder.http_server import launch_server from sglang.srt.disaggregation.encoder.http_server import launch_server
launch_server(server_args) 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 # Legacy SMG gRPC server (--smg-grpc-mode, or the deprecated --grpc-mode
# which __post_init__ folds into smg_grpc_mode). The native Rust gRPC # which __post_init__ folds into smg_grpc_mode). The native Rust gRPC
# server is a separate path, enabled by --grpc-port, that starts # 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 from sglang.srt.entrypoints.grpc_server import serve_grpc
asyncio.run(serve_grpc(server_args)) asyncio.run(serve_grpc(server_args))
elif server_args.use_ray: elif cfg.use_ray:
# Ray mode: HTTP mode with Ray backend. # Ray mode: HTTP mode with Ray backend.
try: try:
from sglang.srt.ray.http_server import launch_server from sglang.srt.ray.http_server import launch_server
@@ -221,17 +221,17 @@ def _native_embedding_spec(
def resolved_embedding_plan( def resolved_embedding_plan(
spec: EmbeddingModelSpec, *, server_args: Any, model_config: Any spec: EmbeddingModelSpec, *, config: Any, model_config: Any
) -> dict[str, Any]: ) -> dict[str, Any]:
"""Combine static capabilities with the effective server configuration. """Combine static capabilities with the effective server configuration.
This boundary deliberately accepts duck-typed arguments so the declarative This boundary deliberately accepts duck-typed arguments so the declarative
registry remains independent of ServerArgs and ModelConfig import cycles. 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( prefill_graph = getattr(getattr(config, "cuda_graph_config", None), "prefill", None)
getattr(server_args, "cuda_graph_config", None), "prefill", None
)
backend = getattr(prefill_graph, "backend", None) backend = getattr(prefill_graph, "backend", None)
backend_value = getattr(backend, "value", backend) backend_value = getattr(backend, "value", backend)
capture_sizes = getattr(prefill_graph, "bs", None) or [] capture_sizes = getattr(prefill_graph, "bs", None) or []
@@ -239,7 +239,7 @@ def resolved_embedding_plan(
return { return {
**spec.as_dict(), **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)), "supports_dimensions": bool(getattr(model_config, "is_matryoshka", False)),
"matryoshka_dimensions": list( "matryoshka_dimensions": list(
getattr(model_config, "matryoshka_dimensions", None) or [] getattr(model_config, "matryoshka_dimensions", None) or []
@@ -252,14 +252,10 @@ def resolved_embedding_plan(
}, },
"cache": { "cache": {
"kv_cache_disabled": bool( "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( "radix_cache_disabled": bool(getattr(config, "disable_radix_cache", False)),
getattr(server_args, "disable_radix_cache", False) "chunked_prefill_disabled": getattr(config, "chunked_prefill_size", None)
),
"chunked_prefill_disabled": getattr(
server_args, "chunked_prefill_size", None
)
== -1, == -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.disaggregation.encoder.server import MMEncoder, launch_encoder
from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle
from sglang.srt.managers.schedule_batch import Modality 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.server_args import PortArgs, ServerArgs
from sglang.srt.utils import random_uuid from sglang.srt.utils import random_uuid
from sglang.srt.utils.network import NetworkAddress, get_zmq_socket 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() dist_init_method = na.to_tcp()
else: else:
dist_init_method = NetworkAddress( 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() ).to_tcp()
send_sockets: List[zmq.Socket] = [] 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}" schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}"
send_sockets.append( send_sockets.append(
get_zmq_socket(zmq_ctx, zmq.PUSH, schedule_path, bind=False) 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) 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) server.add_insecure_port(listen_addr)
await server.start() await server.start()
@@ -113,11 +113,11 @@ def _register_encoder_url_with_bootstrap(server_args: ServerArgs):
instead of serialising sleeps in a single thread. 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", "::"): 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" 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} payload = {"url": encoder_url}
bootstrap_urls = list(server_args.encoder_register_urls) bootstrap_urls = list(server_args.encoder_register_urls)
if not bootstrap_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): 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", "::"): 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" 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} payload = {"url": encoder_url}
for bootstrap_url in server_args.encoder_register_urls: 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() use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get()
self.use_image_processor_gpu = ( self.use_image_processor_gpu = (
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) self._load_mm_processor(server_args)
@@ -158,7 +158,7 @@ class EncoderPreprocessor:
def _load_mm_processor(self, server_args: ServerArgs): def _load_mm_processor(self, server_args: ServerArgs):
from transformers import AutoImageProcessor, AutoVideoProcessor 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 = ( image_processor_kwargs = (
{} {}
if image_processor_backend == "auto" if image_processor_backend == "auto"
@@ -167,7 +167,7 @@ class EncoderPreprocessor:
try: try:
self.image_processor = AutoImageProcessor.from_pretrained( self.image_processor = AutoImageProcessor.from_pretrained(
get_serving().tokenizer_path or get_model().model_path, 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, revision=server_args.revision,
**image_processor_kwargs, **image_processor_kwargs,
) )
@@ -178,7 +178,7 @@ class EncoderPreprocessor:
try: try:
self.video_processor = AutoVideoProcessor.from_pretrained( self.video_processor = AutoVideoProcessor.from_pretrained(
get_serving().tokenizer_path or get_model().model_path, 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, revision=server_args.revision,
) )
except Exception as e: except Exception as e:
@@ -188,7 +188,7 @@ class EncoderPreprocessor:
try: try:
_audio_proc = AutoProcessor.from_pretrained( _audio_proc = AutoProcessor.from_pretrained(
get_serving().tokenizer_path or get_model().model_path, 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, revision=server_args.revision,
) )
if not hasattr(_audio_proc, "feature_extractor"): 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.managers.schedule_batch import Modality, Req
from sglang.srt.multimodal.cache import media_preprocess_kwargs from sglang.srt.multimodal.cache import media_preprocess_kwargs
from sglang.srt.multimodal.transport import determine_tensor_transport_mode 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.server_args import ServerArgs
from sglang.srt.utils import ImageData from sglang.srt.utils import ImageData
from sglang.srt.utils.common import safe_pickle_loads 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 # When None (e.g. in a scheduler subprocess that has no in-process
# bootstrap), fall back to a snapshot of the static --encoder-urls. # bootstrap), fall back to a snapshot of the static --encoder-urls.
self.encode_urls: List[str] = ( 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.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.pp_rank = pp_rank
self.tp_rank = tp_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.tp_group = tp_group
self.nnodes = server_args.nnodes self.nnodes = server_args.nnodes
self.hostname = get_local_ip_auto() self.hostname = get_local_ip_auto()
@@ -1836,9 +1843,9 @@ class MMReceiverBase(ABC):
_processor = get_processor( _processor = get_processor(
get_serving().tokenizer_path, get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode, 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, revision=server_args.revision,
image_processor_backend=resolve_image_processor_backend(server_args), image_processor_backend=resolve_image_processor_backend(get_mm()),
**extra_kwargs, **extra_kwargs,
) )
@@ -2659,7 +2666,7 @@ def create_mm_receiver(
transport_mode = envs.SGLANG_ENCODER_MM_RECEIVER_MODE.get() transport_mode = envs.SGLANG_ENCODER_MM_RECEIVER_MODE.get()
logger.debug(f"MMReceiver transport_mode from env: {transport_mode}") 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}") logger.info(f"EPD MMReceiver: using transport_mode={transport_mode}")
receiver_cls = _MM_RECEIVER_BY_MODE.get(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 HTTP uses this entry point today. gRPC can reuse it later without
importing HTTP application state or Uvicorn. 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( raise ValueError(
"Encoder DP mode requires --dp-size > 1 and --tp-size 1; got " "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 dp_size = get_parallel().config.dp_size
logger.info(f"Launching encoder in DP mode: dp_size={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_exec,
get_mm, get_mm,
get_model, get_model,
get_parallel,
publish, publish,
) )
from sglang.srt.server_args import ServerArgs 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 this instance's value, not a config change, so it travels as an
argument.""" argument."""
assert_published(server_args, role="encoder") 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 self.server_args = server_args
configure_media_url_security( configure_media_url_security(
get_mm().allowed_media_domains, get_mm().allowed_media_domains,
@@ -470,7 +471,7 @@ class MMEncoder:
self.load_config = LoadConfig( self.load_config = LoadConfig(
load_format=get_model().load_format, load_format=get_model().load_format,
download_dir=server_args.download_dir, 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_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_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, 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( init_distributed_environment(
backend=get_default_distributed_backend(self.device), backend=get_default_distributed_backend(self.device),
world_size=server_args.tp_size, world_size=get_parallel().config.tp_size,
rank=rank, rank=rank,
distributed_init_method=dist_init_method, distributed_init_method=dist_init_method,
local_rank=rank, 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) initialize_dp_attention(server_args, self.model_config)
self.model = load_model( self.model = load_model(
@@ -554,7 +557,7 @@ class MMEncoder:
) )
self.mm_global_cache = EmbeddingCacheController( self.mm_global_cache = EmbeddingCacheController(
rank, rank,
server_args.tp_size, get_parallel().config.tp_size,
embedding_store=embedding_store, embedding_store=embedding_store,
hidden_dims=self._embedding_dims, hidden_dims=self._embedding_dims,
tp_group=get_tp_group().cpu_group, tp_group=get_tp_group().cpu_group,
@@ -1032,7 +1035,7 @@ class MMEncoder:
) )
def _broadcast_global_cache_mask(self, mask_tensor: torch.Tensor): 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( torch.distributed.broadcast(
mask_tensor, mask_tensor,
src=0, src=0,
@@ -46,7 +46,7 @@ from sglang.srt.disaggregation.utils import (
resolve_dcp_dst_entry_indices, resolve_dcp_dst_entry_indices,
) )
from sglang.srt.environ import envs 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 from sglang.srt.server_args import ServerArgs
try: try:
@@ -405,7 +405,8 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
): ):
super().__init__(args, disaggregation_mode, server_args, is_mla_backend) super().__init__(args, disaggregation_mode, server_args, is_mla_backend)
self.transfer_source_rank = ( 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( self.kv_args.kv_data_mem_kinds = _normalize_kv_mem_kinds(
getattr(self.kv_args, "kv_data_mem_kinds", None), 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 ( from sglang.srt.runtime_context import (
get_exec, get_exec,
get_parallel, get_parallel,
get_serving,
) )
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ( 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() dist_init_method = na.to_tcp()
else: else:
dist_init_method = NetworkAddress( 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() ).to_tcp()
return dist_init_method 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 NCCL. Covers the spec-decode batch plus, if ``include_prefill``, a prefill
chunk. Returns ``floor`` if server args are unavailable.""" chunk. Returns ``floor`` if server args are unavailable."""
try: 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: schedule, spec = get_schedule(), get_spec()
v = getattr(sa, name, 0) tokens = g(schedule.max_running_requests) * max(
return v if isinstance(v, int) and v > 0 else 0 g(spec.speculative_num_draft_tokens), g(spec.speculative_eagle_topk), 1
tokens = g("max_running_requests") * max(
g("speculative_num_draft_tokens"), g("speculative_eagle_topk"), 1
) )
if include_prefill: 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) return max(tokens, floor)
except Exception: except Exception:
return floor return floor
+21 -22
View File
@@ -48,6 +48,7 @@ import torch
import uvloop import uvloop
import zmq 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.elastic_ep.expert_backup_manager import run_expert_backup_manager
from sglang.srt.entrypoints.engine_info_bootstrap_server import ( from sglang.srt.entrypoints.engine_info_bootstrap_server import (
EngineInfoBootstrapServer, 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.parser.template_manager import TemplateManager
from sglang.srt.plugins import load_plugins from sglang.srt.plugins import load_plugins
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_disagg,
get_exec, get_exec,
get_model, get_model,
get_parallel, get_parallel,
@@ -309,9 +311,9 @@ class Engine(EngineScoreMixin, EngineBase):
trace_modules=server_args.trace_modules, trace_modules=server_args.trace_modules,
) )
thread_label = "Tokenizer" thread_label = "Tokenizer"
if server_args.disaggregation_mode == "prefill": if get_disagg().disaggregation_mode == "prefill":
thread_label = "Prefill Tokenizer" thread_label = "Prefill Tokenizer"
elif server_args.disaggregation_mode == "decode": elif get_disagg().disaggregation_mode == "decode":
thread_label = "Decode Tokenizer" thread_label = "Decode Tokenizer"
trace_set_thread_info(thread_label) 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 # Needs a tokenizer and a chat template, so it cannot live in the
# pipeline; after the plugins, which may register the parser detected. # pipeline; after the plugins, which may register the parser detected.
if ( parsers = resolving_view(server_args)
server_args.reasoning_parser == "auto" if parsers.reasoning_parser == "auto" or parsers.tool_call_parser == "auto":
or server_args.tool_call_parser == "auto"
):
resolve_auto_parsers(server_args) resolve_auto_parsers(server_args)
# This publish replaces whatever was published before it, so the # 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): def _set_envs_and_config(server_args: ServerArgs):
cfg = resolving_view(server_args)
# Set global environments # Set global environments
# MNNVL fabric (GB200/GB300) multi-node: cross-node NVLink needs NCCL's # MNNVL fabric (GB200/GB300) multi-node: cross-node NVLink needs NCCL's
# cuMem-based buffers and MNNVL transport. Default them on (user-set # cuMem-based buffers and MNNVL transport. Default them on (user-set
# values win; the symm-mem override below only fires when unset). # 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_CUMEM_ENABLE", "1")
os.environ.setdefault("NCCL_MNNVL_ENABLE", "1") os.environ.setdefault("NCCL_MNNVL_ENABLE", "1")
if "NCCL_CUMEM_ENABLE" not in os.environ or 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(server_args.enable_symm_mem)) os.environ["NCCL_CUMEM_ENABLE"] = str(int(cfg.enable_symm_mem))
if ( if (
"NCCL_NVLS_ENABLE" not in os.environ "NCCL_NVLS_ENABLE" not in os.environ
or server_args.enable_nccl_nvls or cfg.enable_nccl_nvls
or server_args.enable_symm_mem or cfg.enable_symm_mem
): ):
os.environ["NCCL_NVLS_ENABLE"] = str( 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. # 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 # 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["NCCL_GRAPH_MIXING_SUPPORT"] = "0"
os.environ["CUDA_DEVICE_MAX_CONNECTIONS"] = "8" os.environ["CUDA_DEVICE_MAX_CONNECTIONS"] = "8"
@@ -1669,7 +1670,7 @@ def _set_envs_and_config(server_args: ServerArgs):
) )
# Set prometheus env vars # Set prometheus env vars
if server_args.enable_metrics: if cfg.enable_metrics:
set_prometheus_multiproc_dir() set_prometheus_multiproc_dir()
# Set ulimit # Set ulimit
@@ -1677,7 +1678,7 @@ def _set_envs_and_config(server_args: ServerArgs):
# Check flashinfer version # Check flashinfer version
if not get_bool_env_var("SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK"): 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( assert_pkg_version(
"flashinfer_python", "flashinfer_python",
"0.6.17", "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. # Signal handlers can only be registered from the main thread.
if threading.current_thread() is threading.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. # Register the signal handler.
# The child processes will send SIGQUIT to this process when any error happens # The child processes will send SIGQUIT to this process when any error happens
# This process then clean up the whole process tree # 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) signal.signal(signal.SIGQUIT, launch_phase_sigquit_handler)
else: else:
# Allow users to register a custom SIGQUIT handler for things like crash dump # Allow users to register a custom SIGQUIT handler for things like crash dump
logger.error( logger.error(f"Using custom SIGQUIT handler: {cfg.custom_sigquit_handler}")
f"Using custom SIGQUIT handler: {server_args.custom_sigquit_handler}" signal.signal(signal.SIGQUIT, cfg.custom_sigquit_handler)
)
signal.signal(signal.SIGQUIT, server_args.custom_sigquit_handler)
else: else:
logger.warning( logger.warning(
"Signal handler is not added because the engine is not in the " "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) mp.set_start_method("spawn", force=True)
# Set gc threshold # Set gc threshold
if gc_threshold := server_args.gc_threshold: if gc_threshold := cfg.gc_threshold:
gc.set_threshold(*gc_threshold) gc.set_threshold(*gc_threshold)
_log_legacy_kernel_cache_dirs() _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 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.configs.embedding_model_spec import resolved_embedding_plan
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_lora, get_lora,
@@ -416,7 +417,7 @@ class RuntimeHandle:
if embedding_model_spec is not None: if embedding_model_spec is not None:
result["embedding"] = resolved_embedding_plan( result["embedding"] = resolved_embedding_plan(
embedding_model_spec, embedding_model_spec,
server_args=self.server_args, config=resolving_view(self.server_args),
model_config=model_config, model_config=model_config,
) )
return json.dumps(result, default=str) 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." "version mismatch — see the chained exception above for details."
) from e ) 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_app = web.Application()
sidecar_runner = None sidecar_runner = None
sidecar_port = ( sidecar_port = (
@@ -176,7 +185,7 @@ async def serve_grpc(server_args, model_info=None):
# Metrics setup: must set PROMETHEUS_MULTIPROC_DIR before scheduler # Metrics setup: must set PROMETHEUS_MULTIPROC_DIR before scheduler
# processes import prometheus_client, since the env var is inherited # processes import prometheus_client, since the env var is inherited
# at fork time. # at fork time.
if server_args.enable_metrics: if cfg.enable_metrics:
try: try:
from sglang.srt.observability.func_timer import enable_func_timer from sglang.srt.observability.func_timer import enable_func_timer
from sglang.srt.utils import set_prometheus_multiproc_dir from sglang.srt.utils import set_prometheus_multiproc_dir
@@ -204,7 +213,7 @@ async def serve_grpc(server_args, model_info=None):
) )
try: try:
sidecar_runner = await _start_sidecar_server( sidecar_runner = await _start_sidecar_server(
server_args.host, sidecar_port, sidecar_app cfg.host, sidecar_port, sidecar_app
) )
except OSError as e: except OSError as e:
logger.error( logger.error(
@@ -232,7 +241,7 @@ async def serve_grpc(server_args, model_info=None):
) )
if sidecar_supported: if sidecar_supported:
serve_kwargs["on_request_manager_ready"] = _on_request_manager_ready 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 # User explicitly asked for metrics but the installed servicer can't
# start the sidecar that serves them — fail loud rather than silently # start the sidecar that serves them — fail loud rather than silently
# produce a server with no /metrics endpoint. # 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.responses import ORJSONResponse, Response, StreamingResponse
from fastapi.routing import APIRoute 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.configs.embedding_model_spec import resolved_embedding_plan
from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode 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}" thread_label = f"MultiTokenizer-{_global_state.tokenizer_manager.worker_id}"
# Add prometheus middleware # Add prometheus middleware
if server_args.enable_metrics: if get_observability().enable_metrics:
add_prometheus_middleware(app) add_prometheus_middleware(app)
enable_func_timer() enable_func_timer()
@@ -488,6 +489,7 @@ from sglang.srt.runtime_context import (
get_exec, get_exec,
get_lora, get_lora,
get_model, get_model,
get_observability,
get_parallel, get_parallel,
get_serving, get_serving,
publish, publish,
@@ -771,7 +773,7 @@ async def model_info():
if embedding_model_spec is not None: if embedding_model_spec is not None:
result["embedding"] = resolved_embedding_plan( result["embedding"] = resolved_embedding_plan(
embedding_model_spec, embedding_model_spec,
server_args=_global_state.tokenizer_manager.server_args, config=resolving_view(_global_state.tokenizer_manager.server_args),
model_config=model_config, model_config=model_config,
) )
return result return result
@@ -2540,7 +2542,7 @@ def _setup_and_run_http_server(
if tokenizer_manager is not None: if tokenizer_manager is not None:
tokenizer_manager._subprocess_watchdog = subprocess_watchdog tokenizer_manager._subprocess_watchdog = subprocess_watchdog
if server_args.enable_metrics: if get_observability().enable_metrics:
add_prometheus_track_response_middleware(app) add_prometheus_track_response_middleware(app)
# Pass additional arguments to the lifespan function. # Pass additional arguments to the lifespan function.
@@ -2602,12 +2604,13 @@ def _setup_and_run_http_server(
if server_args.enable_http2: if server_args.enable_http2:
logger.info( logger.info(
f"Starting embedded Granian HTTP/2 server on " 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( _run_granian_server(
host=server_args.host, host=get_serving().host,
port=server_args.port, port=get_serving().port,
log_level=server_args.log_level_http or server_args.log_level, log_level=get_observability().log_level_http
or get_observability().log_level,
http2_max_concurrent_streams=( http2_max_concurrent_streams=(
server_args.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. # Use Config/Server API for access to the SSLContext.
config = uvicorn.Config( config = uvicorn.Config(
app, app,
host=server_args.host, host=get_serving().host,
port=server_args.port, port=get_serving().port,
root_path=server_args.fastapi_root_path, 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_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
loop="uvloop", loop="uvloop",
ssl_keyfile=server_args.ssl_keyfile, ssl_keyfile=server_args.ssl_keyfile,
@@ -2658,10 +2662,11 @@ def _setup_and_run_http_server(
# Default case, one tokenizer process # Default case, one tokenizer process
uvicorn.run( uvicorn.run(
app, app,
host=server_args.host, host=get_serving().host,
port=server_args.port, port=get_serving().port,
root_path=server_args.fastapi_root_path, 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_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
loop="uvloop", loop="uvloop",
ssl_keyfile=server_args.ssl_keyfile, ssl_keyfile=server_args.ssl_keyfile,
@@ -2689,12 +2694,13 @@ def _setup_and_run_http_server(
if server_args.enable_http2: if server_args.enable_http2:
logger.info( logger.info(
f"Starting embedded Granian HTTP/2 server on " 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( _run_granian_server(
host=server_args.host, host=get_serving().host,
port=server_args.port, port=get_serving().port,
log_level=server_args.log_level_http or server_args.log_level, log_level=get_observability().log_level_http
or get_observability().log_level,
http2_max_concurrent_streams=( http2_max_concurrent_streams=(
server_args.http2_max_concurrent_streams server_args.http2_max_concurrent_streams
), ),
@@ -2707,10 +2713,11 @@ def _setup_and_run_http_server(
else: else:
uvicorn.run( uvicorn.run(
"sglang.srt.entrypoints.http_server:app", "sglang.srt.entrypoints.http_server:app",
host=server_args.host, host=get_serving().host,
port=server_args.port, port=get_serving().port,
root_path=server_args.fastapi_root_path, 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_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
timeout_worker_healthcheck=envs.SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT.get(), timeout_worker_healthcheck=envs.SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT.get(),
loop="uvloop", loop="uvloop",
@@ -2748,12 +2755,12 @@ def _start_native_grpc_server_for_runtime(
) )
grpc_handle = grpc_native.start_server( grpc_handle = grpc_native.start_server(
host=server_args.host, host=get_serving().host,
port=grpc_port, port=grpc_port,
runtime_handle=runtime_handle, runtime_handle=runtime_handle,
worker_threads=server_args.grpc_worker_threads, 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 return grpc_handle
+1 -1
View File
@@ -118,7 +118,7 @@ def start_sidecar(server_args) -> Sidecar:
module_name = server_args.sidecar module_name = server_args.sidecar
assert module_name is not None assert module_name is not None
sidecar_args, shutdown_timeout = _parse_sidecar_args(server_args.sidecar_args) 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( proc = mp.get_context("spawn").Process(
name=f"sglang_sidecar_{module_name}", name=f"sglang_sidecar_{module_name}",
target=_run_sidecar, target=_run_sidecar,
+1 -1
View File
@@ -146,7 +146,7 @@ async def get_loads(
"version": __version__, "version": __version__,
"accelerator": _accelerator_name(), "accelerator": _accelerator_name(),
"num_accelerators": _num_accelerators_per_dp_rank( "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.pp_size,
get_parallel().config.dp_size, get_parallel().config.dp_size,
get_parallel().config.enable_dp_attention, get_parallel().config.enable_dp_attention,
@@ -764,7 +764,7 @@ class _UtilizationRateAccumulatorMixin(_Accumulator):
single_pass_global_physical_count, single_pass_global_physical_count,
num_gpu=self._expert_location_metadata.ep_size, 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( torch.distributed.reduce(
gpu_physical_count, dst=0, op=torch.distributed.ReduceOp.SUM 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 # Cannot use local_physical_count to support select_experts
self._expert_location_metadata.num_physical_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, dtype=torch.int32,
device=self._server_args.device, device=get_device_namespace().device,
) )
self._first_dump = True self._first_dump = True
@@ -22,6 +22,7 @@ import torch
from torch.profiler import ProfilerActivity, profile from torch.profiler import ProfilerActivity, profile
from sglang.srt.model_executor.runner import DecodeCudaGraphRunner 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 from sglang.srt.utils import register_xpu_device_properties_for_dynamo
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -116,7 +117,7 @@ class XPUGraphRunner(DecodeCudaGraphRunner):
def __init__(self, model_runner: ModelRunner): def __init__(self, model_runner: ModelRunner):
assert ( assert (
not model_runner.server_args.enable_memory_saver not get_exec().features.enable_memory_saver
), "XPUGraphRunner does not support Torch Memory Saver yet." ), "XPUGraphRunner does not support Torch Memory Saver yet."
register_fake_ops() register_fake_ops()
self._apply_xpu_compile_config() self._apply_xpu_compile_config()
@@ -112,10 +112,10 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
self.device = model_runner.device self.device = model_runner.device
self.max_context_len = model_runner.model_config.context_len self.max_context_len = model_runner.model_config.context_len
self.num_heads = model_runner.model_config.get_num_attention_heads( 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( 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 self.head_size = model_runner.model_config.head_dim
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Optional, Tuple
import torch import torch
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.configs.model_config import ( from sglang.srt.configs.model_config import (
get_minimax_sparse_attention_config, get_minimax_sparse_attention_config,
get_minimax_sparse_disable_value_layer_ids, get_minimax_sparse_disable_value_layer_ids,
@@ -116,7 +117,9 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
self.max_context_len = int(runner.model_config.context_len) self.max_context_len = int(runner.model_config.context_len)
# Per-forward cache for the native decode block table (rebuilt each forward). # Per-forward cache for the native decode block table (rebuilt each forward).
self._native_decode_bt: dict = {} 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: if self.fp8_attn_gemm:
assert self.kv_pool.main_pool.dtype == torch.float8_e4m3fn, ( assert self.kv_pool.main_pool.dtype == torch.float8_e4m3fn, (
"fp8 attn-GEMM mode requires an fp8_e4m3fn main KV pool, got " "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: if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.server_args import ServerArgs
class ContextParallelStrategyKind(IntEnum): class ContextParallelStrategyKind(IntEnum):
@@ -237,23 +236,27 @@ def _is_dsa_active() -> bool:
_STRATEGY: Optional[ContextParallelStrategy] = None _STRATEGY: Optional[ContextParallelStrategy] = None
def init_cp_strategy(server_args: ServerArgs) -> None: def init_cp_strategy(
"""Bind the configured CP strategy for this process.""" *, enable_prefill_cp: bool, cp_size: int, cp_strategy: str
from sglang.srt.arg_groups.overrides import resolving_view ) -> 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 global _STRATEGY
if not cfg.enable_prefill_cp: if not enable_prefill_cp:
_STRATEGY = None _STRATEGY = None
return return
cp_size = cfg.attn_cp_size
if cp_size <= 1: if cp_size <= 1:
_STRATEGY = None _STRATEGY = None
return return
kind = ContextParallelStrategyKind.from_string(cfg.cp_strategy) kind = ContextParallelStrategyKind.from_string(cp_strategy)
if kind == ContextParallelStrategyKind.ZIGZAG: if kind == ContextParallelStrategyKind.ZIGZAG:
from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy 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) _STRATEGY = InterleaveCPStrategy(cp_size=cp_size)
else: else:
raise ValueError( raise ValueError(
f"Unsupported cp_strategy kind {kind} for " f"Unsupported cp_strategy kind {kind} for cp_strategy={cp_strategy!r}"
f"cp_strategy={cfg.cp_strategy!r}"
) )
@@ -280,14 +282,16 @@ def get_cp_strategy() -> Optional[ContextParallelStrategy]:
global _STRATEGY global _STRATEGY
if _STRATEGY is None: if _STRATEGY is None:
from sglang.srt.runtime_context import get_server_args
try: try:
server_args = get_server_args() parallel = get_parallel().config
except ValueError: except ValueError:
return None return None
if server_args is not None and get_parallel().config.enable_prefill_cp: if parallel.enable_prefill_cp:
init_cp_strategy(server_args) init_cp_strategy(
enable_prefill_cp=True,
cp_size=parallel.attn_cp_size,
cp_strategy=parallel.cp_strategy,
)
return _STRATEGY return _STRATEGY
+1 -1
View File
@@ -63,7 +63,7 @@ def is_glm_dsa_cache_layer_split_enabled(model_runner: "ModelRunner") -> bool:
return ( return (
not model_runner.is_draft_worker 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 model_runner.use_mla_backend
and is_deepseek_dsa(model_runner.model_config.hf_config) and is_deepseek_dsa(model_runner.model_config.hf_config)
) )
@@ -13,7 +13,7 @@ from sglang.srt.distributed import (
get_tp_group, get_tp_group,
) )
from sglang.srt.distributed.parallel_state import in_the_same_node_as 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 ( from sglang.srt.utils import (
ceil_align, ceil_align,
get_cuda_driver_bindings, get_cuda_driver_bindings,
@@ -75,12 +75,16 @@ def _resolve_backend(backend: str, is_multi_node: bool = False) -> str:
return backend return backend
def resolve_flashinfer_allreduce_fusion_backend(server_args) -> Optional[str]: def resolve_flashinfer_allreduce_fusion_backend() -> Optional[str]:
backend = getattr(server_args, "flashinfer_allreduce_fusion_backend", None) """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: if backend is None:
return None return None
is_multi_node = getattr(server_args, "nnodes", 1) > 1 return _resolve_backend(backend, get_parallel().config.nnodes > 1)
return _resolve_backend(backend, is_multi_node)
if is_flashinfer_available(): if is_flashinfer_available():
@@ -716,8 +720,7 @@ def ensure_workspace_initialized(
token_num = token_num or max_token_num token_num = token_num or max_token_num
group_key = (device_group, cpu_group) group_key = (device_group, cpu_group)
effective_dtype = dtype or torch.bfloat16 effective_dtype = dtype or torch.bfloat16
server_args = get_server_args() backend = resolve_flashinfer_allreduce_fusion_backend()
backend = resolve_flashinfer_allreduce_fusion_backend(server_args)
if backend is None: if backend is None:
return False return False
+34 -21
View File
@@ -4,7 +4,6 @@ import logging
import os import os
from contextlib import contextmanager from contextlib import contextmanager
from enum import Enum, IntEnum from enum import Enum, IntEnum
from typing import TYPE_CHECKING
import torch import torch
@@ -12,14 +11,18 @@ from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
is_dp_attention_enabled, 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 from sglang.srt.utils import is_cuda, is_npu
_is_npu = 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.runtime_context import get_server_args
from sglang.srt.utils.common import log_info_on_rank0 from sglang.srt.utils.common import log_info_on_rank0
@@ -308,37 +311,47 @@ def get_ascend_dispatcher_output_dtype(dispatcher):
return DispatcherOutputDtype.BF16 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 = get_flags().moe
moe.a2a_backend = MoeA2ABackend(server_args.moe_a2a_backend) moe.a2a_backend = MoeA2ABackend(exec_moe.moe_a2a_backend)
moe.runner_backend = MoeRunnerBackend(server_args.moe_runner_backend) moe.runner_backend = MoeRunnerBackend(exec_moe.moe_runner_backend)
moe.speculative_runner_backend = ( moe.speculative_runner_backend = (
MoeRunnerBackend(server_args.speculative_moe_runner_backend) MoeRunnerBackend(spec.speculative_moe_runner_backend)
if server_args.speculative_moe_runner_backend is not None if spec.speculative_moe_runner_backend is not None
else moe.runner_backend else moe.runner_backend
) )
moe.speculative_a2a_backend = ( moe.speculative_a2a_backend = (
MoeA2ABackend(server_args.speculative_moe_a2a_backend) MoeA2ABackend(spec.speculative_moe_a2a_backend)
if server_args.speculative_moe_a2a_backend is not None if spec.speculative_moe_a2a_backend is not None
else moe.a2a_backend else moe.a2a_backend
) )
moe.deepep_mode = DeepEPMode(server_args.deepep_mode) moe.deepep_mode = DeepEPMode(exec_moe.deepep_mode)
moe.deepep_config = server_args.deepep_config or "" moe.deepep_config = exec_moe.deepep_config or ""
moe.tbo_enabled = server_args.enable_two_batch_overlap moe.tbo_enabled = overlap.enable_two_batch_overlap
moe.sbo_enabled = server_args.enable_single_batch_overlap moe.sbo_enabled = overlap.enable_single_batch_overlap
if moe.sbo_enabled and is_cuda(): if moe.sbo_enabled and is_cuda():
if torch.cuda.get_device_capability()[0] == 9: if torch.cuda.get_device_capability()[0] == 9:
raise ValueError( 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." "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.tbo_token_distribution_threshold = overlap.tbo_token_distribution_threshold
moe.disable_fp4_allgather = server_args.disable_flashinfer_cutlass_moe_fp4_allgather moe.disable_fp4_allgather = exec_moe.disable_flashinfer_cutlass_moe_fp4_allgather
moe.quantization = server_args.quantization moe.quantization = get_model().quantization
# Seeded with the user's intent; each model's gate refines the ACTIVE # Seeded with the user's intent; each model's gate refines the ACTIVE
# value for its own build (install_shared_experts_fusion_decision). # 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 = ( 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 import logging
from enum import Enum from enum import Enum
from typing import TYPE_CHECKING, Optional from typing import Optional
import torch import torch
from sglang.srt.runtime_context import get_exec
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
get_device_capability, get_device_capability,
is_cuda, is_cuda,
@@ -13,9 +14,6 @@ from sglang.srt.utils.common import (
) )
from sglang.srt.utils.custom_op import register_custom_op_from_extern 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__) logger = logging.getLogger(__name__)
@@ -142,11 +140,11 @@ class Fp4GemmRunnerBackend(Enum):
FP4_GEMM_RUNNER_BACKEND: Fp4GemmRunnerBackend | None = None FP4_GEMM_RUNNER_BACKEND: Fp4GemmRunnerBackend | None = None
def initialize_fp4_gemm_config(server_args: ServerArgs) -> None: def initialize_fp4_gemm_config() -> None:
"""Initialize FP4 GEMM configuration from server args.""" """Initialize the FP4 GEMM backend from the published configuration."""
global FP4_GEMM_RUNNER_BACKEND global FP4_GEMM_RUNNER_BACKEND
backend = server_args.fp4_gemm_runner_backend backend = get_exec().kernel.fp4_gemm_runner_backend
if backend == "auto": if backend == "auto":
if is_sm100_supported(): if is_sm100_supported():
backend = "flashinfer_cutedsl" backend = "flashinfer_cutedsl"
@@ -3,23 +3,10 @@ from __future__ import annotations
import logging import logging
from enum import Enum from enum import Enum
from functools import lru_cache, partial 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 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 ( from sglang.kernels.ops.quantization.fp8_kernel import (
fp8_dtype, fp8_dtype,
fp8_max, fp8_max,
@@ -28,12 +15,18 @@ from sglang.kernels.ops.quantization.fp8_kernel import (
is_fp8_fnuz, is_fp8_fnuz,
per_token_group_quant_fp8, per_token_group_quant_fp8,
scaled_fp8_quant, scaled_fp8_quant,
sglang_per_token_group_quant_fp8,
sglang_per_token_group_quant_fp8_row_padded,
sglang_per_token_quant_fp8, sglang_per_token_quant_fp8,
static_quant_fp8, static_quant_fp8,
triton_scaled_mm, triton_scaled_mm,
w8a8_block_fp8_matmul_deepgemm, w8a8_block_fp8_matmul_deepgemm,
w8a8_block_fp8_matmul_triton, 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 ( from sglang.srt.utils import (
ceil_align, ceil_align,
ceil_div, ceil_div,
@@ -54,6 +47,7 @@ from sglang.srt.utils import (
is_xpu, is_xpu,
offloader, offloader,
) )
from sglang.srt.utils.common import torch_release
from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.custom_op import register_custom_op
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -799,11 +793,11 @@ def _dispatch_auto_backend() -> Callable:
return triton_w8a8_block_fp8_linear 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.""" """Initialize FP8 GEMM configuration."""
global FP8_GEMM_RUNNER_BACKEND 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(): if backend == "auto" and is_sm120_supported():
backend = "cutlass" 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.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread 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.server_args import PortArgs, ServerArgs
from sglang.srt.utils import configure_logger, freeze_gc, kill_itself_when_parent_died from sglang.srt.utils import configure_logger, freeze_gc, kill_itself_when_parent_died
from sglang.srt.utils.hf_transformers_utils import get_tokenizer from sglang.srt.utils.hf_transformers_utils import get_tokenizer
@@ -130,7 +136,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
self.tokenizer = get_tokenizer( self.tokenizer = get_tokenizer(
get_serving().tokenizer_path, get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode, 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, revision=server_args.revision,
tokenizer_backend=server_args.tokenizer_backend, tokenizer_backend=server_args.tokenizer_backend,
) )
@@ -151,7 +157,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(), test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(),
) )
if server_args.enable_metrics: if get_observability().enable_metrics:
start_cpu_monitor_thread("detokenizer") start_cpu_monitor_thread("detokenizer")
def init_request_dispatcher(self): 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 import msgspec
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import TokenizedGenerateReqInput from sglang.srt.managers.io_struct import TokenizedGenerateReqInput
from sglang.srt.managers.utils import ( from sglang.srt.managers.utils import (
@@ -29,6 +30,8 @@ from sglang.srt.managers.utils import (
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_mm, get_mm,
get_observability,
get_parallel,
get_serving, get_serving,
) )
from sglang.srt.utils.flatten import ( from sglang.srt.utils.flatten import (
@@ -272,7 +275,7 @@ class NativeMmHost:
) )
return ( return (
self.server_args.tp_size > 1 get_parallel().config.tp_size > 1
and determine_tensor_transport_mode() != "default" and determine_tensor_transport_mode() != "default"
and not self.server_args.skip_tokenizer_init and not self.server_args.skip_tokenizer_init
) )
@@ -395,13 +398,13 @@ class RustServer:
"ingress has no equivalent). Launch without SGLANG_RUST_SERVER, or " "ingress has no equivalent). Launch without SGLANG_RUST_SERVER, or "
"drop --preferred-sampling-params and send those values per request." "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, # 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. # 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 dp_rank = scheduler.ps.attn_dp_rank if scheduler.ps.dp_size > 1 else None
if dp_rank is not 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( launch_cores, server_cores = cls._partition_cores(
mm_workers=( mm_workers=(
@@ -754,7 +757,7 @@ class RustServer:
ext = load_rust_extension("sglang.srt.rust_extensions._server") ext = load_rust_extension("sglang.srt.rust_extensions._server")
sa = scheduler.server_args sa = resolving_view(scheduler.server_args)
mc = scheduler.model_config mc = scheduler.model_config
disaggregation_mode = { disaggregation_mode = {
"null": ext.DisaggregationMode.Null, "null": ext.DisaggregationMode.Null,
@@ -768,9 +771,9 @@ class RustServer:
revision=sa.revision, revision=sa.revision,
load_format=sa.load_format, load_format=sa.load_format,
weight_version=sa.weight_version, weight_version=sa.weight_version,
host=sa.host, host=get_serving().host,
port=sa.port, port=sa.port,
log_level=sa.log_level, log_level=get_observability().log_level,
log_level_http=sa.log_level_http, log_level_http=sa.log_level_http,
chat_template=sa.chat_template, chat_template=sa.chat_template,
tool_call_parser=sa.tool_call_parser, tool_call_parser=sa.tool_call_parser,
+3 -3
View File
@@ -903,11 +903,11 @@ class Scheduler(
"moe_topk", "moe_topk",
) )
if any(hasattr(config_to_check, attr) for attr in moe_topk_attrs): 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 GEMM-related configuration for FP8 and FP4 backends.
initialize_fp8_gemm_config(self.server_args) initialize_fp8_gemm_config()
initialize_fp4_gemm_config(self.server_args) initialize_fp4_gemm_config()
initialize_bf16_gemm_config(self.server_args) initialize_bf16_gemm_config(self.server_args)
# This must be called after initialize_moe_config # This must be called after initialize_moe_config
@@ -179,8 +179,8 @@ class TokenizerControlMixin:
) )
if primary_group_control: if primary_group_control:
control_fan_out = ( control_fan_out = (
worker_count + self.server_args.tp_size - 1 worker_count + get_parallel().config.tp_size - 1
) // self.server_args.tp_size ) // get_parallel().config.tp_size
else: else:
control_fan_out = worker_count control_fan_out = worker_count
@@ -3591,7 +3591,7 @@ def get_processor_wrapper(server_args):
tokenizer_mode=get_serving().tokenizer_mode, tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=get_model().trust_remote_code, trust_remote_code=get_model().trust_remote_code,
revision=get_model().revision, 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, tokenizer_backend=get_serving().tokenizer_backend,
model_name=get_model().model_path, model_name=get_model().model_path,
) )
+1 -1
View File
@@ -86,7 +86,7 @@ class HiRadixCache(RadixCache):
self.page_size = params.page_size self.page_size = params.page_size
self.kv_cache = params.token_to_kv_pool_allocator.get_kvcache() 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): if isinstance(self.kv_cache, MHATokenToKVPool):
self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)( 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__) logger = logging.getLogger(__name__)
def _get_allocator_type(server_args: ServerArgs) -> str: def _get_allocator_type() -> str:
return get_allocator_type(server_args) return get_allocator_type()
def _evict_swa_for_device_alloc(cache: UnifiedRadixCache, required_size: int) -> None: 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, get_memory().hicache_size if host_size is None else host_size,
page_size, page_size,
get_memory().hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
pool_label=pool_label, pool_label=pool_label,
**kwargs, **kwargs,
) )
@@ -540,7 +540,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=swa_num_host_pages, num_host_pages=swa_num_host_pages,
slot_page_size=kvcache.swa_page_size, slot_page_size=kvcache.swa_page_size,
layout=get_memory().hicache_mem_layout, 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 swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator
entries.append( entries.append(
@@ -567,7 +567,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=num_host_pages, num_host_pages=num_host_pages,
slot_page_size=page_size, slot_page_size=page_size,
layout=get_memory().hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
) )
c4_indexer_host_pool = DeepSeekV4PagedHostPool( c4_indexer_host_pool = DeepSeekV4PagedHostPool(
pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER), pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER),
@@ -579,7 +579,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=num_host_pages, num_host_pages=num_host_pages,
slot_page_size=page_size, slot_page_size=page_size,
layout=get_memory().hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
) )
entries.extend( entries.extend(
[ [
@@ -610,7 +610,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=swa_num_host_pages, num_host_pages=swa_num_host_pages,
swa_page_size=kvcache.swa_page_size, swa_page_size=kvcache.swa_page_size,
layout=get_memory().hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
) )
c4_indexer_state_host_pool = DeepSeekV4StateHostPool( c4_indexer_state_host_pool = DeepSeekV4StateHostPool(
pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER_STATE), 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, num_host_pages=swa_num_host_pages,
swa_page_size=kvcache.swa_page_size, swa_page_size=kvcache.swa_page_size,
layout=get_memory().hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
) )
entries.extend( entries.extend(
[ [
@@ -653,7 +653,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=num_host_pages, num_host_pages=num_host_pages,
slot_page_size=page_size, slot_page_size=page_size,
layout=get_memory().hicache_mem_layout, 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. # C128 state pool is intentionally not registered with hicache.
# page_size=256 % 128 == 0, so state pool is not consumed on load. # page_size=256 % 128 == 0, so state pool is not consumed on load.
@@ -739,7 +739,7 @@ def build_hybrid_mamba_stack(
mamba_pool, mamba_pool,
get_memory().hicache_ratio, get_memory().hicache_ratio,
mamba_host_size, mamba_host_size,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
layout=get_memory().hicache_mem_layout, layout=get_memory().hicache_mem_layout,
) )
entries = [ entries = [
@@ -1038,7 +1038,7 @@ def build_full_draft_pools(
host_to_device_ratio=host_pool_group.logical_size / pool.size, host_to_device_ratio=host_pool_group.logical_size / pool.size,
page_size=controller.page_size, page_size=controller.page_size,
layout=get_memory().hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
pool_label="draft", pool_label="draft",
) )
draft_layer_mapping = {i: i for i in range(pool.layer_num)} draft_layer_mapping = {i: i for i in range(pool.layer_num)}
@@ -1064,7 +1064,7 @@ def build_full_draft_pools(
pool, pool,
draft_host_pool, draft_host_pool,
get_memory().hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
) )
specs.append( specs.append(
SidecarPoolSpec( SidecarPoolSpec(
@@ -1111,7 +1111,7 @@ def build_swa_draft_pools(
num_host_pages=target_swa_host_pool.num_host_pages, num_host_pages=target_swa_host_pool.num_host_pages,
slot_page_size=draft_swa_pool.page_size, slot_page_size=draft_swa_pool.page_size,
layout=target_swa_host_pool.layout, layout=target_swa_host_pool.layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
) )
else: else:
host_pool = _build_mha_mla_host_pool( 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, host_to_device_ratio=target_swa_host_pool.size / draft_swa_pool.size,
page_size=target_swa_host_pool.page_size, page_size=target_swa_host_pool.page_size,
layout=target_swa_host_pool.layout, layout=target_swa_host_pool.layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
pool_label="draft_swa", pool_label="draft_swa",
) )
@@ -1513,7 +1513,7 @@ class _DsaStrategy(StackStrategy):
full_kv_pool, full_kv_pool,
kv_host_pool, kv_host_pool,
get_memory().hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
), ),
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=model_name, model_name=model_name,
@@ -1951,7 +1951,7 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
kv, kv,
kv_host_pool, kv_host_pool,
get_memory().hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(),
), ),
model_name=get_serving().served_model_name, model_name=get_serving().served_model_name,
storage_backend_extra_config=extra_config, storage_backend_extra_config=extra_config,
@@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Any, Optional
import msgspec import msgspec
import torch import torch
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.configs.hybrid_arch import ( from sglang.srt.configs.hybrid_arch import (
hybrid_gdn_config, hybrid_gdn_config,
kimi_linear_config, kimi_linear_config,
@@ -1262,7 +1263,7 @@ class KVCacheConfigurator:
sparse_layer_ids=sparse_layer_ids, sparse_layer_ids=sparse_layer_ids,
disable_value_sparse_layer_ids=disable_value_sparse_layer_ids, disable_value_sparse_layer_ids=disable_value_sparse_layer_ids,
device=self.device, 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, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
) )
@@ -1533,7 +1534,7 @@ class KVCacheConfigurator:
# with the widening-dequant contract. # with the widening-dequant contract.
index_dtype=( index_dtype=(
self.kv_cache_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 else self.model_dtype
), ),
head_num=self.model_config.get_num_kv_heads( head_num=self.model_config.get_num_kv_heads(
@@ -99,14 +99,15 @@ def get_allocator_from_storage(allocator_type):
return HostTensorAllocator() return HostTensorAllocator()
def get_allocator_type(server_args) -> str: def get_allocator_type() -> str:
backend = getattr(server_args, "hicache_storage_backend", None) """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": if backend == "shm":
return "shm" return "shm"
if backend == "dynamic": if backend == "dynamic":
extra_config_str = getattr( extra_config_str = get_memory().hicache_storage_backend_extra_config
server_args, "hicache_storage_backend_extra_config", None
)
if extra_config_str: if extra_config_str:
try: try:
config = json.loads(extra_config_str) config = json.loads(extra_config_str)
@@ -81,6 +81,7 @@ from sglang.srt.observability.metrics_collector import (
StorageMetrics, StorageMetrics,
StorageMetricsCollector, StorageMetricsCollector,
) )
from sglang.srt.runtime_context import get_memory
from sglang.srt.session.streaming_session import StreamingSession from sglang.srt.session.streaming_session import StreamingSession
from sglang.srt.utils.common import ceil_align from sglang.srt.utils.common import ceil_align
@@ -385,7 +386,7 @@ class UnifiedRadixCache(BasePrefixCache):
self.extra_metric_labels = server_args.extra_metric_labels self.extra_metric_labels = server_args.extra_metric_labels
# Parse storage config once, share with assembler and tree # 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_extra_config = None
storage_prefetch_threshold = 256 storage_prefetch_threshold = 256
prefetch_timeout_base = 1.0 prefetch_timeout_base = 1.0
@@ -399,7 +400,7 @@ class UnifiedRadixCache(BasePrefixCache):
prefetch_timeout_per_ki_token, prefetch_timeout_per_ki_token,
hicache_storage_pass_prefix_keys, hicache_storage_pass_prefix_keys,
) = HybridCacheController.parse_storage_backend_extra_config( ) = 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( attach_hybrid_pool_to_unified_cache(
@@ -442,7 +443,7 @@ class UnifiedRadixCache(BasePrefixCache):
# State initialization # State initialization
self.write_through_threshold = ( 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.is_write_back = (
self.cache_controller is not None self.cache_controller is not None
@@ -457,7 +458,7 @@ class UnifiedRadixCache(BasePrefixCache):
pool=_COMPONENT_POOL_LABEL[ct], pool=_COMPONENT_POOL_LABEL[ct],
) )
self.load_back_threshold = 10 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). # Runtime attach/detach of the L3 backend (startup, admin API, atexit).
self._storage_attachment = StorageAttachment(self) self._storage_attachment = StorageAttachment(self)
@@ -609,7 +609,7 @@ class CPUGraphRunner:
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = (
model_runner.server_args.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.dp_size = get_parallel().config.dp_size
self.pp_size = get_parallel().config.pp_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 mindspore.communication import create_group
from sglang.srt.distributed.parallel_state import _groups from sglang.srt.distributed.parallel_state import _groups
from sglang.srt.runtime_context import get_serving
logger = logging.getLogger(__name__) 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: if server_args.dist_init_addr:
dist_init_method = f"tcp://{server_args.dist_init_addr}" dist_init_method = f"tcp://{server_args.dist_init_addr}"
else: 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) set_ms_parallel_env(rank, local_rank, world_size, dist_init_method)
ms.set_context(infer_boost="on", jit_level="O0") ms.set_context(infer_boost="on", jit_level="O0")
@@ -1202,7 +1202,6 @@ class ModelRunner:
# Pre-expand RoPE cache before CUDA Graph capture # Pre-expand RoPE cache before CUDA Graph capture
reserve_rope_cache_for_long_sequences( reserve_rope_cache_for_long_sequences(
self.model, self.model,
self.server_args,
self.model_config, self.model_config,
logger, 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.hash_topk import HashTopK
from sglang.srt.layers.moe.topk import TopK 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 from sglang.srt.utils import get_bool_env_var, is_hip, log_info_on_rank0
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -54,7 +55,7 @@ def prepare_moe_topk(
# Redundant experts therefore need to be included in the per-rank # Redundant experts therefore need to be included in the per-rank
# expert count used for Waterfill's shared-expert slot remapping. # expert count used for Waterfill's shared-expert slot remapping.
num_physical_routed_experts = ( 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): if isinstance(module, TopK):
routed_scaling_factor = module.topk_config.routed_scaling_factor routed_scaling_factor = module.topk_config.routed_scaling_factor
@@ -216,7 +216,7 @@ class BaseRunner(ABC):
self.model_runner = model_runner self.model_runner = model_runner
self.device = model_runner.device self.device = model_runner.device
self.device_module = torch.get_device_module(self.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 # elastic-EP scale-up rewrites dp_size on the published config
self.dp_size = get_parallel().config.dp_size self.dp_size = get_parallel().config.dp_size
self.pp_size = get_parallel().config.pp_size self.pp_size = get_parallel().config.pp_size
@@ -64,7 +64,7 @@ def resolve_decode_backend(
cfg = get_exec().graph.cuda_graph_config cfg = get_exec().graph.cuda_graph_config
backend_name = cfg.decode.backend if cfg is not None else Backend.FULL 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": if model_runner.device == "npu":
from sglang.srt.hardware_backend.npu.graph_runner.npu_cudagraph_backend import ( 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: if backend_name == Backend.BREAKABLE:
return BreakableCudaGraphBackend( return BreakableCudaGraphBackend(
cuda_graph_runner, 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, debug_eager=get_exec().graph.debug_cuda_graph,
) )
if backend_name == Backend.FULL: if backend_name == Backend.FULL:
return FullCudaGraphBackend( return FullCudaGraphBackend(
cuda_graph_runner, 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. # Default: tc_piecewise.
return TcPiecewiseCudaGraphBackend(cuda_graph_runner) return TcPiecewiseCudaGraphBackend(cuda_graph_runner)
@@ -88,7 +88,7 @@ class RayDataParallelController(DataParallelController):
dp_port_args_list.append(tmp_port_args) dp_port_args_list.append(tmp_port_args)
# Create ZMQ PUSH socket for this DP rank (controller → scheduler) # 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.workers[dp_rank] = get_zmq_socket(
self.context, self.context,
zmq.PUSH, zmq.PUSH,
@@ -139,7 +139,7 @@ class RayDataParallelController(DataParallelController):
dp_rank: DP rank for regular DP; None for DP attention (derived from tp_rank). 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. 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) batch_start_idx = len(self.scheduler_actors)
if not self.is_custom_pg: 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( pp_range, tp_range, pp_per_node, tp_per_node = _calculate_rank_ranges(
nnodes, nnodes,
get_parallel().config.pp_size, get_parallel().config.pp_size,
server_args.tp_size, get_parallel().config.tp_size,
node_rank=node_idx, node_rank=node_idx,
) )
for pp_rank in pp_range: for pp_rank in pp_range:
@@ -160,13 +160,14 @@ class RayDataParallelController(DataParallelController):
tp_rank % tp_per_node 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( _, _, actual_dp_rank, _ = compute_dp_attention_world_info(
get_parallel().config.enable_dp_attention, parallel.enable_dp_attention,
tp_rank, tp_rank,
server_args.tp_size, parallel.tp_size,
get_parallel().config.dp_size, parallel.dp_size,
get_parallel().config.attn_cp_size, parallel.attn_cp_size,
) )
rank_port_args = PortArgs.init_new( rank_port_args = PortArgs.init_new(
server_args, actual_dp_rank, worker_ports server_args, actual_dp_rank, worker_ports
@@ -204,10 +205,11 @@ class RayDataParallelController(DataParallelController):
self.scheduler_actors.append(actor) self.scheduler_actors.append(actor)
else: else:
world_size = _compute_world_size(server_args) world_size = _compute_world_size()
bundle_indices = _resolve_bundle_indices(self.pg, 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: if dp_rank is not None:
start_rank = dp_rank * ranks_per_tp_group start_rank = dp_rank * ranks_per_tp_group
end_rank = start_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): for global_rank in range(start_rank, end_rank):
local_rank = global_rank % ranks_per_tp_group local_rank = global_rank % ranks_per_tp_group
pp_rank = local_rank // server_args.tp_size pp_rank = local_rank // parallel.tp_size
tp_rank = local_rank % server_args.tp_size tp_rank = local_rank % parallel.tp_size
rank_port_args = port_args rank_port_args = port_args
actual_dp_rank = dp_rank actual_dp_rank = dp_rank
@@ -235,7 +237,7 @@ class RayDataParallelController(DataParallelController):
_, _, actual_dp_rank, _ = compute_dp_attention_world_info( _, _, actual_dp_rank, _ = compute_dp_attention_world_info(
get_parallel().config.enable_dp_attention, get_parallel().config.enable_dp_attention,
tp_rank, tp_rank,
server_args.tp_size, get_parallel().config.tp_size,
get_parallel().config.dp_size, get_parallel().config.dp_size,
get_parallel().config.attn_cp_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). """Compute world_size (total number of scheduler actors/GPUs needed).
Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size. 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: parallel = get_parallel().config
return server_args.tp_size * get_parallel().config.pp_size if parallel.enable_dp_attention:
return ( return parallel.tp_size * parallel.pp_size
get_parallel().config.dp_size return parallel.dp_size * parallel.tp_size * parallel.pp_size
* server_args.tp_size
* get_parallel().config.pp_size
)
def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]: def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]:
@@ -274,16 +273,13 @@ class RayEngine(Engine):
placement_group as create_placement_group, placement_group as create_placement_group,
) )
if get_parallel().config.enable_dp_attention: parallel = get_parallel().config
total_gpus = server_args.tp_size * get_parallel().config.pp_size if parallel.enable_dp_attention:
total_gpus = parallel.tp_size * parallel.pp_size
else: else:
total_gpus = ( total_gpus = parallel.dp_size * parallel.tp_size * parallel.pp_size
get_parallel().config.dp_size
* server_args.tp_size
* get_parallel().config.pp_size
)
nnodes = server_args.nnodes nnodes = parallel.nnodes
gpus_per_node = total_gpus // nnodes gpus_per_node = total_gpus // nnodes
strategy = "STRICT_PACK" if nnodes == 1 else "SPREAD" strategy = "STRICT_PACK" if nnodes == 1 else "SPREAD"
@@ -300,8 +296,8 @@ class RayEngine(Engine):
ray.get(pg.ready()) ray.get(pg.ready())
is_custom_pg = placement_group is not None is_custom_pg = placement_group is not None
nnodes = server_args.nnodes nnodes = get_parallel().config.nnodes
world_size = _compute_world_size(server_args) world_size = _compute_world_size()
if not is_custom_pg: if not is_custom_pg:
engine_bundle, engine_ip = _find_engine_bundle(pg, nnodes) engine_bundle, engine_ip = _find_engine_bundle(pg, nnodes)
@@ -341,7 +337,7 @@ class RayEngine(Engine):
_calculate_rank_ranges( _calculate_rank_ranges(
nnodes, nnodes,
get_parallel().config.pp_size, get_parallel().config.pp_size,
server_args.tp_size, get_parallel().config.tp_size,
node_rank=node_idx, node_rank=node_idx,
) )
) )
@@ -377,9 +373,10 @@ class RayEngine(Engine):
f"bundle_indices={bundle_indices}" f"bundle_indices={bundle_indices}"
) )
tp_size = get_parallel().config.tp_size
for rank in range(world_size): for rank in range(world_size):
pp_rank = rank // server_args.tp_size pp_rank = rank // tp_size
tp_rank = rank % server_args.tp_size tp_rank = rank % tp_size
bundle_idx = bundle_indices[rank] bundle_idx = bundle_indices[rank]
actor = _create_scheduler_actor( actor = _create_scheduler_actor(
@@ -455,21 +452,18 @@ class RayEngine(Engine):
RayDataParallelController, 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 # 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: else:
total_gpus = ( total_gpus = parallel.dp_size * parallel.tp_size * parallel.pp_size
get_parallel().config.dp_size gpus_per_node = total_gpus // parallel.nnodes
* server_args.tp_size
* get_parallel().config.pp_size
)
gpus_per_node = total_gpus // server_args.nnodes
logger.info( logger.info(
f"Ray DP cluster: {server_args.nnodes} nodes, " f"Ray DP cluster: {parallel.nnodes} nodes, "
f"{gpus_per_node} GPUs/node, dp_size={get_parallel().config.dp_size}, " f"{gpus_per_node} GPUs/node, dp_size={parallel.dp_size}, "
f"tp_size={server_args.tp_size}, pp_size={get_parallel().config.pp_size}, " f"tp_size={parallel.tp_size}, pp_size={parallel.pp_size}, "
f"enable_dp_attention={get_parallel().config.enable_dp_attention}" f"enable_dp_attention={parallel.enable_dp_attention}"
) )
# Set dist_init_addr on server_args so PortArgs.init_new() can compute # 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 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): def _handle_dwdp(self):
cfg = resolving_view(self) cfg = resolving_view(self)
@@ -9,6 +9,7 @@ import torch
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode 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.server_args import DRAFT_ATTENTION_BACKEND_CHOICES, ServerArgs
from sglang.srt.speculative.dflash_info import DFlashVerifyInput from sglang.srt.speculative.dflash_info import DFlashVerifyInput
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
@@ -28,12 +29,16 @@ class DraftWorkerBundle(msgspec.Struct, frozen=True):
resolved_attention_backend: str resolved_attention_backend: str
def _resolve_draft_attention_backend_fallback( def _resolve_draft_attention_backend_fallback(*, algo_label: str) -> str:
*, server_args: ServerArgs, algo_label: str """The draft's attention backend, from the published leaves.
) -> str:
draft_backend = server_args.speculative_draft_attention_backend `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: if draft_backend is None:
draft_backend, _ = server_args.get_attention_backends() draft_backend, _ = attention_backends()
if draft_backend is None: if draft_backend is None:
return "triton" if torch.version.hip else "flashinfer" return "triton" if torch.version.hip else "flashinfer"
if draft_backend not in DRAFT_ATTENTION_BACKEND_CHOICES: 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 # validated (e.g. a self-drafting architecture); it skips the generic
# supported-backend fallback below. # supported-backend fallback below.
draft_backend = attention_backend_override or ( draft_backend = attention_backend_override or (
_resolve_draft_attention_backend_fallback( _resolve_draft_attention_backend_fallback(algo_label=algo_label)
server_args=server_args, algo_label=algo_label
)
) )
from sglang.srt.layers.moe.utils import draft_model_build_scope 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 return decorator
def reserve_rope_cache_for_long_sequences( def reserve_rope_cache_for_long_sequences(model, model_config, logger=None):
model, server_args, model_config, logger=None """Pre-expand RoPE cache for long sequences and speculative decoding.
):
"""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.environ import envs
from sglang.srt.runtime_context import get_model, get_spec
SAFETY_FACTOR = envs.SGLANG_SPEC_EXPANSION_SAFETY_FACTOR.get() SAFETY_FACTOR = envs.SGLANG_SPEC_EXPANSION_SAFETY_FACTOR.get()
MARGIN = envs.SGLANG_ROPE_CACHE_SAFETY_MARGIN.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 # 1) Estimate base context upper bound
base_ctx = ( base_ctx = (
getattr(server_args, "context_length", None) get_model().context_length
or getattr(model_config, "context_len", None) or getattr(model_config, "context_len", None)
or getattr(model_config, "max_model_len", None) or getattr(model_config, "max_model_len", None)
or getattr(model_config.hf_text_config, "max_position_embeddings", 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 # 2) Speculative decoding expansion
steps = int(getattr(server_args, "speculative_num_steps", 0) or 0) steps = int(get_spec().speculative_num_steps or 0)
draft = int(getattr(server_args, "speculative_num_draft_tokens", 0) or 0) draft = int(get_spec().speculative_num_draft_tokens or 0)
reserve = base_ctx + steps * draft * SAFETY_FACTOR + MARGIN reserve = base_ctx + steps * draft * SAFETY_FACTOR + MARGIN
# 3) Align to reduce reallocation frequency # 3) Align to reduce reallocation frequency
@@ -54,11 +54,16 @@ from .tokenizer import (
_IMAGE_PROCESSOR_BACKENDS = {"auto", "torchvision", "pil"} _IMAGE_PROCESSOR_BACKENDS = {"auto", "torchvision", "pil"}
def resolve_image_processor_backend(server_args) -> str: def resolve_image_processor_backend(mm_config) -> str:
"""Resolve the new backend option while honoring the legacy disable flag.""" """Resolve the new backend option while honoring the legacy disable flag.
if getattr(server_args, "disable_fast_image_processor", False):
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 "pil"
return getattr(server_args, "image_processor_backend", "auto") return getattr(mm_config, "image_processor_backend", "auto")
def _normalize_image_processor_backend( 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
import torch.distributed as dist 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.configs.load_config import LoadConfig
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_parallel, publish from sglang.srt.runtime_context import get_parallel, publish
@@ -148,27 +149,28 @@ class WeightCacheDaemon:
dist_init_method: Optional[str] = None, dist_init_method: Optional[str] = None,
): ):
self.server_args = server_args 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.gpu_id = gpu_id
self.tp_size = server_args.tp_size self.tp_size = cfg.tp_size
self.tp_rank = tp_rank self.tp_rank = tp_rank
self.pp_size = server_args.pp_size self.pp_size = cfg.pp_size
self.pp_rank = pp_rank self.pp_rank = pp_rank
self.dp_size = server_args.dp_size self.dp_size = cfg.dp_size
self.ep_size = server_args.ep_size self.ep_size = cfg.ep_size
self.moe_dp_size = server_args.moe_dp_size self.moe_dp_size = cfg.moe_dp_size
self.enable_dp_attention = server_args.enable_dp_attention self.enable_dp_attention = cfg.enable_dp_attention
self.enable_dp_lm_head = server_args.enable_dp_lm_head self.enable_dp_lm_head = cfg.enable_dp_lm_head
self.attn_cp_size = server_args.attn_cp_size self.attn_cp_size = cfg.attn_cp_size
self.moe_dense_tp_size = server_args.moe_dense_tp_size self.moe_dense_tp_size = cfg.moe_dense_tp_size
self.moe_a2a_backend = server_args.moe_a2a_backend self.moe_a2a_backend = cfg.moe_a2a_backend
self.deepep_mode = server_args.deepep_mode self.deepep_mode = cfg.deepep_mode
self.load_format = server_args.load_format self.load_format = cfg.load_format
self.dtype = server_args.dtype self.dtype = cfg.dtype
self.quantization = server_args.quantization self.quantization = cfg.quantization
self.model_loader_extra_config = server_args.model_loader_extra_config self.model_loader_extra_config = cfg.model_loader_extra_config
self.trust_remote_code = server_args.trust_remote_code self.trust_remote_code = cfg.trust_remote_code
self.revision = server_args.revision self.revision = cfg.revision
self.dist_init_method = dist_init_method self.dist_init_method = dist_init_method
self.socket_path = get_socket_path( self.socket_path = get_socket_path(
@@ -223,7 +225,7 @@ class WeightCacheDaemon:
distributed_init_method=self.dist_init_method, distributed_init_method=self.dist_init_method,
local_rank=self.gpu_id, local_rank=self.gpu_id,
backend=current_platform.get_torch_distributed_backend_str(), 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( initialize_model_parallel(
@@ -281,7 +283,7 @@ class WeightCacheDaemon:
from sglang.srt.layers.moe import initialize_moe_config from sglang.srt.layers.moe import initialize_moe_config
initialize_moe_config(server_args) initialize_moe_config()
# Initialize distributed backend for model loading # Initialize distributed backend for model loading
# (must be done after server_args and model_config are available) # (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 \\ --nnodes 2 --node-rank 1 \\
--dist-init-method tcp://node0-ip:29500 --dist-init-method tcp://node0-ip:29500
""" """
cfg = resolving_view(server_args)
import socket as sock_mod import socket as sock_mod
# Replicate _calculate_rank_ranges logic from engine.py # Replicate _calculate_rank_ranges logic from engine.py
pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1) pp_size_per_node = max(cfg.pp_size // cfg.nnodes, 1)
nnodes_per_pp_rank = max(server_args.nnodes // server_args.pp_size, 1) nnodes_per_pp_rank = max(cfg.nnodes // cfg.pp_size, 1)
pp_rank_range = range( pp_rank_range = range(
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank), pp_size_per_node * (cfg.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 + 1),
) )
nnodes_per_tp_group = nnodes_per_pp_rank 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_rank_range = range(
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group), tp_size_per_node * (cfg.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 + 1),
) )
if server_args.nnodes > 1 and dist_init_method is None: if cfg.nnodes > 1 and dist_init_method is None:
raise ValueError( raise ValueError(
"dist_init_method is required for multi-node weight cache daemons. " "dist_init_method is required for multi-node weight cache daemons. "
"Use --dist-init-method tcp://<node0-ip>:<port> to specify the " "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. # Validate and clean up stale .ready/.sock files from prior runs.
for pp_rank in pp_rank_range: for pp_rank in pp_rank_range:
for tp_rank in tp_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) cleanup_stale_daemon_files(global_rank, force=force)
procs = [] procs = []
@@ -716,8 +719,8 @@ def launch_weight_cache_daemons(
tp_rank, tp_rank,
pp_size_per_node, pp_size_per_node,
tp_size_per_node, tp_size_per_node,
base_gpu_id=server_args.base_gpu_id, base_gpu_id=cfg.base_gpu_id,
gpu_id_step=server_args.gpu_id_step, gpu_id_step=cfg.gpu_id_step,
) )
proc = spawn_weight_cache_daemon( proc = spawn_weight_cache_daemon(
server_args, server_args,
@@ -738,7 +741,7 @@ def launch_weight_cache_daemons(
start_time = time.time() start_time = time.time()
for pp_rank in pp_rank_range: for pp_rank in pp_rank_range:
for tp_rank in tp_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) ready_path = get_ready_path(global_rank)
while not os.path.exists(ready_path): while not os.path.exists(ready_path):
time.sleep(check_interval) time.sleep(check_interval)
@@ -772,7 +775,7 @@ def launch_weight_cache_daemons(
) )
logger.info( 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"(pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, "
f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, " f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, "
f"dist_init_method={dist_init_method})" f"dist_init_method={dist_init_method})"
@@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Generator, List, Optional, Tuple
import zmq import zmq
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
from sglang.srt.utils.network import get_zmq_socket 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 """Run the engine until the server warmup request has been received and
fully processed, so scripts never observe foreign warmup traffic.""" fully processed, so scripts never observe foreign warmup traffic."""
scheduler = ctx.scheduler scheduler = ctx.scheduler
server_args = scheduler.server_args cfg = resolving_view(scheduler.server_args)
if server_args.skip_server_warmup: if cfg.skip_server_warmup:
logger.info("scripted_runtime: skip_server_warmup set, not driving warmup") logger.info("scripted_runtime: skip_server_warmup set, not driving warmup")
return 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_fully_idle() can transiently report idle while a PP microbatch result
# is still in flight, so require it to hold for two full microbatch # is still in flight, so require it to hold for two full microbatch
# rotations after the warmup request was observed on the recv socket. # 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 proxy = ctx._tokenizer_recv_proxy
deadline = start_time + WARMUP_DRIVE_TIMEOUT_S 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.dp_attention import set_dp_buffer_len
from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatcher from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatcher
from sglang.srt.layers.moe.utils import initialize_moe_config 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.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -22,7 +23,8 @@ class TestFlashinferDispatcher(CustomTestCase):
server_args.moe_runner_backend = "flashinfer_cutlass" server_args.moe_runner_backend = "flashinfer_cutlass"
server_args.moe_a2a_backend = "flashinfer" server_args.moe_a2a_backend = "flashinfer"
set_global_server_args_for_scheduler(server_args) set_global_server_args_for_scheduler(server_args)
initialize_moe_config(server_args) publish(server_args, role="scheduler")
initialize_moe_config()
init_distributed_environment( init_distributed_environment(
world_size=-1, # Auto-detect from environment world_size=-1, # Auto-detect from environment
+9 -21
View File
@@ -60,7 +60,7 @@ class _FakeCPGroup:
class TestCPStrategyUnit(CustomTestCase): class TestCPStrategyUnit(CustomTestCase):
def tearDown(self): 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): def test_strategy_kind_maps_cli_values(self):
self.assertEqual(ContextParallelStrategyKind.NONE.value, 0) self.assertEqual(ContextParallelStrategyKind.NONE.value, 0)
@@ -77,11 +77,9 @@ class TestCPStrategyUnit(CustomTestCase):
def test_init_cp_strategy_binds_zigzag_strategy(self): def test_init_cp_strategy_binds_zigzag_strategy(self):
init_cp_strategy( init_cp_strategy(
SimpleNamespace(
enable_prefill_cp=True, enable_prefill_cp=True,
cp_size=4,
cp_strategy="zigzag", cp_strategy="zigzag",
attn_cp_size=4,
)
) )
self.assertTrue(is_cp_enabled()) self.assertTrue(is_cp_enabled())
@@ -91,11 +89,9 @@ class TestCPStrategyUnit(CustomTestCase):
def test_get_cp_strategy_is_initialized_under_cp_v2(self): def test_get_cp_strategy_is_initialized_under_cp_v2(self):
init_cp_strategy( init_cp_strategy(
SimpleNamespace(
enable_prefill_cp=True, enable_prefill_cp=True,
cp_size=4,
cp_strategy="interleave", cp_strategy="interleave",
attn_cp_size=4,
)
) )
with patch( with patch(
@@ -108,7 +104,7 @@ class TestCPStrategyUnit(CustomTestCase):
class TestPrefillCPBCGReplay(CustomTestCase): class TestPrefillCPBCGReplay(CustomTestCase):
def tearDown(self): 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): def _make_runner(self):
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner) runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
@@ -141,11 +137,9 @@ class TestPrefillCPBCGReplay(CustomTestCase):
def _enable_zigzag(self): def _enable_zigzag(self):
init_cp_strategy( init_cp_strategy(
SimpleNamespace(
enable_prefill_cp=True, enable_prefill_cp=True,
cp_size=4,
cp_strategy="zigzag", cp_strategy="zigzag",
attn_cp_size=4,
)
) )
def test_local_capacity_overflow_uses_next_capture_bucket(self): def test_local_capacity_overflow_uses_next_capture_bucket(self):
@@ -267,16 +261,13 @@ class TestPrefillCPBCGReplay(CustomTestCase):
class TestCPZigzagStrategy(CustomTestCase): class TestCPZigzagStrategy(CustomTestCase):
def setUp(self): def setUp(self):
init_cp_strategy( init_cp_strategy(
SimpleNamespace(
enable_prefill_cp=True, enable_prefill_cp=True,
cp_size=4,
cp_strategy="zigzag", cp_strategy="zigzag",
attn_cp_size=4,
attention_backend="fa3",
)
) )
def tearDown(self): 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): def _metadata_for_rank(self, rank, *, cp_size, seq_lens, extend_seq_lens):
strategy = ZigzagCPStrategy(cp_size=cp_size) strategy = ZigzagCPStrategy(cp_size=cp_size)
@@ -809,16 +800,13 @@ class TestCPZigzagStrategy(CustomTestCase):
class TestCPInterleaveStrategy(CustomTestCase): class TestCPInterleaveStrategy(CustomTestCase):
def setUp(self): def setUp(self):
init_cp_strategy( init_cp_strategy(
SimpleNamespace(
enable_prefill_cp=True, enable_prefill_cp=True,
cp_size=4,
cp_strategy="interleave", cp_strategy="interleave",
attn_cp_size=4,
attention_backend="fa3",
)
) )
def tearDown(self): 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): def _metadata_for_rank(self, rank, *, cp_size, seq_lens, extend_seq_lens):
strategy = InterleaveCPStrategy(cp_size=cp_size) strategy = InterleaveCPStrategy(cp_size=cp_size)
@@ -71,7 +71,7 @@ class TestEmbeddingModelSpec(unittest.TestCase):
) )
plan = resolved_embedding_plan( plan = resolved_embedding_plan(
spec, spec,
server_args=SimpleNamespace( config=SimpleNamespace(
is_embedding=True, is_embedding=True,
cuda_graph_config=SimpleNamespace( cuda_graph_config=SimpleNamespace(
prefill=SimpleNamespace( prefill=SimpleNamespace(
@@ -1,5 +1,4 @@
import contextlib import contextlib
import types
import unittest import unittest
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
@@ -90,23 +89,24 @@ def _torch_allreduce_residual_rmsnorm_baseline(
class TestFlashInferCommFusion(CustomTestCase): 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): def test_auto_backend_resolves_by_arch(self):
single_node = types.SimpleNamespace( single_node = ("auto", False)
flashinfer_allreduce_fusion_backend="auto", nnodes=1 multi_node = ("auto", True)
)
multi_node = types.SimpleNamespace(
flashinfer_allreduce_fusion_backend="auto", nnodes=2
)
# Blackwell: mnnvl on both single-node and multi-node. # Blackwell: mnnvl on both single-node and multi-node.
with patch.object(fusion, "is_sm100_supported", return_value=True): with patch.object(fusion, "is_sm100_supported", return_value=True):
self.assertEqual( self.assertEqual(
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node), fusion._resolve_backend(*single_node),
"mnnvl", "mnnvl",
) )
self.assertEqual( self.assertEqual(fusion._resolve_backend(*multi_node), "mnnvl")
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node), "mnnvl"
)
# SM90: auto uses trtllm on single-node, multi-node is unsupported. # SM90: auto uses trtllm on single-node, multi-node is unsupported.
with ( with (
@@ -114,11 +114,11 @@ class TestFlashInferCommFusion(CustomTestCase):
patch.object(fusion, "is_sm90_supported", return_value=True), patch.object(fusion, "is_sm90_supported", return_value=True),
): ):
self.assertEqual( self.assertEqual(
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node), fusion._resolve_backend(*single_node),
"trtllm", "trtllm",
) )
with self.assertRaises(ValueError): 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 # Architectures outside SM90/SM10X are unsupported. Both pre-SM90
# and post-SM10X devices (e.g. SM120) must fail closed. # 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), patch.object(fusion, "is_sm90_supported", return_value=False),
): ):
with self.assertRaises(ValueError): with self.assertRaises(ValueError):
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node) fusion._resolve_backend(*single_node)
with self.assertRaises(ValueError): with self.assertRaises(ValueError):
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node) fusion._resolve_backend(*multi_node)
def test_explicit_backend_validation(self): def test_explicit_backend_validation(self):
single_node_mnnvl = types.SimpleNamespace( single_node_mnnvl = ("mnnvl", False)
flashinfer_allreduce_fusion_backend="mnnvl", nnodes=1 multi_node_mnnvl = ("mnnvl", True)
) single_node_trtllm = ("trtllm", False)
multi_node_mnnvl = types.SimpleNamespace( multi_node_trtllm = ("trtllm", True)
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
)
with ( with (
patch.object(fusion, "is_sm100_supported", return_value=False), patch.object(fusion, "is_sm100_supported", return_value=False),
patch.object(fusion, "is_sm90_supported", return_value=True), patch.object(fusion, "is_sm90_supported", return_value=True),
): ):
self.assertEqual( self.assertEqual(
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node_mnnvl), fusion._resolve_backend(*single_node_mnnvl),
"mnnvl", "mnnvl",
) )
self.assertEqual( self.assertEqual(
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node_trtllm), fusion._resolve_backend(*single_node_trtllm),
"trtllm", "trtllm",
) )
with self.assertRaises(ValueError): with self.assertRaises(ValueError):
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_mnnvl) fusion._resolve_backend(*multi_node_mnnvl)
with self.assertRaises(ValueError): 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): with patch.object(fusion, "is_sm100_supported", return_value=True):
self.assertEqual( self.assertEqual(
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_mnnvl), fusion._resolve_backend(*multi_node_mnnvl),
"mnnvl", "mnnvl",
) )
with self.assertRaises(ValueError): 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"): for arch in ("pre_sm90", "post_sm10x"):
with ( with (
@@ -184,9 +176,9 @@ class TestFlashInferCommFusion(CustomTestCase):
single_node_trtllm, single_node_trtllm,
multi_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): with self.assertRaises(ValueError):
fusion.resolve_flashinfer_allreduce_fusion_backend(args) fusion._resolve_backend(*args)
def test_allreduce_fusion_backends_match_torch_baseline(self): def test_allreduce_fusion_backends_match_torch_baseline(self):
fake_comm = _FakeFlashInferComm() fake_comm = _FakeFlashInferComm()
@@ -27,6 +27,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
Phase, Phase,
PhaseConfig, 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 import PortArgs, ServerArgs, prepare_server_args
from sglang.srt.server_args_config_parser import ConfigArgumentMerger from sglang.srt.server_args_config_parser import ConfigArgumentMerger
from sglang.test.ci.ci_register import register_cpu_ci 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_core = SimpleNamespace(start_server=MagicMock(return_value="handle"))
fake_bridge = SimpleNamespace(RuntimeHandle=MagicMock(return_value="rt")) fake_bridge = SimpleNamespace(RuntimeHandle=MagicMock(return_value="rt"))
server_args = SimpleNamespace( # The host comes from the `serving` bag; `grpc_worker_threads` is not a
host="127.0.0.1", grpc_port=50051, grpc_worker_threads=4 # 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 ( with (
patch( patch(
"sglang.srt.rust_extensions.load_rust_extension", "sglang.srt.rust_extensions.load_rust_extension",
@@ -2373,7 +2378,7 @@ class TestGrpcServerArgs(CustomTestCase):
tokenizer_manager=MagicMock(), tokenizer_manager=MagicMock(),
template_manager=MagicMock(), template_manager=MagicMock(),
scheduler_info={}, scheduler_info={},
grpc_port=resolution_result(server_args, "grpc_port"), grpc_port=get_serving().grpc_port,
) )
self.assertEqual(handle, "handle") self.assertEqual(handle, "handle")
@@ -19,7 +19,7 @@ from sglang.srt.layers.moe.utils import (
speculative_moe_backend_context, speculative_moe_backend_context,
) )
from sglang.srt.model_executor.model_runner import ModelRunner 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.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -145,9 +145,11 @@ class TestFusionDecisionFlag(CustomTestCase):
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
self._seed() self._seed()
initialize_moe_config( publish(
ServerArgs(model_path="dummy", disable_shared_experts_fusion=True) ServerArgs(model_path="dummy", disable_shared_experts_fusion=True),
role="scheduler",
) )
initialize_moe_config()
moe = get_flags().moe moe = get_flags().moe
self.assertTrue(moe.disable_shared_experts_fusion) self.assertTrue(moe.disable_shared_experts_fusion)
self.assertTrue(moe.speculative_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 # The test below asserts this map is exactly the set of such reads, so the
# reasons cannot drift away from the code. # reasons cannot drift away from the code.
_CONFIGURED_SIZE_CALL_SITES = { _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"): ( ("srt/entrypoints/engine.py", "pp_size"): (
"the launch path decides how many scheduler processes to spawn; it runs " "the launch path decides how many scheduler processes to spawn; it runs "
"before any of them exists, so there is no group to ask" "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 " "the Ray driver sizes the actor placement group; the actors it is about "
"to create are the ones that will hold the process groups" "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"): ( ("srt/ray/data_parallel_controller.py", "pp_size"): (
"same placement arithmetic on the DP path -- ranks per TP group, " "same placement arithmetic on the DP path -- ranks per TP group, "
"computed in the driver before the actors start" "computed in the driver before the actors start"
@@ -115,6 +134,26 @@ _CONFIGURED_SIZE_CALL_SITES = {
("srt/managers/scheduler.py", "dcp_size"): ( ("srt/managers/scheduler.py", "dcp_size"): (
"same pre-distributed-init arithmetic in configure_scheduler_process" "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"): ( ("srt/model_executor/runner/base_runner.py", "pp_size"): (
"the runner's layer window is arithmetic over the configured stage " "the runner's layer window is arithmetic over the configured stage "
"count; a draft runner shares the target's groups, so the live " "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 " "the encode server's launch entry sizes its workers before it has "
"spawned any of them" "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"): ( ("srt/utils/common.py", "tp_size"): (
"the require_*_tp_gather predicates compared the configured tp_size " "the require_*_tp_gather predicates compared the configured tp_size "
"when they read the record; the live property answers a different " "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.""" swap under the speculative contexts and restore on exit."""
def _init(self, **kw): def _init(self, **kw):
from types import SimpleNamespace
from sglang.srt.layers.moe.utils import initialize_moe_config from sglang.srt.layers.moe.utils import initialize_moe_config
defaults = dict( defaults = dict(
@@ -538,7 +536,12 @@ class TestMoeFlagsGroup(_IsolatedServerArgs):
disable_shared_experts_fusion=False, disable_shared_experts_fusion=False,
) )
defaults.update(kw) 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): def test_lazy_defaults_before_initialize(self):
from sglang.srt.layers.moe.utils import ( 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. # are step-12 exposure like any other pair.
_PASSED = frozenset({"model_path", "device", "random_seed"}) _PASSED = frozenset({"model_path", "device", "random_seed"})
# The reads that still take a value off the supplied instance. `initialize_moe_config` # Empty. A pair belongs here when a reader has no bag to read -- it runs before
# is handed the record until the replay goes away; the rest are pre-publish launcher # its process publishes -- and cannot use `resolving_view` either. The launcher's
# reads. # pre-publish reads (`_set_envs_and_config`, the auto-parser gate) and the
_EXPOSED = { # late-resolution detection it calls all read the declarations now, so nothing
("disaggregation/encoder/server.py", "model_loader_extra_config"), # qualifies. A new entry needs that kind of reason next to it.
("layers/moe/utils.py", "deepep_mode"), _EXPOSED: frozenset = frozenset()
("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"),
}
# Pairs whose resolution write only happens on a CUDA host (capability or # Pairs whose resolution write only happens on a CUDA host (capability or
# `is_cuda()` gated): asserted on the CUDA registration, invisible to the CPU # `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 # 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 # 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. # conversion; the list exists so a new one is a decision made when it is written.
_OVERRIDDEN_AND_READ = { _OVERRIDDEN_AND_READ: frozenset = frozenset()
("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"),
}
def _expanded_override_keys(rel, tree, call, kw) -> set: def _expanded_override_keys(rel, tree, call, kw) -> set: