config: the runtime readers take the published bags (#36254)
This commit is contained in:
@@ -139,9 +139,9 @@ bag to override at all.
|
||||
`Engine`s can share one process, bags are last-publish-wins across them") is
|
||||
**retracted** — owner ruling (2026-08-15): a process holds at most one live
|
||||
config at a time (concurrent multi-Engine is unsupported; sequential rebuild
|
||||
stays legal, unit tests rely on it). What still reads the instance in those
|
||||
files is pinned pair by pair in the exposure ratchet, each with its own
|
||||
disposition; none of it is a boundary to imitate. What
|
||||
stays legal, unit tests rely on it). Nothing in those files reads the instance
|
||||
any more -- the exposure ratchet's pin set is empty, so the next such read is a
|
||||
new entry that has to argue for itself. What
|
||||
genuinely stays per-instance is what differs per *worker* within one engine:
|
||||
`base_gpu_id` travels as a constructor argument (`MMEncoder(gpu_id=...)`;
|
||||
`BaseMultimodalProcessor._fast_image_processor_device` is the shape to copy).
|
||||
@@ -149,18 +149,17 @@ bag to override at all.
|
||||
supplied-instance contract; don't rewrite the parameter reads unless the
|
||||
field is runtime-mutated (see the elastic-EP `ep_size` case in
|
||||
`eplb/expert_location.py`) — **or the field is one that resolution fills in
|
||||
and the callee runs in a process that has published.** That second case is
|
||||
pinned debt, not a style question: the record is destined to carry the
|
||||
user's raw input, so `server_args.page_size` inside a runner-owned
|
||||
constructor will read the raw pre-resolution value instead of the effective
|
||||
one. Debt means a decision, not automatically a bag read: pick where the
|
||||
and the callee runs in a process that has published.** That second case is a
|
||||
decision, not a style question: the record carries the user's raw input, so a
|
||||
resolution-filled field read off it inside a runner-owned constructor answers
|
||||
with the pre-resolution value instead of the effective one. Debt means a decision, not automatically a bag read: pick where the
|
||||
value should come from — usually the `get_*()` bag, sometimes a runner stamp
|
||||
or a constructor argument (the per-mode attention pair and the encode-server
|
||||
`gpu_id` above are dispositions of exactly this debt). The per-instance
|
||||
boundaries above are **not** exempt from this unless-clause (the multi-Engine
|
||||
exemption is retracted); each one gets its own disposition.
|
||||
`test_supplied_instance_exposure_ratchet.py`
|
||||
pins the remaining set — three spellings of the read: `server_args.field`,
|
||||
pins that set (empty today) — three spellings of the read: `server_args.field`,
|
||||
literal-name `getattr(server_args, "field", default)`, and the parked form
|
||||
(`self.x = server_args` in a method that takes the parameter, read as
|
||||
`self.x.field` anywhere in the class) — and fails on a new one, so the
|
||||
@@ -411,7 +410,8 @@ probes, swappable ACTIVE values. Not for config mirrors (read the bag leaf inste
|
||||
|
||||
- Groups are typed dataclasses on `Flags` (`capture` / `moe` / `dp`): typo-safe writes,
|
||||
transactional test-only `override(**kw)` context manager.
|
||||
- `flags.moe` is materialized by `initialize_moe_config(server_args)` at scheduler init;
|
||||
- `flags.moe` is materialized by `initialize_moe_config()` at scheduler init (it
|
||||
reads `exec.moe` / `spec` / `model`, and takes no record);
|
||||
accessors (`get_moe_a2a_backend` etc.) are thin shims with lazy defaults. The speculative
|
||||
contexts (`speculative_moe_backend_context`) swap the ACTIVE leaves around draft forwards.
|
||||
- `flags.dp` is materialized by `initialize_dp_attention`; `is_dp_attention_enabled()` is a
|
||||
|
||||
@@ -30,6 +30,7 @@ from sglang.benchmark.datasets import DatasetRow, get_dataset
|
||||
from sglang.benchmark.datasets.random import sample_random_requests
|
||||
from sglang.benchmark.utils import get_tokenizer, set_ulimit
|
||||
from sglang.lang.backend.runtime_endpoint import Runtime
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.entrypoints.engine import Engine
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
@@ -366,6 +367,7 @@ def _create_ray_engine_backend(server_args: ServerArgs):
|
||||
RayEngine requires a placement group, so we launch it inside a Ray actor
|
||||
and return a lightweight proxy that forwards calls via ray.get().
|
||||
"""
|
||||
cfg = resolving_view(server_args)
|
||||
import ray
|
||||
from ray.runtime_env import RuntimeEnv
|
||||
from ray.util.placement_group import placement_group
|
||||
@@ -377,7 +379,7 @@ def _create_ray_engine_backend(server_args: ServerArgs):
|
||||
if not ray.is_initialized():
|
||||
ray.init(runtime_env=RuntimeEnv(env_vars=env_vars))
|
||||
|
||||
total_gpus = server_args.tp_size * server_args.pp_size
|
||||
total_gpus = cfg.tp_size * cfg.pp_size
|
||||
pg = placement_group([{"CPU": 1, "GPU": total_gpus}], strategy="STRICT_PACK")
|
||||
ray.get(pg.ready())
|
||||
|
||||
@@ -398,7 +400,7 @@ def _create_ray_engine_backend(server_args: ServerArgs):
|
||||
placement_group=pg,
|
||||
placement_group_bundle_index=0,
|
||||
),
|
||||
).remote(**dict(server_args._raw_input))
|
||||
).remote(**dict(cfg._raw_input))
|
||||
|
||||
class _Proxy:
|
||||
"""Forwards method calls to the remote RayEngine actor."""
|
||||
@@ -434,20 +436,21 @@ def throughput_test(
|
||||
):
|
||||
# A programmatic caller may hand over a freshly constructed record, and
|
||||
# the backends below read the resolved paths and the raw snapshot.
|
||||
server_args.resolve_once()
|
||||
cfg = resolving_view(server_args)
|
||||
cfg.resolve_once()
|
||||
if bench_args.backend == "engine":
|
||||
if server_args.use_ray:
|
||||
if cfg.use_ray:
|
||||
backend = _create_ray_engine_backend(server_args)
|
||||
else:
|
||||
backend = Engine(server_args=server_args)
|
||||
if not backend:
|
||||
raise ValueError("Please provide valid engine arguments")
|
||||
elif bench_args.backend == "runtime":
|
||||
backend = Runtime(**dict(server_args._raw_input))
|
||||
backend = Runtime(**dict(cfg._raw_input))
|
||||
else:
|
||||
raise ValueError('Please set backend to either "engine" or "runtime"')
|
||||
|
||||
tokenizer_id = server_args.tokenizer_path or server_args.model_path
|
||||
tokenizer_id = cfg.tokenizer_path or cfg.model_path
|
||||
tokenizer = get_tokenizer(tokenizer_id)
|
||||
|
||||
# Set global environments
|
||||
|
||||
@@ -64,6 +64,7 @@ import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolution_result, resolving_view
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
destroy_distributed_environment,
|
||||
@@ -79,7 +80,11 @@ from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config
|
||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||
from sglang.srt.managers.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw
|
||||
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||
from sglang.srt.model_executor.cuda_graph_config import Phase, cuda_graph_fully_disabled
|
||||
from sglang.srt.model_executor.cuda_graph_config import (
|
||||
CudaGraphConfig,
|
||||
Phase,
|
||||
cuda_graph_fully_disabled,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.runtime_context import get_parallel, get_schedule, publish
|
||||
@@ -297,44 +302,45 @@ class BenchArgs:
|
||||
|
||||
|
||||
def load_model(server_args, port_args, gpu_id, tp_rank):
|
||||
cfg = resolving_view(server_args)
|
||||
suppress_other_loggers()
|
||||
rank_print = print if tp_rank == 0 else lambda *args, **kwargs: None
|
||||
moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size)
|
||||
moe_ep_rank = tp_rank // (cfg.tp_size // cfg.ep_size)
|
||||
|
||||
model_config = ModelConfig.from_server_args(server_args)
|
||||
attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = (
|
||||
compute_dp_attention_world_info(
|
||||
server_args.enable_dp_attention,
|
||||
cfg.enable_dp_attention,
|
||||
tp_rank,
|
||||
server_args.tp_size,
|
||||
server_args.dp_size,
|
||||
server_args.attn_cp_size,
|
||||
cfg.tp_size,
|
||||
cfg.dp_size,
|
||||
cfg.attn_cp_size,
|
||||
)
|
||||
)
|
||||
ps = ParallelState(
|
||||
tp_rank=tp_rank,
|
||||
tp_size=server_args.tp_size,
|
||||
tp_size=cfg.tp_size,
|
||||
pp_rank=0,
|
||||
pp_size=1,
|
||||
dp_rank=None,
|
||||
dp_size=server_args.dp_size,
|
||||
dp_size=cfg.dp_size,
|
||||
attn_tp_rank=attn_tp_rank,
|
||||
attn_tp_size=attn_tp_size,
|
||||
attn_cp_rank=0,
|
||||
attn_cp_size=server_args.attn_cp_size,
|
||||
attn_dcp_rank=tp_rank % server_args.dcp_size,
|
||||
attn_dcp_size=server_args.dcp_size,
|
||||
attn_cp_size=cfg.attn_cp_size,
|
||||
attn_dcp_rank=tp_rank % cfg.dcp_size,
|
||||
attn_dcp_size=cfg.dcp_size,
|
||||
attn_dp_rank=attn_dp_rank,
|
||||
attn_dp_size=attn_dp_size,
|
||||
moe_ep_rank=moe_ep_rank,
|
||||
moe_ep_size=server_args.ep_size,
|
||||
moe_ep_size=cfg.ep_size,
|
||||
moe_dp_rank=None,
|
||||
moe_dp_size=server_args.moe_dp_size,
|
||||
moe_dp_size=cfg.moe_dp_size,
|
||||
gpu_id=gpu_id,
|
||||
)
|
||||
runner_kwargs = dict(
|
||||
model_config=model_config,
|
||||
mem_fraction_static=server_args.mem_fraction_static,
|
||||
mem_fraction_static=cfg.mem_fraction_static,
|
||||
gpu_id=gpu_id,
|
||||
ps=ps,
|
||||
nccl_port=port_args.nccl_port,
|
||||
@@ -350,20 +356,20 @@ def load_model(server_args, port_args, gpu_id, tp_rank):
|
||||
model_runner = MlxModelRunnerStub(**runner_kwargs)
|
||||
else:
|
||||
model_runner = ModelRunner(**runner_kwargs)
|
||||
if server_args.is_startup_weight_load_overlap:
|
||||
if cfg.is_startup_weight_load_overlap:
|
||||
model_runner.start_startup_weight_load()
|
||||
model_runner.alloc_memory_pool()
|
||||
model_runner.init_attention_backends()
|
||||
model_runner.init_cuda_graphs()
|
||||
if server_args.is_startup_weight_load_overlap:
|
||||
if cfg.is_startup_weight_load_overlap:
|
||||
model_runner.finalize_startup_weight_load()
|
||||
rank_print(f"max_total_num_tokens={model_runner.max_total_num_tokens}")
|
||||
tokenizer = get_tokenizer(
|
||||
server_args.tokenizer_path,
|
||||
tokenizer_mode=server_args.tokenizer_mode,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
cfg.tokenizer_path,
|
||||
tokenizer_mode=cfg.tokenizer_mode,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
)
|
||||
if server_args.tp_size > 1:
|
||||
if cfg.tp_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
if _use_mlx:
|
||||
@@ -584,19 +590,20 @@ class _MlxBenchRunner:
|
||||
"""Wraps MlxModelRunner for the MLX benchmark path."""
|
||||
|
||||
def __init__(self, model_runner, server_args):
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner
|
||||
|
||||
# Radix cache requires the scheduler's allocator/trie; disable in
|
||||
# standalone bench mode where no scheduler is present.
|
||||
init_kwargs = dict(
|
||||
model_path=server_args.model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
model_path=cfg.model_path,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
disable_radix_cache=True,
|
||||
mem_fraction_static=server_args.mem_fraction_static,
|
||||
quantization=server_args.quantization,
|
||||
mem_fraction_static=cfg.mem_fraction_static,
|
||||
quantization=cfg.quantization,
|
||||
)
|
||||
if server_args.max_total_tokens is not None:
|
||||
init_kwargs["pool_size"] = server_args.max_total_tokens
|
||||
if cfg.max_total_tokens is not None:
|
||||
init_kwargs["pool_size"] = cfg.max_total_tokens
|
||||
self.mlx_runner = MlxModelRunner(**init_kwargs)
|
||||
self.mlx_runner.init_cache_pools(req_to_token_pool=None)
|
||||
self.fake_torch_runner = model_runner
|
||||
@@ -883,17 +890,18 @@ def latency_test(
|
||||
gpu_id,
|
||||
tp_rank,
|
||||
):
|
||||
cfg = resolving_view(server_args)
|
||||
# `main` runs this inline for tp_size == 1 and spawns it per rank otherwise;
|
||||
# a spawned child arrives with nothing published.
|
||||
publish(server_args, role="scheduler")
|
||||
initialize_moe_config(server_args)
|
||||
initialize_fp8_gemm_config(server_args)
|
||||
initialize_fp4_gemm_config(server_args)
|
||||
initialize_moe_config()
|
||||
initialize_fp8_gemm_config()
|
||||
initialize_fp4_gemm_config()
|
||||
|
||||
# Set CPU affinity
|
||||
if get_bool_env_var("SGLANG_SET_CPU_AFFINITY"):
|
||||
parallel = get_parallel().config
|
||||
set_gpu_proc_affinity(
|
||||
server_args.pp_size, server_args.tp_size, server_args.nnodes, tp_rank
|
||||
parallel.pp_size, parallel.tp_size, parallel.nnodes, tp_rank
|
||||
)
|
||||
|
||||
# Configure the logger
|
||||
@@ -988,22 +996,45 @@ def latency_test(
|
||||
for result in result_list:
|
||||
fout.write(json.dumps(result) + "\n")
|
||||
|
||||
if server_args.tp_size > 1:
|
||||
if cfg.tp_size > 1:
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
|
||||
|
||||
def main(server_args, bench_args):
|
||||
# The decode phase has to capture the batch sizes this run benchmarks, and
|
||||
# the per-phase convenience knob loses to an explicit --cuda-graph-config
|
||||
# JSON (resolution applies that last), so the size is merged into that JSON.
|
||||
if getattr(server_args, "_declarations_materialized", False):
|
||||
# A record the caller already resolved: nothing will parse a raw dict
|
||||
# again, so the declaration has to be the finished typed config.
|
||||
merged = resolution_result(server_args, "cuda_graph_config")
|
||||
merged = (
|
||||
merged.to_dict()
|
||||
if isinstance(merged, CudaGraphConfig)
|
||||
else dict(merged or {})
|
||||
)
|
||||
decode = dict(merged.get(Phase.DECODE) or {})
|
||||
decode["max_bs"] = max(bench_args.batch_size)
|
||||
merged[Phase.DECODE] = decode
|
||||
graph_config = CudaGraphConfig.from_dict(merged)
|
||||
else:
|
||||
explicit = server_args.cuda_graph_config
|
||||
if isinstance(explicit, CudaGraphConfig):
|
||||
explicit = explicit.to_dict()
|
||||
graph_config = dict(explicit or {})
|
||||
decode = dict(graph_config.get(Phase.DECODE) or {})
|
||||
decode["max_bs"] = max(bench_args.batch_size)
|
||||
graph_config[Phase.DECODE] = decode
|
||||
server_args = server_args.replace_resolved(
|
||||
"benchmark.one_batch", cuda_graph_config=graph_config
|
||||
)
|
||||
server_args.resolve_once()
|
||||
|
||||
# The legacy cuda_graph_max_bs_decode field does not propagate; set the
|
||||
# decode phase.
|
||||
if server_args.cuda_graph_config is not None:
|
||||
server_args.cuda_graph_config[Phase.DECODE].max_bs = max(bench_args.batch_size)
|
||||
cfg = resolving_view(server_args)
|
||||
|
||||
_set_envs_and_config(server_args)
|
||||
|
||||
if server_args.model_path:
|
||||
if cfg.model_path:
|
||||
if bench_args.correctness_test:
|
||||
work_func = correctness_test
|
||||
else:
|
||||
|
||||
@@ -33,6 +33,7 @@ from sglang.benchmark.datasets import get_dataset
|
||||
from sglang.benchmark.endpoint import acquire_endpoint
|
||||
from sglang.benchmark.utils import get_processor, get_tokenizer
|
||||
from sglang.profiler import run_profile
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST
|
||||
from sglang.srt.entrypoints.http_server import launch_server
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -1234,6 +1235,7 @@ def run_benchmark_internal(
|
||||
|
||||
|
||||
def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
|
||||
cfg = resolving_view(server_args)
|
||||
results, server_info = run_benchmark_internal(server_args, bench_args)
|
||||
|
||||
# Save results as pydantic models in the JSON format
|
||||
@@ -1241,7 +1243,7 @@ def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
|
||||
save_results_as_pydantic_models(
|
||||
results,
|
||||
pydantic_result_filename=bench_args.pydantic_result_filename,
|
||||
model_path=server_args.model_path,
|
||||
model_path=cfg.model_path,
|
||||
server_args=bench_args.server_args_for_metrics,
|
||||
)
|
||||
|
||||
|
||||
@@ -17,13 +17,19 @@ import time
|
||||
|
||||
import requests
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST
|
||||
from sglang.srt.entrypoints.http_server import launch_server
|
||||
from sglang.srt.entrypoints.warmup import warmup
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase
|
||||
from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Backend,
|
||||
CudaGraphConfig,
|
||||
Phase,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
|
||||
@@ -59,8 +65,7 @@ async def warm_up_compile(
|
||||
disaggregation_mode: str, tokenizer_manager: TokenizerManager
|
||||
):
|
||||
print("\nGenerate warm up request for compiling DeepGEMM...\n")
|
||||
server_args = tokenizer_manager.server_args
|
||||
dp_size = server_args.dp_size
|
||||
dp_size = get_parallel().config.dp_size
|
||||
base_ids = [0, 1, 2, 3]
|
||||
sampling_params = {
|
||||
"temperature": 0.0,
|
||||
@@ -76,7 +81,8 @@ async def warm_up_compile(
|
||||
)
|
||||
generate_req_input.bootstrap_host = [FAKE_BOOTSTRAP_HOST] * dp_size
|
||||
generate_req_input.bootstrap_room = [
|
||||
i * (2**63 // dp_size) + (i % server_args.tp_size) for i in range(dp_size)
|
||||
i * (2**63 // dp_size) + (i % get_parallel().config.tp_size)
|
||||
for i in range(dp_size)
|
||||
]
|
||||
else:
|
||||
input_ids = (
|
||||
@@ -105,6 +111,7 @@ def launch_server_process_and_send_one_request(
|
||||
# Keeps the device probe out of the fork below, for a caller that reaches
|
||||
# this without resolving first.
|
||||
server_args.resolve_once()
|
||||
cfg = resolving_view(server_args)
|
||||
|
||||
proc = multiprocessing.Process(target=launch_server_internal, args=(server_args,))
|
||||
proc.start()
|
||||
@@ -125,7 +132,7 @@ def launch_server_process_and_send_one_request(
|
||||
if response.status_code == 200:
|
||||
# Rank-0 node send a request to sync with other node and then return.
|
||||
if server_args.node_rank == 0:
|
||||
dp_size = server_args.dp_size
|
||||
dp_size = cfg.dp_size
|
||||
base_ids = [0, 1, 2, 3]
|
||||
payload = {
|
||||
"sampling_params": {
|
||||
@@ -133,11 +140,11 @@ def launch_server_process_and_send_one_request(
|
||||
"temperature": 0,
|
||||
},
|
||||
}
|
||||
if server_args.disaggregation_mode != "null":
|
||||
if cfg.disaggregation_mode != "null":
|
||||
payload["input_ids"] = [list(base_ids) for _ in range(dp_size)]
|
||||
payload["bootstrap_host"] = [FAKE_BOOTSTRAP_HOST] * dp_size
|
||||
payload["bootstrap_room"] = [
|
||||
i * (2**63 // dp_size) + (i % server_args.tp_size)
|
||||
i * (2**63 // dp_size) + (i % cfg.tp_size)
|
||||
for i in range(dp_size)
|
||||
]
|
||||
else:
|
||||
@@ -177,16 +184,24 @@ def compile_server_args(args, compile_args: CompileArgs) -> ServerArgs:
|
||||
"""The config this script serves with: no cuda graph, no torch compile, and a
|
||||
watchdog that outlives the compilation."""
|
||||
args.enable_torch_compile = False
|
||||
# The convenience flags lose to an explicit --cuda-graph-config JSON, which
|
||||
# resolution applies last, so this tool's "no cuda graph" guarantee is
|
||||
# merged into that JSON instead -- an operator serving with their own config
|
||||
# still compiles without capture.
|
||||
explicit = args.cuda_graph_config
|
||||
if isinstance(explicit, CudaGraphConfig):
|
||||
explicit = explicit.to_dict()
|
||||
explicit = dict(explicit or {})
|
||||
for phase in (Phase.DECODE, Phase.PREFILL):
|
||||
phase_config = dict(explicit.get(phase) or {})
|
||||
phase_config["backend"] = Backend.DISABLED
|
||||
explicit[phase] = phase_config
|
||||
args.cuda_graph_config = explicit
|
||||
# Watchdog timeout follows compile_args.timeout because compilation takes long.
|
||||
args.watchdog_timeout = compile_args.timeout
|
||||
args.warmups = "compile-deep-gemm"
|
||||
server_args = ServerArgs.from_cli_args(args)
|
||||
# `cuda_graph_config` is None until resolution parses it.
|
||||
server_args.resolve_once()
|
||||
server_args.cuda_graph_config[Phase.DECODE].backend = Backend.DISABLED
|
||||
server_args.cuda_graph_config[Phase.PREFILL].backend = Backend.DISABLED
|
||||
print(f"Disable CUDA Graph and Torch Compile to save time...")
|
||||
return server_args
|
||||
return ServerArgs.from_cli_args(args)
|
||||
|
||||
|
||||
def run_compile(server_args: ServerArgs, compile_args: CompileArgs):
|
||||
|
||||
@@ -456,13 +456,15 @@ class Runtime:
|
||||
self.endpoint.cache_prefix(prefix)
|
||||
|
||||
def get_tokenizer(self):
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||
|
||||
cfg = resolving_view(self.server_args)
|
||||
return get_tokenizer(
|
||||
self.server_args.tokenizer_path or self.server_args.model_path,
|
||||
tokenizer_mode=self.server_args.tokenizer_mode,
|
||||
trust_remote_code=self.server_args.trust_remote_code,
|
||||
revision=self.server_args.revision,
|
||||
cfg.tokenizer_path or cfg.model_path,
|
||||
tokenizer_mode=cfg.tokenizer_mode,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
revision=cfg.revision,
|
||||
)
|
||||
|
||||
async def async_generate(
|
||||
|
||||
@@ -5,6 +5,7 @@ import os
|
||||
import sys
|
||||
import warnings
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.plugins import load_plugins
|
||||
from sglang.srt.server_args import prepare_server_args
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
@@ -18,10 +19,11 @@ def run_server(server_args):
|
||||
# The flags dispatched on below are decided by resolution (`--grpc-mode`
|
||||
# folds into `smg_grpc_mode`), and `prepare_server_args` returns raw input.
|
||||
server_args.resolve_once()
|
||||
cfg = resolving_view(server_args)
|
||||
|
||||
if server_args.encoder_only:
|
||||
if cfg.encoder_only:
|
||||
# For encoder disaggregation
|
||||
if server_args.smg_grpc_mode or server_args.grpc_mode:
|
||||
if cfg.smg_grpc_mode or cfg.grpc_mode:
|
||||
from sglang.srt.disaggregation.encoder.grpc_server import (
|
||||
serve_grpc_encoder,
|
||||
)
|
||||
@@ -31,7 +33,7 @@ def run_server(server_args):
|
||||
from sglang.srt.disaggregation.encoder.http_server import launch_server
|
||||
|
||||
launch_server(server_args)
|
||||
elif server_args.smg_grpc_mode:
|
||||
elif cfg.smg_grpc_mode:
|
||||
# Legacy SMG gRPC server (--smg-grpc-mode, or the deprecated --grpc-mode
|
||||
# which __post_init__ folds into smg_grpc_mode). The native Rust gRPC
|
||||
# server is a separate path, enabled by --grpc-port, that starts
|
||||
@@ -39,7 +41,7 @@ def run_server(server_args):
|
||||
from sglang.srt.entrypoints.grpc_server import serve_grpc
|
||||
|
||||
asyncio.run(serve_grpc(server_args))
|
||||
elif server_args.use_ray:
|
||||
elif cfg.use_ray:
|
||||
# Ray mode: HTTP mode with Ray backend.
|
||||
try:
|
||||
from sglang.srt.ray.http_server import launch_server
|
||||
|
||||
@@ -221,17 +221,17 @@ def _native_embedding_spec(
|
||||
|
||||
|
||||
def resolved_embedding_plan(
|
||||
spec: EmbeddingModelSpec, *, server_args: Any, model_config: Any
|
||||
spec: EmbeddingModelSpec, *, config: Any, model_config: Any
|
||||
) -> dict[str, Any]:
|
||||
"""Combine static capabilities with the effective server configuration.
|
||||
|
||||
This boundary deliberately accepts duck-typed arguments so the declarative
|
||||
registry remains independent of ServerArgs and ModelConfig import cycles.
|
||||
`config` must answer with the *resolved* configuration -- the readback
|
||||
callers pass `resolving_view(record)`, which is where a decision lives.
|
||||
"""
|
||||
|
||||
prefill_graph = getattr(
|
||||
getattr(server_args, "cuda_graph_config", None), "prefill", None
|
||||
)
|
||||
prefill_graph = getattr(getattr(config, "cuda_graph_config", None), "prefill", None)
|
||||
backend = getattr(prefill_graph, "backend", None)
|
||||
backend_value = getattr(backend, "value", backend)
|
||||
capture_sizes = getattr(prefill_graph, "bs", None) or []
|
||||
@@ -239,7 +239,7 @@ def resolved_embedding_plan(
|
||||
|
||||
return {
|
||||
**spec.as_dict(),
|
||||
"enabled": bool(getattr(server_args, "is_embedding", False)),
|
||||
"enabled": bool(getattr(config, "is_embedding", False)),
|
||||
"supports_dimensions": bool(getattr(model_config, "is_matryoshka", False)),
|
||||
"matryoshka_dimensions": list(
|
||||
getattr(model_config, "matryoshka_dimensions", None) or []
|
||||
@@ -252,14 +252,10 @@ def resolved_embedding_plan(
|
||||
},
|
||||
"cache": {
|
||||
"kv_cache_disabled": bool(
|
||||
getattr(server_args, "prefill_only_disable_kv_cache", False)
|
||||
getattr(config, "prefill_only_disable_kv_cache", False)
|
||||
),
|
||||
"radix_cache_disabled": bool(
|
||||
getattr(server_args, "disable_radix_cache", False)
|
||||
),
|
||||
"chunked_prefill_disabled": getattr(
|
||||
server_args, "chunked_prefill_size", None
|
||||
)
|
||||
"radix_cache_disabled": bool(getattr(config, "disable_radix_cache", False)),
|
||||
"chunked_prefill_disabled": getattr(config, "chunked_prefill_size", None)
|
||||
== -1,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -24,7 +24,12 @@ from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc
|
||||
from sglang.srt.disaggregation.encoder.server import MMEncoder, launch_encoder
|
||||
from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle
|
||||
from sglang.srt.managers.schedule_batch import Modality
|
||||
from sglang.srt.runtime_context import get_disagg, publish
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
publish,
|
||||
)
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import random_uuid
|
||||
from sglang.srt.utils.network import NetworkAddress, get_zmq_socket
|
||||
@@ -212,11 +217,11 @@ async def serve_grpc_encoder(server_args: ServerArgs):
|
||||
dist_init_method = na.to_tcp()
|
||||
else:
|
||||
dist_init_method = NetworkAddress(
|
||||
server_args.host or "127.0.0.1", port_args.nccl_port
|
||||
get_serving().host or "127.0.0.1", port_args.nccl_port
|
||||
).to_tcp()
|
||||
|
||||
send_sockets: List[zmq.Socket] = []
|
||||
for rank in range(1, server_args.tp_size):
|
||||
for rank in range(1, get_parallel().config.tp_size):
|
||||
schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}"
|
||||
send_sockets.append(
|
||||
get_zmq_socket(zmq_ctx, zmq.PUSH, schedule_path, bind=False)
|
||||
@@ -254,7 +259,9 @@ async def serve_grpc_encoder(server_args: ServerArgs):
|
||||
)
|
||||
reflection.enable_server_reflection(SERVICE_NAMES, server)
|
||||
|
||||
listen_addr = NetworkAddress(server_args.host, server_args.port).to_host_port_str()
|
||||
listen_addr = NetworkAddress(
|
||||
get_serving().host, get_serving().port
|
||||
).to_host_port_str()
|
||||
server.add_insecure_port(listen_addr)
|
||||
|
||||
await server.start()
|
||||
|
||||
@@ -113,11 +113,11 @@ def _register_encoder_url_with_bootstrap(server_args: ServerArgs):
|
||||
instead of serialising sleeps in a single thread.
|
||||
"""
|
||||
|
||||
host = server_args.host
|
||||
host = get_serving().host
|
||||
if not host or host in ("0.0.0.0", "::"):
|
||||
host = get_local_ip_auto(server_args.host)
|
||||
host = get_local_ip_auto(get_serving().host)
|
||||
scheme = "https" if server_args.ssl_certfile else "http"
|
||||
encoder_url = NetworkAddress(host, server_args.port).to_url(scheme)
|
||||
encoder_url = NetworkAddress(host, get_serving().port).to_url(scheme)
|
||||
payload = {"url": encoder_url}
|
||||
bootstrap_urls = list(server_args.encoder_register_urls)
|
||||
if not bootstrap_urls:
|
||||
@@ -174,11 +174,11 @@ def _register_encoder_url_with_bootstrap(server_args: ServerArgs):
|
||||
|
||||
|
||||
def _unregister_encoder_url_from_bootstrap(server_args: ServerArgs):
|
||||
host = server_args.host
|
||||
host = get_serving().host
|
||||
if not host or host in ("0.0.0.0", "::"):
|
||||
host = get_local_ip_auto(server_args.host)
|
||||
host = get_local_ip_auto(get_serving().host)
|
||||
scheme = "https" if server_args.ssl_certfile else "http"
|
||||
encoder_url = NetworkAddress(host, server_args.port).to_url(scheme)
|
||||
encoder_url = NetworkAddress(host, get_serving().port).to_url(scheme)
|
||||
payload = {"url": encoder_url}
|
||||
|
||||
for bootstrap_url in server_args.encoder_register_urls:
|
||||
|
||||
@@ -127,7 +127,7 @@ class EncoderPreprocessor:
|
||||
use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get()
|
||||
self.use_image_processor_gpu = (
|
||||
use_image_processor_gpu
|
||||
and resolve_image_processor_backend(server_args) != "pil"
|
||||
and resolve_image_processor_backend(get_mm()) != "pil"
|
||||
)
|
||||
|
||||
self._load_mm_processor(server_args)
|
||||
@@ -158,7 +158,7 @@ class EncoderPreprocessor:
|
||||
def _load_mm_processor(self, server_args: ServerArgs):
|
||||
from transformers import AutoImageProcessor, AutoVideoProcessor
|
||||
|
||||
image_processor_backend = resolve_image_processor_backend(server_args)
|
||||
image_processor_backend = resolve_image_processor_backend(get_mm())
|
||||
image_processor_kwargs = (
|
||||
{}
|
||||
if image_processor_backend == "auto"
|
||||
@@ -167,7 +167,7 @@ class EncoderPreprocessor:
|
||||
try:
|
||||
self.image_processor = AutoImageProcessor.from_pretrained(
|
||||
get_serving().tokenizer_path or get_model().model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
trust_remote_code=get_model().trust_remote_code,
|
||||
revision=server_args.revision,
|
||||
**image_processor_kwargs,
|
||||
)
|
||||
@@ -178,7 +178,7 @@ class EncoderPreprocessor:
|
||||
try:
|
||||
self.video_processor = AutoVideoProcessor.from_pretrained(
|
||||
get_serving().tokenizer_path or get_model().model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
trust_remote_code=get_model().trust_remote_code,
|
||||
revision=server_args.revision,
|
||||
)
|
||||
except Exception as e:
|
||||
@@ -188,7 +188,7 @@ class EncoderPreprocessor:
|
||||
try:
|
||||
_audio_proc = AutoProcessor.from_pretrained(
|
||||
get_serving().tokenizer_path or get_model().model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
trust_remote_code=get_model().trust_remote_code,
|
||||
revision=server_args.revision,
|
||||
)
|
||||
if not hasattr(_audio_proc, "feature_extractor"):
|
||||
|
||||
@@ -35,7 +35,14 @@ from sglang.srt.managers.multimodal_processor import get_mm_processor, import_pr
|
||||
from sglang.srt.managers.schedule_batch import Modality, Req
|
||||
from sglang.srt.multimodal.cache import media_preprocess_kwargs
|
||||
from sglang.srt.multimodal.transport import determine_tensor_transport_mode
|
||||
from sglang.srt.runtime_context import get_disagg, get_exec, get_serving
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_mm,
|
||||
get_model,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import ImageData
|
||||
from sglang.srt.utils.common import safe_pickle_loads
|
||||
@@ -1721,13 +1728,13 @@ class MMReceiverBase(ABC):
|
||||
# When None (e.g. in a scheduler subprocess that has no in-process
|
||||
# bootstrap), fall back to a snapshot of the static --encoder-urls.
|
||||
self.encode_urls: List[str] = (
|
||||
encode_urls if encode_urls is not None else list(server_args.encoder_urls)
|
||||
encode_urls if encode_urls is not None else list(get_disagg().encoder_urls)
|
||||
)
|
||||
self.recv_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get()
|
||||
self.host = get_local_ip_auto(server_args.host)
|
||||
self.host = get_local_ip_auto(get_serving().host)
|
||||
self.pp_rank = pp_rank
|
||||
self.tp_rank = tp_rank
|
||||
self.tp_size = server_args.tp_size
|
||||
self.tp_size = get_parallel().config.tp_size
|
||||
self.tp_group = tp_group
|
||||
self.nnodes = server_args.nnodes
|
||||
self.hostname = get_local_ip_auto()
|
||||
@@ -1836,9 +1843,9 @@ class MMReceiverBase(ABC):
|
||||
_processor = get_processor(
|
||||
get_serving().tokenizer_path,
|
||||
tokenizer_mode=server_args.tokenizer_mode,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
trust_remote_code=get_model().trust_remote_code,
|
||||
revision=server_args.revision,
|
||||
image_processor_backend=resolve_image_processor_backend(server_args),
|
||||
image_processor_backend=resolve_image_processor_backend(get_mm()),
|
||||
**extra_kwargs,
|
||||
)
|
||||
|
||||
@@ -2659,7 +2666,7 @@ def create_mm_receiver(
|
||||
transport_mode = envs.SGLANG_ENCODER_MM_RECEIVER_MODE.get()
|
||||
logger.debug(f"MMReceiver transport_mode from env: {transport_mode}")
|
||||
|
||||
_validate_transport_mode(transport_mode, encode_urls or server_args.encoder_urls)
|
||||
_validate_transport_mode(transport_mode, encode_urls or get_disagg().encoder_urls)
|
||||
logger.info(f"EPD MMReceiver: using transport_mode={transport_mode}")
|
||||
|
||||
receiver_cls = _MM_RECEIVER_BY_MODE.get(transport_mode)
|
||||
|
||||
@@ -1572,10 +1572,10 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher:
|
||||
HTTP uses this entry point today. gRPC can reuse it later without
|
||||
importing HTTP application state or Uvicorn.
|
||||
"""
|
||||
if get_parallel().config.dp_size <= 1 or server_args.tp_size != 1:
|
||||
if get_parallel().config.dp_size <= 1 or get_parallel().config.tp_size != 1:
|
||||
raise ValueError(
|
||||
"Encoder DP mode requires --dp-size > 1 and --tp-size 1; got "
|
||||
f"dp_size={get_parallel().config.dp_size}, tp_size={server_args.tp_size}."
|
||||
f"dp_size={get_parallel().config.dp_size}, tp_size={get_parallel().config.tp_size}."
|
||||
)
|
||||
dp_size = get_parallel().config.dp_size
|
||||
logger.info(f"Launching encoder in DP mode: dp_size={dp_size}")
|
||||
|
||||
@@ -63,6 +63,7 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_mm,
|
||||
get_model,
|
||||
get_parallel,
|
||||
publish,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -450,7 +451,7 @@ class MMEncoder:
|
||||
this instance's value, not a config change, so it travels as an
|
||||
argument."""
|
||||
assert_published(server_args, role="encoder")
|
||||
logger.info(f"init MMEncoder {rank}/{server_args.tp_size}")
|
||||
logger.info(f"init MMEncoder {rank}/{get_parallel().config.tp_size}")
|
||||
self.server_args = server_args
|
||||
configure_media_url_security(
|
||||
get_mm().allowed_media_domains,
|
||||
@@ -470,7 +471,7 @@ class MMEncoder:
|
||||
self.load_config = LoadConfig(
|
||||
load_format=get_model().load_format,
|
||||
download_dir=server_args.download_dir,
|
||||
model_loader_extra_config=server_args.model_loader_extra_config,
|
||||
model_loader_extra_config=get_model().model_loader_extra_config,
|
||||
remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip,
|
||||
remote_instance_weight_loader_seed_instance_service_port=server_args.remote_instance_weight_loader_seed_instance_service_port,
|
||||
remote_instance_weight_loader_send_weights_group_ports=server_args.remote_instance_weight_loader_send_weights_group_ports,
|
||||
@@ -491,12 +492,14 @@ class MMEncoder:
|
||||
|
||||
init_distributed_environment(
|
||||
backend=get_default_distributed_backend(self.device),
|
||||
world_size=server_args.tp_size,
|
||||
world_size=get_parallel().config.tp_size,
|
||||
rank=rank,
|
||||
distributed_init_method=dist_init_method,
|
||||
local_rank=rank,
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=server_args.tp_size)
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=get_parallel().config.tp_size
|
||||
)
|
||||
initialize_dp_attention(server_args, self.model_config)
|
||||
|
||||
self.model = load_model(
|
||||
@@ -554,7 +557,7 @@ class MMEncoder:
|
||||
)
|
||||
self.mm_global_cache = EmbeddingCacheController(
|
||||
rank,
|
||||
server_args.tp_size,
|
||||
get_parallel().config.tp_size,
|
||||
embedding_store=embedding_store,
|
||||
hidden_dims=self._embedding_dims,
|
||||
tp_group=get_tp_group().cpu_group,
|
||||
@@ -1032,7 +1035,7 @@ class MMEncoder:
|
||||
)
|
||||
|
||||
def _broadcast_global_cache_mask(self, mask_tensor: torch.Tensor):
|
||||
if self.server_args.tp_size > 1:
|
||||
if get_parallel().config.tp_size > 1:
|
||||
torch.distributed.broadcast(
|
||||
mask_tensor,
|
||||
src=0,
|
||||
|
||||
@@ -46,7 +46,7 @@ from sglang.srt.disaggregation.utils import (
|
||||
resolve_dcp_dst_entry_indices,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
try:
|
||||
@@ -405,7 +405,8 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
):
|
||||
super().__init__(args, disaggregation_mode, server_args, is_mla_backend)
|
||||
self.transfer_source_rank = (
|
||||
self.kv_args.pp_rank * self.server_args.tp_size + self.kv_args.engine_rank
|
||||
self.kv_args.pp_rank * get_parallel().config.tp_size
|
||||
+ self.kv_args.engine_rank
|
||||
)
|
||||
self.kv_args.kv_data_mem_kinds = _normalize_kv_mem_kinds(
|
||||
getattr(self.kv_args, "kv_data_mem_kinds", None),
|
||||
|
||||
@@ -31,6 +31,7 @@ from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
@@ -187,7 +188,7 @@ def _resolve_dist_init_method(*, server_args: ServerArgs, dist_port: int) -> str
|
||||
dist_init_method = na.to_tcp()
|
||||
else:
|
||||
dist_init_method = NetworkAddress(
|
||||
server_args.host or "127.0.0.1", dist_port
|
||||
get_serving().host or "127.0.0.1", dist_port
|
||||
).to_tcp()
|
||||
return dist_init_method
|
||||
|
||||
|
||||
@@ -423,19 +423,19 @@ def recommended_max_tokens(include_prefill: bool, floor: int = 0) -> int:
|
||||
NCCL. Covers the spec-decode batch plus, if ``include_prefill``, a prefill
|
||||
chunk. Returns ``floor`` if server args are unavailable."""
|
||||
try:
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_schedule, get_spec
|
||||
|
||||
sa = get_server_args()
|
||||
def g(value) -> int:
|
||||
return value if isinstance(value, int) and value > 0 else 0
|
||||
|
||||
def g(name: str) -> int:
|
||||
v = getattr(sa, name, 0)
|
||||
return v if isinstance(v, int) and v > 0 else 0
|
||||
|
||||
tokens = g("max_running_requests") * max(
|
||||
g("speculative_num_draft_tokens"), g("speculative_eagle_topk"), 1
|
||||
schedule, spec = get_schedule(), get_spec()
|
||||
tokens = g(schedule.max_running_requests) * max(
|
||||
g(spec.speculative_num_draft_tokens), g(spec.speculative_eagle_topk), 1
|
||||
)
|
||||
if include_prefill:
|
||||
tokens = max(tokens, g("chunked_prefill_size"), g("max_prefill_tokens"))
|
||||
tokens = max(
|
||||
tokens, g(schedule.chunked_prefill_size), g(schedule.max_prefill_tokens)
|
||||
)
|
||||
return max(tokens, floor)
|
||||
except Exception:
|
||||
return floor
|
||||
|
||||
@@ -48,6 +48,7 @@ import torch
|
||||
import uvloop
|
||||
import zmq
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.elastic_ep.expert_backup_manager import run_expert_backup_manager
|
||||
from sglang.srt.entrypoints.engine_info_bootstrap_server import (
|
||||
EngineInfoBootstrapServer,
|
||||
@@ -98,6 +99,7 @@ from sglang.srt.parser.template_detection import resolve_auto_parsers
|
||||
from sglang.srt.parser.template_manager import TemplateManager
|
||||
from sglang.srt.plugins import load_plugins
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_model,
|
||||
get_parallel,
|
||||
@@ -309,9 +311,9 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
trace_modules=server_args.trace_modules,
|
||||
)
|
||||
thread_label = "Tokenizer"
|
||||
if server_args.disaggregation_mode == "prefill":
|
||||
if get_disagg().disaggregation_mode == "prefill":
|
||||
thread_label = "Prefill Tokenizer"
|
||||
elif server_args.disaggregation_mode == "decode":
|
||||
elif get_disagg().disaggregation_mode == "decode":
|
||||
thread_label = "Decode Tokenizer"
|
||||
trace_set_thread_info(thread_label)
|
||||
|
||||
@@ -1056,10 +1058,8 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
|
||||
# Needs a tokenizer and a chat template, so it cannot live in the
|
||||
# pipeline; after the plugins, which may register the parser detected.
|
||||
if (
|
||||
server_args.reasoning_parser == "auto"
|
||||
or server_args.tool_call_parser == "auto"
|
||||
):
|
||||
parsers = resolving_view(server_args)
|
||||
if parsers.reasoning_parser == "auto" or parsers.tool_call_parser == "auto":
|
||||
resolve_auto_parsers(server_args)
|
||||
|
||||
# This publish replaces whatever was published before it, so the
|
||||
@@ -1627,27 +1627,28 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
|
||||
|
||||
def _set_envs_and_config(server_args: ServerArgs):
|
||||
cfg = resolving_view(server_args)
|
||||
# Set global environments
|
||||
# MNNVL fabric (GB200/GB300) multi-node: cross-node NVLink needs NCCL's
|
||||
# cuMem-based buffers and MNNVL transport. Default them on (user-set
|
||||
# values win; the symm-mem override below only fires when unset).
|
||||
if server_args.nnodes > 1 and is_mnnvl_fabric_device():
|
||||
if cfg.nnodes > 1 and is_mnnvl_fabric_device():
|
||||
os.environ.setdefault("NCCL_CUMEM_ENABLE", "1")
|
||||
os.environ.setdefault("NCCL_MNNVL_ENABLE", "1")
|
||||
if "NCCL_CUMEM_ENABLE" not in os.environ or server_args.enable_symm_mem:
|
||||
os.environ["NCCL_CUMEM_ENABLE"] = str(int(server_args.enable_symm_mem))
|
||||
if "NCCL_CUMEM_ENABLE" not in os.environ or cfg.enable_symm_mem:
|
||||
os.environ["NCCL_CUMEM_ENABLE"] = str(int(cfg.enable_symm_mem))
|
||||
if (
|
||||
"NCCL_NVLS_ENABLE" not in os.environ
|
||||
or server_args.enable_nccl_nvls
|
||||
or server_args.enable_symm_mem
|
||||
or cfg.enable_nccl_nvls
|
||||
or cfg.enable_symm_mem
|
||||
):
|
||||
os.environ["NCCL_NVLS_ENABLE"] = str(
|
||||
int(server_args.enable_nccl_nvls or server_args.enable_symm_mem)
|
||||
int(cfg.enable_nccl_nvls or cfg.enable_symm_mem)
|
||||
)
|
||||
if "NCCL_GRAPH_MIXING_SUPPORT" not in os.environ or server_args.enable_symm_mem:
|
||||
if "NCCL_GRAPH_MIXING_SUPPORT" not in os.environ or cfg.enable_symm_mem:
|
||||
# Note(wh): NCCL_GRAPH_MIXING_SUPPORT=0 can help improve performance for symmetric kernels.
|
||||
# details in https://github.com/NVIDIA/nccl-tests/issues/333#issuecomment-3103636985
|
||||
if server_args.dcp_size > 1:
|
||||
if cfg.dcp_size > 1:
|
||||
os.environ["NCCL_GRAPH_MIXING_SUPPORT"] = "0"
|
||||
os.environ["CUDA_DEVICE_MAX_CONNECTIONS"] = "8"
|
||||
|
||||
@@ -1669,7 +1670,7 @@ def _set_envs_and_config(server_args: ServerArgs):
|
||||
)
|
||||
|
||||
# Set prometheus env vars
|
||||
if server_args.enable_metrics:
|
||||
if cfg.enable_metrics:
|
||||
set_prometheus_multiproc_dir()
|
||||
|
||||
# Set ulimit
|
||||
@@ -1677,7 +1678,7 @@ def _set_envs_and_config(server_args: ServerArgs):
|
||||
|
||||
# Check flashinfer version
|
||||
if not get_bool_env_var("SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK"):
|
||||
if "flashinfer" in server_args.get_attention_backends():
|
||||
if "flashinfer" in cfg.get_attention_backends():
|
||||
assert_pkg_version(
|
||||
"flashinfer_python",
|
||||
"0.6.17",
|
||||
@@ -1694,7 +1695,7 @@ def _set_envs_and_config(server_args: ServerArgs):
|
||||
|
||||
# Signal handlers can only be registered from the main thread.
|
||||
if threading.current_thread() is threading.main_thread():
|
||||
if server_args.custom_sigquit_handler is None:
|
||||
if cfg.custom_sigquit_handler is None:
|
||||
# Register the signal handler.
|
||||
# The child processes will send SIGQUIT to this process when any error happens
|
||||
# This process then clean up the whole process tree
|
||||
@@ -1709,10 +1710,8 @@ def _set_envs_and_config(server_args: ServerArgs):
|
||||
signal.signal(signal.SIGQUIT, launch_phase_sigquit_handler)
|
||||
else:
|
||||
# Allow users to register a custom SIGQUIT handler for things like crash dump
|
||||
logger.error(
|
||||
f"Using custom SIGQUIT handler: {server_args.custom_sigquit_handler}"
|
||||
)
|
||||
signal.signal(signal.SIGQUIT, server_args.custom_sigquit_handler)
|
||||
logger.error(f"Using custom SIGQUIT handler: {cfg.custom_sigquit_handler}")
|
||||
signal.signal(signal.SIGQUIT, cfg.custom_sigquit_handler)
|
||||
else:
|
||||
logger.warning(
|
||||
"Signal handler is not added because the engine is not in the "
|
||||
@@ -1724,7 +1723,7 @@ def _set_envs_and_config(server_args: ServerArgs):
|
||||
mp.set_start_method("spawn", force=True)
|
||||
|
||||
# Set gc threshold
|
||||
if gc_threshold := server_args.gc_threshold:
|
||||
if gc_threshold := cfg.gc_threshold:
|
||||
gc.set_threshold(*gc_threshold)
|
||||
|
||||
_log_legacy_kernel_cache_dirs()
|
||||
|
||||
@@ -14,6 +14,7 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.configs.embedding_model_spec import resolved_embedding_plan
|
||||
from sglang.srt.runtime_context import (
|
||||
get_lora,
|
||||
@@ -416,7 +417,7 @@ class RuntimeHandle:
|
||||
if embedding_model_spec is not None:
|
||||
result["embedding"] = resolved_embedding_plan(
|
||||
embedding_model_spec,
|
||||
server_args=self.server_args,
|
||||
config=resolving_view(self.server_args),
|
||||
model_config=model_config,
|
||||
)
|
||||
return json.dumps(result, default=str)
|
||||
|
||||
@@ -165,6 +165,15 @@ async def serve_grpc(server_args, model_info=None):
|
||||
"version mismatch — see the chained exception above for details."
|
||||
) from e
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
# The integrated servicer builds an `Engine`, which validates and publishes
|
||||
# on its own. Validating here would run `check_server_args` twice, and the
|
||||
# LoRA normalization is not idempotent -- the second pass sees the `LoRARef`
|
||||
# objects the first one declared and rejects them. So this entry reads the
|
||||
# declarations for what it needs before the engine exists.
|
||||
cfg = resolving_view(server_args)
|
||||
|
||||
sidecar_app = web.Application()
|
||||
sidecar_runner = None
|
||||
sidecar_port = (
|
||||
@@ -176,7 +185,7 @@ async def serve_grpc(server_args, model_info=None):
|
||||
# Metrics setup: must set PROMETHEUS_MULTIPROC_DIR before scheduler
|
||||
# processes import prometheus_client, since the env var is inherited
|
||||
# at fork time.
|
||||
if server_args.enable_metrics:
|
||||
if cfg.enable_metrics:
|
||||
try:
|
||||
from sglang.srt.observability.func_timer import enable_func_timer
|
||||
from sglang.srt.utils import set_prometheus_multiproc_dir
|
||||
@@ -204,7 +213,7 @@ async def serve_grpc(server_args, model_info=None):
|
||||
)
|
||||
try:
|
||||
sidecar_runner = await _start_sidecar_server(
|
||||
server_args.host, sidecar_port, sidecar_app
|
||||
cfg.host, sidecar_port, sidecar_app
|
||||
)
|
||||
except OSError as e:
|
||||
logger.error(
|
||||
@@ -232,7 +241,7 @@ async def serve_grpc(server_args, model_info=None):
|
||||
)
|
||||
if sidecar_supported:
|
||||
serve_kwargs["on_request_manager_ready"] = _on_request_manager_ready
|
||||
elif server_args.enable_metrics:
|
||||
elif cfg.enable_metrics:
|
||||
# User explicitly asked for metrics but the installed servicer can't
|
||||
# start the sidecar that serves them — fail loud rather than silently
|
||||
# produce a server with no /metrics endpoint.
|
||||
|
||||
@@ -63,6 +63,7 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import ORJSONResponse, Response, StreamingResponse
|
||||
from fastapi.routing import APIRoute
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.configs.embedding_model_spec import resolved_embedding_plan
|
||||
from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX
|
||||
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode
|
||||
@@ -284,7 +285,7 @@ async def lifespan(fast_api_app: FastAPI):
|
||||
thread_label = f"MultiTokenizer-{_global_state.tokenizer_manager.worker_id}"
|
||||
|
||||
# Add prometheus middleware
|
||||
if server_args.enable_metrics:
|
||||
if get_observability().enable_metrics:
|
||||
add_prometheus_middleware(app)
|
||||
enable_func_timer()
|
||||
|
||||
@@ -488,6 +489,7 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_lora,
|
||||
get_model,
|
||||
get_observability,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
publish,
|
||||
@@ -771,7 +773,7 @@ async def model_info():
|
||||
if embedding_model_spec is not None:
|
||||
result["embedding"] = resolved_embedding_plan(
|
||||
embedding_model_spec,
|
||||
server_args=_global_state.tokenizer_manager.server_args,
|
||||
config=resolving_view(_global_state.tokenizer_manager.server_args),
|
||||
model_config=model_config,
|
||||
)
|
||||
return result
|
||||
@@ -2540,7 +2542,7 @@ def _setup_and_run_http_server(
|
||||
if tokenizer_manager is not None:
|
||||
tokenizer_manager._subprocess_watchdog = subprocess_watchdog
|
||||
|
||||
if server_args.enable_metrics:
|
||||
if get_observability().enable_metrics:
|
||||
add_prometheus_track_response_middleware(app)
|
||||
|
||||
# Pass additional arguments to the lifespan function.
|
||||
@@ -2602,12 +2604,13 @@ def _setup_and_run_http_server(
|
||||
if server_args.enable_http2:
|
||||
logger.info(
|
||||
f"Starting embedded Granian HTTP/2 server on "
|
||||
f"{server_args.host}:{server_args.port}"
|
||||
f"{get_serving().host}:{get_serving().port}"
|
||||
)
|
||||
_run_granian_server(
|
||||
host=server_args.host,
|
||||
port=server_args.port,
|
||||
log_level=server_args.log_level_http or server_args.log_level,
|
||||
host=get_serving().host,
|
||||
port=get_serving().port,
|
||||
log_level=get_observability().log_level_http
|
||||
or get_observability().log_level,
|
||||
http2_max_concurrent_streams=(
|
||||
server_args.http2_max_concurrent_streams
|
||||
),
|
||||
@@ -2621,10 +2624,11 @@ def _setup_and_run_http_server(
|
||||
# Use Config/Server API for access to the SSLContext.
|
||||
config = uvicorn.Config(
|
||||
app,
|
||||
host=server_args.host,
|
||||
port=server_args.port,
|
||||
host=get_serving().host,
|
||||
port=get_serving().port,
|
||||
root_path=server_args.fastapi_root_path,
|
||||
log_level=server_args.log_level_http or server_args.log_level,
|
||||
log_level=get_observability().log_level_http
|
||||
or get_observability().log_level,
|
||||
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
|
||||
loop="uvloop",
|
||||
ssl_keyfile=server_args.ssl_keyfile,
|
||||
@@ -2658,10 +2662,11 @@ def _setup_and_run_http_server(
|
||||
# Default case, one tokenizer process
|
||||
uvicorn.run(
|
||||
app,
|
||||
host=server_args.host,
|
||||
port=server_args.port,
|
||||
host=get_serving().host,
|
||||
port=get_serving().port,
|
||||
root_path=server_args.fastapi_root_path,
|
||||
log_level=server_args.log_level_http or server_args.log_level,
|
||||
log_level=get_observability().log_level_http
|
||||
or get_observability().log_level,
|
||||
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
|
||||
loop="uvloop",
|
||||
ssl_keyfile=server_args.ssl_keyfile,
|
||||
@@ -2689,12 +2694,13 @@ def _setup_and_run_http_server(
|
||||
if server_args.enable_http2:
|
||||
logger.info(
|
||||
f"Starting embedded Granian HTTP/2 server on "
|
||||
f"{server_args.host}:{server_args.port}"
|
||||
f"{get_serving().host}:{get_serving().port}"
|
||||
)
|
||||
_run_granian_server(
|
||||
host=server_args.host,
|
||||
port=server_args.port,
|
||||
log_level=server_args.log_level_http or server_args.log_level,
|
||||
host=get_serving().host,
|
||||
port=get_serving().port,
|
||||
log_level=get_observability().log_level_http
|
||||
or get_observability().log_level,
|
||||
http2_max_concurrent_streams=(
|
||||
server_args.http2_max_concurrent_streams
|
||||
),
|
||||
@@ -2707,10 +2713,11 @@ def _setup_and_run_http_server(
|
||||
else:
|
||||
uvicorn.run(
|
||||
"sglang.srt.entrypoints.http_server:app",
|
||||
host=server_args.host,
|
||||
port=server_args.port,
|
||||
host=get_serving().host,
|
||||
port=get_serving().port,
|
||||
root_path=server_args.fastapi_root_path,
|
||||
log_level=server_args.log_level_http or server_args.log_level,
|
||||
log_level=get_observability().log_level_http
|
||||
or get_observability().log_level,
|
||||
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
|
||||
timeout_worker_healthcheck=envs.SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT.get(),
|
||||
loop="uvloop",
|
||||
@@ -2748,12 +2755,12 @@ def _start_native_grpc_server_for_runtime(
|
||||
)
|
||||
|
||||
grpc_handle = grpc_native.start_server(
|
||||
host=server_args.host,
|
||||
host=get_serving().host,
|
||||
port=grpc_port,
|
||||
runtime_handle=runtime_handle,
|
||||
worker_threads=server_args.grpc_worker_threads,
|
||||
)
|
||||
logger.info(f"Native gRPC server started on {server_args.host}:{grpc_port}")
|
||||
logger.info(f"Native gRPC server started on {get_serving().host}:{grpc_port}")
|
||||
return grpc_handle
|
||||
|
||||
|
||||
|
||||
@@ -118,7 +118,7 @@ def start_sidecar(server_args) -> Sidecar:
|
||||
module_name = server_args.sidecar
|
||||
assert module_name is not None
|
||||
sidecar_args, shutdown_timeout = _parse_sidecar_args(server_args.sidecar_args)
|
||||
endpoint = build_sidecar_endpoint(server_args.host, get_serving().grpc_port)
|
||||
endpoint = build_sidecar_endpoint(get_serving().host, get_serving().grpc_port)
|
||||
proc = mp.get_context("spawn").Process(
|
||||
name=f"sglang_sidecar_{module_name}",
|
||||
target=_run_sidecar,
|
||||
|
||||
@@ -146,7 +146,7 @@ async def get_loads(
|
||||
"version": __version__,
|
||||
"accelerator": _accelerator_name(),
|
||||
"num_accelerators": _num_accelerators_per_dp_rank(
|
||||
tokenizer_manager.server_args.tp_size,
|
||||
get_parallel().config.tp_size,
|
||||
get_parallel().config.pp_size,
|
||||
get_parallel().config.dp_size,
|
||||
get_parallel().config.enable_dp_attention,
|
||||
|
||||
@@ -764,7 +764,7 @@ class _UtilizationRateAccumulatorMixin(_Accumulator):
|
||||
single_pass_global_physical_count,
|
||||
num_gpu=self._expert_location_metadata.ep_size,
|
||||
)
|
||||
gpu_physical_count = gpu_physical_count.to(self._server_args.device)
|
||||
gpu_physical_count = gpu_physical_count.to(get_device_namespace().device)
|
||||
torch.distributed.reduce(
|
||||
gpu_physical_count, dst=0, op=torch.distributed.ReduceOp.SUM
|
||||
)
|
||||
@@ -898,9 +898,9 @@ class _StatAccumulator(_UtilizationRateAccumulatorMixin):
|
||||
# Cannot use local_physical_count to support select_experts
|
||||
self._expert_location_metadata.num_physical_experts,
|
||||
),
|
||||
buffer_size=self._server_args.expert_distribution_recorder_buffer_size,
|
||||
buffer_size=get_exec().moe.expert_distribution_recorder_buffer_size,
|
||||
dtype=torch.int32,
|
||||
device=self._server_args.device,
|
||||
device=get_device_namespace().device,
|
||||
)
|
||||
self._first_dump = True
|
||||
|
||||
|
||||
@@ -22,6 +22,7 @@ import torch
|
||||
from torch.profiler import ProfilerActivity, profile
|
||||
|
||||
from sglang.srt.model_executor.runner import DecodeCudaGraphRunner
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import register_xpu_device_properties_for_dynamo
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -116,7 +117,7 @@ class XPUGraphRunner(DecodeCudaGraphRunner):
|
||||
|
||||
def __init__(self, model_runner: ModelRunner):
|
||||
assert (
|
||||
not model_runner.server_args.enable_memory_saver
|
||||
not get_exec().features.enable_memory_saver
|
||||
), "XPUGraphRunner does not support Torch Memory Saver yet."
|
||||
register_fake_ops()
|
||||
self._apply_xpu_compile_config()
|
||||
|
||||
@@ -112,10 +112,10 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
|
||||
self.device = model_runner.device
|
||||
self.max_context_len = model_runner.model_config.context_len
|
||||
self.num_heads = model_runner.model_config.get_num_attention_heads(
|
||||
model_runner.server_args.tp_size
|
||||
get_parallel().tp_size
|
||||
)
|
||||
self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
|
||||
model_runner.server_args.tp_size
|
||||
get_parallel().tp_size
|
||||
)
|
||||
self.head_size = model_runner.model_config.head_dim
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.configs.model_config import (
|
||||
get_minimax_sparse_attention_config,
|
||||
get_minimax_sparse_disable_value_layer_ids,
|
||||
@@ -116,7 +117,9 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
self.max_context_len = int(runner.model_config.context_len)
|
||||
# Per-forward cache for the native decode block table (rebuilt each forward).
|
||||
self._native_decode_bt: dict = {}
|
||||
self.fp8_attn_gemm = m3_fp8_attn_gemm_enabled(runner.server_args)
|
||||
self.fp8_attn_gemm = m3_fp8_attn_gemm_enabled(
|
||||
resolving_view(runner.server_args)
|
||||
)
|
||||
if self.fp8_attn_gemm:
|
||||
assert self.kv_pool.main_pool.dtype == torch.float8_e4m3fn, (
|
||||
"fp8 attn-GEMM mode requires an fp8_e4m3fn main KV pool, got "
|
||||
|
||||
@@ -33,7 +33,6 @@ from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
|
||||
class ContextParallelStrategyKind(IntEnum):
|
||||
@@ -237,23 +236,27 @@ def _is_dsa_active() -> bool:
|
||||
_STRATEGY: Optional[ContextParallelStrategy] = None
|
||||
|
||||
|
||||
def init_cp_strategy(server_args: ServerArgs) -> None:
|
||||
"""Bind the configured CP strategy for this process."""
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
def init_cp_strategy(
|
||||
*, enable_prefill_cp: bool, cp_size: int, cp_strategy: str
|
||||
) -> None:
|
||||
"""Bind the CP strategy for this process.
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
Takes the three values: resolution calls this from inside `__post_init__`,
|
||||
where the bags do not exist yet, and `get_cp_strategy` calls it lazily in a
|
||||
worker, which reads them off the published bags. Each caller reads from the
|
||||
source it has.
|
||||
"""
|
||||
global _STRATEGY
|
||||
|
||||
if not cfg.enable_prefill_cp:
|
||||
if not enable_prefill_cp:
|
||||
_STRATEGY = None
|
||||
return
|
||||
|
||||
cp_size = cfg.attn_cp_size
|
||||
if cp_size <= 1:
|
||||
_STRATEGY = None
|
||||
return
|
||||
|
||||
kind = ContextParallelStrategyKind.from_string(cfg.cp_strategy)
|
||||
kind = ContextParallelStrategyKind.from_string(cp_strategy)
|
||||
if kind == ContextParallelStrategyKind.ZIGZAG:
|
||||
from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy
|
||||
|
||||
@@ -264,8 +267,7 @@ def init_cp_strategy(server_args: ServerArgs) -> None:
|
||||
_STRATEGY = InterleaveCPStrategy(cp_size=cp_size)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported cp_strategy kind {kind} for "
|
||||
f"cp_strategy={cfg.cp_strategy!r}"
|
||||
f"Unsupported cp_strategy kind {kind} for cp_strategy={cp_strategy!r}"
|
||||
)
|
||||
|
||||
|
||||
@@ -280,14 +282,16 @@ def get_cp_strategy() -> Optional[ContextParallelStrategy]:
|
||||
global _STRATEGY
|
||||
|
||||
if _STRATEGY is None:
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
try:
|
||||
server_args = get_server_args()
|
||||
parallel = get_parallel().config
|
||||
except ValueError:
|
||||
return None
|
||||
if server_args is not None and get_parallel().config.enable_prefill_cp:
|
||||
init_cp_strategy(server_args)
|
||||
if parallel.enable_prefill_cp:
|
||||
init_cp_strategy(
|
||||
enable_prefill_cp=True,
|
||||
cp_size=parallel.attn_cp_size,
|
||||
cp_strategy=parallel.cp_strategy,
|
||||
)
|
||||
return _STRATEGY
|
||||
|
||||
|
||||
|
||||
@@ -63,7 +63,7 @@ def is_glm_dsa_cache_layer_split_enabled(model_runner: "ModelRunner") -> bool:
|
||||
|
||||
return (
|
||||
not model_runner.is_draft_worker
|
||||
and model_runner.server_args.enable_dsa_cache_layer_split
|
||||
and get_parallel().config.enable_dsa_cache_layer_split
|
||||
and model_runner.use_mla_backend
|
||||
and is_deepseek_dsa(model_runner.model_config.hf_config)
|
||||
)
|
||||
|
||||
@@ -13,7 +13,7 @@ from sglang.srt.distributed import (
|
||||
get_tp_group,
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state import in_the_same_node_as
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils import (
|
||||
ceil_align,
|
||||
get_cuda_driver_bindings,
|
||||
@@ -75,12 +75,16 @@ def _resolve_backend(backend: str, is_multi_node: bool = False) -> str:
|
||||
return backend
|
||||
|
||||
|
||||
def resolve_flashinfer_allreduce_fusion_backend(server_args) -> Optional[str]:
|
||||
backend = getattr(server_args, "flashinfer_allreduce_fusion_backend", None)
|
||||
def resolve_flashinfer_allreduce_fusion_backend() -> Optional[str]:
|
||||
"""The fusion backend for this process, or None when fusion is off.
|
||||
|
||||
Reads the published leaves (`exec.comm`, `parallel`): the backend is a
|
||||
resolution decision, and the node count is launch topology.
|
||||
"""
|
||||
backend = get_exec().comm.flashinfer_allreduce_fusion_backend
|
||||
if backend is None:
|
||||
return None
|
||||
is_multi_node = getattr(server_args, "nnodes", 1) > 1
|
||||
return _resolve_backend(backend, is_multi_node)
|
||||
return _resolve_backend(backend, get_parallel().config.nnodes > 1)
|
||||
|
||||
|
||||
if is_flashinfer_available():
|
||||
@@ -716,8 +720,7 @@ def ensure_workspace_initialized(
|
||||
token_num = token_num or max_token_num
|
||||
group_key = (device_group, cpu_group)
|
||||
effective_dtype = dtype or torch.bfloat16
|
||||
server_args = get_server_args()
|
||||
backend = resolve_flashinfer_allreduce_fusion_backend(server_args)
|
||||
backend = resolve_flashinfer_allreduce_fusion_backend()
|
||||
if backend is None:
|
||||
return False
|
||||
|
||||
|
||||
@@ -4,7 +4,6 @@ import logging
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from enum import Enum, IntEnum
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
@@ -12,14 +11,18 @@ from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_exec, get_flags, get_forward, get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_forward,
|
||||
get_model,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.utils import is_cuda, is_npu
|
||||
|
||||
_is_npu = is_npu()
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.utils.common import log_info_on_rank0
|
||||
|
||||
@@ -308,37 +311,47 @@ def get_ascend_dispatcher_output_dtype(dispatcher):
|
||||
return DispatcherOutputDtype.BF16
|
||||
|
||||
|
||||
def initialize_moe_config(server_args: ServerArgs):
|
||||
def initialize_moe_config():
|
||||
"""Seed the MoE runtime flags from the published configuration.
|
||||
|
||||
Reads the bags: `moe_a2a_backend` and its siblings are resolution's
|
||||
answers, and the record carries them only while declarations materialize
|
||||
onto it. Called once per process after publish
|
||||
(scheduler init, the benchmark work functions).
|
||||
"""
|
||||
exec_moe = get_exec().moe
|
||||
overlap = get_exec().overlap
|
||||
spec = get_spec()
|
||||
moe = get_flags().moe
|
||||
moe.a2a_backend = MoeA2ABackend(server_args.moe_a2a_backend)
|
||||
moe.runner_backend = MoeRunnerBackend(server_args.moe_runner_backend)
|
||||
moe.a2a_backend = MoeA2ABackend(exec_moe.moe_a2a_backend)
|
||||
moe.runner_backend = MoeRunnerBackend(exec_moe.moe_runner_backend)
|
||||
moe.speculative_runner_backend = (
|
||||
MoeRunnerBackend(server_args.speculative_moe_runner_backend)
|
||||
if server_args.speculative_moe_runner_backend is not None
|
||||
MoeRunnerBackend(spec.speculative_moe_runner_backend)
|
||||
if spec.speculative_moe_runner_backend is not None
|
||||
else moe.runner_backend
|
||||
)
|
||||
moe.speculative_a2a_backend = (
|
||||
MoeA2ABackend(server_args.speculative_moe_a2a_backend)
|
||||
if server_args.speculative_moe_a2a_backend is not None
|
||||
MoeA2ABackend(spec.speculative_moe_a2a_backend)
|
||||
if spec.speculative_moe_a2a_backend is not None
|
||||
else moe.a2a_backend
|
||||
)
|
||||
moe.deepep_mode = DeepEPMode(server_args.deepep_mode)
|
||||
moe.deepep_config = server_args.deepep_config or ""
|
||||
moe.tbo_enabled = server_args.enable_two_batch_overlap
|
||||
moe.sbo_enabled = server_args.enable_single_batch_overlap
|
||||
moe.deepep_mode = DeepEPMode(exec_moe.deepep_mode)
|
||||
moe.deepep_config = exec_moe.deepep_config or ""
|
||||
moe.tbo_enabled = overlap.enable_two_batch_overlap
|
||||
moe.sbo_enabled = overlap.enable_single_batch_overlap
|
||||
if moe.sbo_enabled and is_cuda():
|
||||
if torch.cuda.get_device_capability()[0] == 9:
|
||||
raise ValueError(
|
||||
"SBO (single batch overlap) is not supported on SM90 GPUs with latest sgl-deep-gemm wheel. Please try removing --enable-single-batch-overlap argument."
|
||||
)
|
||||
moe.tbo_token_distribution_threshold = server_args.tbo_token_distribution_threshold
|
||||
moe.disable_fp4_allgather = server_args.disable_flashinfer_cutlass_moe_fp4_allgather
|
||||
moe.quantization = server_args.quantization
|
||||
moe.tbo_token_distribution_threshold = overlap.tbo_token_distribution_threshold
|
||||
moe.disable_fp4_allgather = exec_moe.disable_flashinfer_cutlass_moe_fp4_allgather
|
||||
moe.quantization = get_model().quantization
|
||||
# Seeded with the user's intent; each model's gate refines the ACTIVE
|
||||
# value for its own build (install_shared_experts_fusion_decision).
|
||||
moe.disable_shared_experts_fusion = server_args.disable_shared_experts_fusion
|
||||
moe.disable_shared_experts_fusion = exec_moe.disable_shared_experts_fusion
|
||||
moe.speculative_disable_shared_experts_fusion = (
|
||||
server_args.disable_shared_experts_fusion
|
||||
exec_moe.disable_shared_experts_fusion
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -2,10 +2,11 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils.common import (
|
||||
get_device_capability,
|
||||
is_cuda,
|
||||
@@ -13,9 +14,6 @@ from sglang.srt.utils.common import (
|
||||
)
|
||||
from sglang.srt.utils.custom_op import register_custom_op_from_extern
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -142,11 +140,11 @@ class Fp4GemmRunnerBackend(Enum):
|
||||
FP4_GEMM_RUNNER_BACKEND: Fp4GemmRunnerBackend | None = None
|
||||
|
||||
|
||||
def initialize_fp4_gemm_config(server_args: ServerArgs) -> None:
|
||||
"""Initialize FP4 GEMM configuration from server args."""
|
||||
def initialize_fp4_gemm_config() -> None:
|
||||
"""Initialize the FP4 GEMM backend from the published configuration."""
|
||||
global FP4_GEMM_RUNNER_BACKEND
|
||||
|
||||
backend = server_args.fp4_gemm_runner_backend
|
||||
backend = get_exec().kernel.fp4_gemm_runner_backend
|
||||
if backend == "auto":
|
||||
if is_sm100_supported():
|
||||
backend = "flashinfer_cutedsl"
|
||||
|
||||
@@ -3,23 +3,10 @@ from __future__ import annotations
|
||||
import logging
|
||||
from enum import Enum
|
||||
from functools import lru_cache, partial
|
||||
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union
|
||||
from typing import Callable, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
sglang_per_token_group_quant_fp8,
|
||||
sglang_per_token_group_quant_fp8_row_padded,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils.common import torch_release
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
fp8_dtype,
|
||||
fp8_max,
|
||||
@@ -28,12 +15,18 @@ from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
is_fp8_fnuz,
|
||||
per_token_group_quant_fp8,
|
||||
scaled_fp8_quant,
|
||||
sglang_per_token_group_quant_fp8,
|
||||
sglang_per_token_group_quant_fp8_row_padded,
|
||||
sglang_per_token_quant_fp8,
|
||||
static_quant_fp8,
|
||||
triton_scaled_mm,
|
||||
w8a8_block_fp8_matmul_deepgemm,
|
||||
w8a8_block_fp8_matmul_triton,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils import (
|
||||
ceil_align,
|
||||
ceil_div,
|
||||
@@ -54,6 +47,7 @@ from sglang.srt.utils import (
|
||||
is_xpu,
|
||||
offloader,
|
||||
)
|
||||
from sglang.srt.utils.common import torch_release
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -799,11 +793,11 @@ def _dispatch_auto_backend() -> Callable:
|
||||
return triton_w8a8_block_fp8_linear
|
||||
|
||||
|
||||
def initialize_fp8_gemm_config(server_args: ServerArgs) -> None:
|
||||
def initialize_fp8_gemm_config() -> None:
|
||||
"""Initialize FP8 GEMM configuration."""
|
||||
global FP8_GEMM_RUNNER_BACKEND
|
||||
|
||||
backend = server_args.fp8_gemm_runner_backend
|
||||
backend = get_exec().kernel.fp8_gemm_runner_backend
|
||||
if backend == "auto" and is_sm120_supported():
|
||||
backend = "cutlass"
|
||||
|
||||
|
||||
@@ -39,7 +39,13 @@ from sglang.srt.managers.io_struct import (
|
||||
)
|
||||
from sglang.srt.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin
|
||||
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
|
||||
from sglang.srt.runtime_context import get_device, get_serving, publish
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_model,
|
||||
get_observability,
|
||||
get_serving,
|
||||
publish,
|
||||
)
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import configure_logger, freeze_gc, kill_itself_when_parent_died
|
||||
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||
@@ -130,7 +136,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
||||
self.tokenizer = get_tokenizer(
|
||||
get_serving().tokenizer_path,
|
||||
tokenizer_mode=server_args.tokenizer_mode,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
trust_remote_code=get_model().trust_remote_code,
|
||||
revision=server_args.revision,
|
||||
tokenizer_backend=server_args.tokenizer_backend,
|
||||
)
|
||||
@@ -151,7 +157,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
||||
test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(),
|
||||
)
|
||||
|
||||
if server_args.enable_metrics:
|
||||
if get_observability().enable_metrics:
|
||||
start_cpu_monitor_thread("detokenizer")
|
||||
|
||||
def init_request_dispatcher(self):
|
||||
|
||||
@@ -20,6 +20,7 @@ from typing import TYPE_CHECKING, Any, Dict, FrozenSet, List, Optional, Tuple
|
||||
|
||||
import msgspec
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.io_struct import TokenizedGenerateReqInput
|
||||
from sglang.srt.managers.utils import (
|
||||
@@ -29,6 +30,8 @@ from sglang.srt.managers.utils import (
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_mm,
|
||||
get_observability,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.utils.flatten import (
|
||||
@@ -272,7 +275,7 @@ class NativeMmHost:
|
||||
)
|
||||
|
||||
return (
|
||||
self.server_args.tp_size > 1
|
||||
get_parallel().config.tp_size > 1
|
||||
and determine_tensor_transport_mode() != "default"
|
||||
and not self.server_args.skip_tokenizer_init
|
||||
)
|
||||
@@ -395,13 +398,13 @@ class RustServer:
|
||||
"ingress has no equivalent). Launch without SGLANG_RUST_SERVER, or "
|
||||
"drop --preferred-sampling-params and send those values per request."
|
||||
)
|
||||
http_addr = f"{server_args.host}:{server_args.port}"
|
||||
http_addr = f"{get_serving().host}:{server_args.port}"
|
||||
|
||||
# Per-DP-rank HTTP port with client load balancing. `None` when DP is off,
|
||||
# so the rank is not conflated with rank 0 of a one-rank group.
|
||||
dp_rank = scheduler.ps.attn_dp_rank if scheduler.ps.dp_size > 1 else None
|
||||
if dp_rank is not None:
|
||||
http_addr = f"{server_args.host}:{server_args.port + dp_rank}"
|
||||
http_addr = f"{get_serving().host}:{server_args.port + dp_rank}"
|
||||
|
||||
launch_cores, server_cores = cls._partition_cores(
|
||||
mm_workers=(
|
||||
@@ -754,7 +757,7 @@ class RustServer:
|
||||
|
||||
ext = load_rust_extension("sglang.srt.rust_extensions._server")
|
||||
|
||||
sa = scheduler.server_args
|
||||
sa = resolving_view(scheduler.server_args)
|
||||
mc = scheduler.model_config
|
||||
disaggregation_mode = {
|
||||
"null": ext.DisaggregationMode.Null,
|
||||
@@ -768,9 +771,9 @@ class RustServer:
|
||||
revision=sa.revision,
|
||||
load_format=sa.load_format,
|
||||
weight_version=sa.weight_version,
|
||||
host=sa.host,
|
||||
host=get_serving().host,
|
||||
port=sa.port,
|
||||
log_level=sa.log_level,
|
||||
log_level=get_observability().log_level,
|
||||
log_level_http=sa.log_level_http,
|
||||
chat_template=sa.chat_template,
|
||||
tool_call_parser=sa.tool_call_parser,
|
||||
|
||||
@@ -903,11 +903,11 @@ class Scheduler(
|
||||
"moe_topk",
|
||||
)
|
||||
if any(hasattr(config_to_check, attr) for attr in moe_topk_attrs):
|
||||
initialize_moe_config(self.server_args)
|
||||
initialize_moe_config()
|
||||
|
||||
# Initialize GEMM-related configuration for FP8 and FP4 backends.
|
||||
initialize_fp8_gemm_config(self.server_args)
|
||||
initialize_fp4_gemm_config(self.server_args)
|
||||
initialize_fp8_gemm_config()
|
||||
initialize_fp4_gemm_config()
|
||||
initialize_bf16_gemm_config(self.server_args)
|
||||
|
||||
# This must be called after initialize_moe_config
|
||||
|
||||
@@ -179,8 +179,8 @@ class TokenizerControlMixin:
|
||||
)
|
||||
if primary_group_control:
|
||||
control_fan_out = (
|
||||
worker_count + self.server_args.tp_size - 1
|
||||
) // self.server_args.tp_size
|
||||
worker_count + get_parallel().config.tp_size - 1
|
||||
) // get_parallel().config.tp_size
|
||||
else:
|
||||
control_fan_out = worker_count
|
||||
|
||||
|
||||
@@ -3591,7 +3591,7 @@ def get_processor_wrapper(server_args):
|
||||
tokenizer_mode=get_serving().tokenizer_mode,
|
||||
trust_remote_code=get_model().trust_remote_code,
|
||||
revision=get_model().revision,
|
||||
image_processor_backend=resolve_image_processor_backend(server_args),
|
||||
image_processor_backend=resolve_image_processor_backend(get_mm()),
|
||||
tokenizer_backend=get_serving().tokenizer_backend,
|
||||
model_name=get_model().model_path,
|
||||
)
|
||||
|
||||
@@ -86,7 +86,7 @@ class HiRadixCache(RadixCache):
|
||||
self.page_size = params.page_size
|
||||
self.kv_cache = params.token_to_kv_pool_allocator.get_kvcache()
|
||||
|
||||
allocator_type = get_allocator_type(server_args)
|
||||
allocator_type = get_allocator_type()
|
||||
|
||||
if isinstance(self.kv_cache, MHATokenToKVPool):
|
||||
self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)(
|
||||
|
||||
@@ -44,8 +44,8 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _get_allocator_type(server_args: ServerArgs) -> str:
|
||||
return get_allocator_type(server_args)
|
||||
def _get_allocator_type() -> str:
|
||||
return get_allocator_type()
|
||||
|
||||
|
||||
def _evict_swa_for_device_alloc(cache: UnifiedRadixCache, required_size: int) -> None:
|
||||
@@ -126,7 +126,7 @@ def build_kv_host_pool(
|
||||
get_memory().hicache_size if host_size is None else host_size,
|
||||
page_size,
|
||||
get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
pool_label=pool_label,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -540,7 +540,7 @@ def build_deepseek_v4_hicache_stack(
|
||||
num_host_pages=swa_num_host_pages,
|
||||
slot_page_size=kvcache.swa_page_size,
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
)
|
||||
swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator
|
||||
entries.append(
|
||||
@@ -567,7 +567,7 @@ def build_deepseek_v4_hicache_stack(
|
||||
num_host_pages=num_host_pages,
|
||||
slot_page_size=page_size,
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
)
|
||||
c4_indexer_host_pool = DeepSeekV4PagedHostPool(
|
||||
pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER),
|
||||
@@ -579,7 +579,7 @@ def build_deepseek_v4_hicache_stack(
|
||||
num_host_pages=num_host_pages,
|
||||
slot_page_size=page_size,
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
)
|
||||
entries.extend(
|
||||
[
|
||||
@@ -610,7 +610,7 @@ def build_deepseek_v4_hicache_stack(
|
||||
num_host_pages=swa_num_host_pages,
|
||||
swa_page_size=kvcache.swa_page_size,
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
)
|
||||
c4_indexer_state_host_pool = DeepSeekV4StateHostPool(
|
||||
pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER_STATE),
|
||||
@@ -621,7 +621,7 @@ def build_deepseek_v4_hicache_stack(
|
||||
num_host_pages=swa_num_host_pages,
|
||||
swa_page_size=kvcache.swa_page_size,
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
)
|
||||
entries.extend(
|
||||
[
|
||||
@@ -653,7 +653,7 @@ def build_deepseek_v4_hicache_stack(
|
||||
num_host_pages=num_host_pages,
|
||||
slot_page_size=page_size,
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
)
|
||||
# C128 state pool is intentionally not registered with hicache.
|
||||
# page_size=256 % 128 == 0, so state pool is not consumed on load.
|
||||
@@ -739,7 +739,7 @@ def build_hybrid_mamba_stack(
|
||||
mamba_pool,
|
||||
get_memory().hicache_ratio,
|
||||
mamba_host_size,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
)
|
||||
entries = [
|
||||
@@ -1038,7 +1038,7 @@ def build_full_draft_pools(
|
||||
host_to_device_ratio=host_pool_group.logical_size / pool.size,
|
||||
page_size=controller.page_size,
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
pool_label="draft",
|
||||
)
|
||||
draft_layer_mapping = {i: i for i in range(pool.layer_num)}
|
||||
@@ -1064,7 +1064,7 @@ def build_full_draft_pools(
|
||||
pool,
|
||||
draft_host_pool,
|
||||
get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
)
|
||||
specs.append(
|
||||
SidecarPoolSpec(
|
||||
@@ -1111,7 +1111,7 @@ def build_swa_draft_pools(
|
||||
num_host_pages=target_swa_host_pool.num_host_pages,
|
||||
slot_page_size=draft_swa_pool.page_size,
|
||||
layout=target_swa_host_pool.layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
)
|
||||
else:
|
||||
host_pool = _build_mha_mla_host_pool(
|
||||
@@ -1119,7 +1119,7 @@ def build_swa_draft_pools(
|
||||
host_to_device_ratio=target_swa_host_pool.size / draft_swa_pool.size,
|
||||
page_size=target_swa_host_pool.page_size,
|
||||
layout=target_swa_host_pool.layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
pool_label="draft_swa",
|
||||
)
|
||||
|
||||
@@ -1513,7 +1513,7 @@ class _DsaStrategy(StackStrategy):
|
||||
full_kv_pool,
|
||||
kv_host_pool,
|
||||
get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
),
|
||||
prefetch_threshold=prefetch_threshold,
|
||||
model_name=model_name,
|
||||
@@ -1951,7 +1951,7 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
|
||||
kv,
|
||||
kv_host_pool,
|
||||
get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
allocator_type=_get_allocator_type(),
|
||||
),
|
||||
model_name=get_serving().served_model_name,
|
||||
storage_backend_extra_config=extra_config,
|
||||
|
||||
@@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Any, Optional
|
||||
import msgspec
|
||||
import torch
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.configs.hybrid_arch import (
|
||||
hybrid_gdn_config,
|
||||
kimi_linear_config,
|
||||
@@ -1262,7 +1263,7 @@ class KVCacheConfigurator:
|
||||
sparse_layer_ids=sparse_layer_ids,
|
||||
disable_value_sparse_layer_ids=disable_value_sparse_layer_ids,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
)
|
||||
@@ -1533,7 +1534,7 @@ class KVCacheConfigurator:
|
||||
# with the widening-dequant contract.
|
||||
index_dtype=(
|
||||
self.kv_cache_dtype
|
||||
if m3_fp8_attn_gemm_enabled(self.server_args)
|
||||
if m3_fp8_attn_gemm_enabled(resolving_view(self.server_args))
|
||||
else self.model_dtype
|
||||
),
|
||||
head_num=self.model_config.get_num_kv_heads(
|
||||
|
||||
@@ -99,14 +99,15 @@ def get_allocator_from_storage(allocator_type):
|
||||
return HostTensorAllocator()
|
||||
|
||||
|
||||
def get_allocator_type(server_args) -> str:
|
||||
backend = getattr(server_args, "hicache_storage_backend", None)
|
||||
def get_allocator_type() -> str:
|
||||
"""The host-allocator kind the published HiCache configuration asks for."""
|
||||
from sglang.srt.runtime_context import get_memory
|
||||
|
||||
backend = get_memory().hicache_storage_backend
|
||||
if backend == "shm":
|
||||
return "shm"
|
||||
if backend == "dynamic":
|
||||
extra_config_str = getattr(
|
||||
server_args, "hicache_storage_backend_extra_config", None
|
||||
)
|
||||
extra_config_str = get_memory().hicache_storage_backend_extra_config
|
||||
if extra_config_str:
|
||||
try:
|
||||
config = json.loads(extra_config_str)
|
||||
|
||||
@@ -81,6 +81,7 @@ from sglang.srt.observability.metrics_collector import (
|
||||
StorageMetrics,
|
||||
StorageMetricsCollector,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_memory
|
||||
from sglang.srt.session.streaming_session import StreamingSession
|
||||
from sglang.srt.utils.common import ceil_align
|
||||
|
||||
@@ -385,7 +386,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
self.extra_metric_labels = server_args.extra_metric_labels
|
||||
|
||||
# Parse storage config once, share with assembler and tree
|
||||
storage_backend = server_args.hicache_storage_backend
|
||||
storage_backend = get_memory().hicache_storage_backend
|
||||
storage_extra_config = None
|
||||
storage_prefetch_threshold = 256
|
||||
prefetch_timeout_base = 1.0
|
||||
@@ -399,7 +400,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
prefetch_timeout_per_ki_token,
|
||||
hicache_storage_pass_prefix_keys,
|
||||
) = HybridCacheController.parse_storage_backend_extra_config(
|
||||
server_args.hicache_storage_backend_extra_config
|
||||
get_memory().hicache_storage_backend_extra_config
|
||||
)
|
||||
|
||||
attach_hybrid_pool_to_unified_cache(
|
||||
@@ -442,7 +443,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
|
||||
# State initialization
|
||||
self.write_through_threshold = (
|
||||
1 if server_args.hicache_write_policy == "write_through" else 2
|
||||
1 if get_memory().hicache_write_policy == "write_through" else 2
|
||||
)
|
||||
self.is_write_back = (
|
||||
self.cache_controller is not None
|
||||
@@ -457,7 +458,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
pool=_COMPONENT_POOL_LABEL[ct],
|
||||
)
|
||||
self.load_back_threshold = 10
|
||||
self.prefetch_stop_policy = server_args.hicache_storage_prefetch_policy
|
||||
self.prefetch_stop_policy = get_memory().hicache_storage_prefetch_policy
|
||||
|
||||
# Runtime attach/detach of the L3 backend (startup, admin API, atexit).
|
||||
self._storage_attachment = StorageAttachment(self)
|
||||
|
||||
@@ -609,7 +609,7 @@ class CPUGraphRunner:
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
)
|
||||
self.tp_size = model_runner.server_args.tp_size
|
||||
self.tp_size = get_parallel().config.tp_size
|
||||
self.dp_size = get_parallel().config.dp_size
|
||||
self.pp_size = get_parallel().config.pp_size
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ from mindspore._c_expression import GroupOptions
|
||||
from mindspore.communication import create_group
|
||||
|
||||
from sglang.srt.distributed.parallel_state import _groups
|
||||
from sglang.srt.runtime_context import get_serving
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -109,7 +110,7 @@ def init_ms_distributed(world_size, rank, local_rank, server_args, port):
|
||||
if server_args.dist_init_addr:
|
||||
dist_init_method = f"tcp://{server_args.dist_init_addr}"
|
||||
else:
|
||||
dist_init_method = f"tcp://{server_args.host}:{port}"
|
||||
dist_init_method = f"tcp://{get_serving().host}:{port}"
|
||||
set_ms_parallel_env(rank, local_rank, world_size, dist_init_method)
|
||||
|
||||
ms.set_context(infer_boost="on", jit_level="O0")
|
||||
|
||||
@@ -1202,7 +1202,6 @@ class ModelRunner:
|
||||
# Pre-expand RoPE cache before CUDA Graph capture
|
||||
reserve_rope_cache_for_long_sequences(
|
||||
self.model,
|
||||
self.server_args,
|
||||
self.model_config,
|
||||
logger,
|
||||
)
|
||||
|
||||
@@ -13,6 +13,7 @@ from sglang.srt.eplb.lplb_solver import (
|
||||
)
|
||||
from sglang.srt.layers.moe.hash_topk import HashTopK
|
||||
from sglang.srt.layers.moe.topk import TopK
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import get_bool_env_var, is_hip, log_info_on_rank0
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -54,7 +55,7 @@ def prepare_moe_topk(
|
||||
# Redundant experts therefore need to be included in the per-rank
|
||||
# expert count used for Waterfill's shared-expert slot remapping.
|
||||
num_physical_routed_experts = (
|
||||
num_routed_experts + server_args.ep_num_redundant_experts
|
||||
num_routed_experts + get_exec().moe.ep_num_redundant_experts
|
||||
)
|
||||
if isinstance(module, TopK):
|
||||
routed_scaling_factor = module.topk_config.routed_scaling_factor
|
||||
|
||||
@@ -216,7 +216,7 @@ class BaseRunner(ABC):
|
||||
self.model_runner = model_runner
|
||||
self.device = model_runner.device
|
||||
self.device_module = torch.get_device_module(self.device)
|
||||
self.tp_size = model_runner.server_args.tp_size
|
||||
self.tp_size = get_parallel().config.tp_size
|
||||
# elastic-EP scale-up rewrites dp_size on the published config
|
||||
self.dp_size = get_parallel().config.dp_size
|
||||
self.pp_size = get_parallel().config.pp_size
|
||||
|
||||
@@ -64,7 +64,7 @@ def resolve_decode_backend(
|
||||
cfg = get_exec().graph.cuda_graph_config
|
||||
backend_name = cfg.decode.backend if cfg is not None else Backend.FULL
|
||||
|
||||
enable_memory_saver = model_runner.server_args.enable_memory_saver
|
||||
enable_memory_saver = get_exec().features.enable_memory_saver
|
||||
|
||||
if model_runner.device == "npu":
|
||||
from sglang.srt.hardware_backend.npu.graph_runner.npu_cudagraph_backend import (
|
||||
@@ -115,13 +115,13 @@ def resolve_prefill_backend(
|
||||
if backend_name == Backend.BREAKABLE:
|
||||
return BreakableCudaGraphBackend(
|
||||
cuda_graph_runner,
|
||||
enable_memory_saver=model_runner.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
debug_eager=get_exec().graph.debug_cuda_graph,
|
||||
)
|
||||
if backend_name == Backend.FULL:
|
||||
return FullCudaGraphBackend(
|
||||
cuda_graph_runner,
|
||||
enable_memory_saver=model_runner.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
)
|
||||
# Default: tc_piecewise.
|
||||
return TcPiecewiseCudaGraphBackend(cuda_graph_runner)
|
||||
|
||||
@@ -88,7 +88,7 @@ class RayDataParallelController(DataParallelController):
|
||||
dp_port_args_list.append(tmp_port_args)
|
||||
|
||||
# Create ZMQ PUSH socket for this DP rank (controller → scheduler)
|
||||
if server_args.node_rank == 0:
|
||||
if get_parallel().config.node_rank == 0:
|
||||
self.workers[dp_rank] = get_zmq_socket(
|
||||
self.context,
|
||||
zmq.PUSH,
|
||||
@@ -139,7 +139,7 @@ class RayDataParallelController(DataParallelController):
|
||||
dp_rank: DP rank for regular DP; None for DP attention (derived from tp_rank).
|
||||
worker_ports: Pre-allocated ports for DP attention; None for regular DP.
|
||||
"""
|
||||
nnodes = server_args.nnodes
|
||||
nnodes = get_parallel().config.nnodes
|
||||
batch_start_idx = len(self.scheduler_actors)
|
||||
|
||||
if not self.is_custom_pg:
|
||||
@@ -148,7 +148,7 @@ class RayDataParallelController(DataParallelController):
|
||||
pp_range, tp_range, pp_per_node, tp_per_node = _calculate_rank_ranges(
|
||||
nnodes,
|
||||
get_parallel().config.pp_size,
|
||||
server_args.tp_size,
|
||||
get_parallel().config.tp_size,
|
||||
node_rank=node_idx,
|
||||
)
|
||||
for pp_rank in pp_range:
|
||||
@@ -160,13 +160,14 @@ class RayDataParallelController(DataParallelController):
|
||||
tp_rank % tp_per_node
|
||||
)
|
||||
|
||||
if get_parallel().config.enable_dp_attention:
|
||||
parallel = get_parallel().config
|
||||
if parallel.enable_dp_attention:
|
||||
_, _, actual_dp_rank, _ = compute_dp_attention_world_info(
|
||||
get_parallel().config.enable_dp_attention,
|
||||
parallel.enable_dp_attention,
|
||||
tp_rank,
|
||||
server_args.tp_size,
|
||||
get_parallel().config.dp_size,
|
||||
get_parallel().config.attn_cp_size,
|
||||
parallel.tp_size,
|
||||
parallel.dp_size,
|
||||
parallel.attn_cp_size,
|
||||
)
|
||||
rank_port_args = PortArgs.init_new(
|
||||
server_args, actual_dp_rank, worker_ports
|
||||
@@ -204,10 +205,11 @@ class RayDataParallelController(DataParallelController):
|
||||
self.scheduler_actors.append(actor)
|
||||
|
||||
else:
|
||||
world_size = _compute_world_size(server_args)
|
||||
world_size = _compute_world_size()
|
||||
bundle_indices = _resolve_bundle_indices(self.pg, world_size)
|
||||
|
||||
ranks_per_tp_group = server_args.tp_size * get_parallel().config.pp_size
|
||||
parallel = get_parallel().config
|
||||
ranks_per_tp_group = parallel.tp_size * parallel.pp_size
|
||||
if dp_rank is not None:
|
||||
start_rank = dp_rank * ranks_per_tp_group
|
||||
end_rank = start_rank + ranks_per_tp_group
|
||||
@@ -224,8 +226,8 @@ class RayDataParallelController(DataParallelController):
|
||||
|
||||
for global_rank in range(start_rank, end_rank):
|
||||
local_rank = global_rank % ranks_per_tp_group
|
||||
pp_rank = local_rank // server_args.tp_size
|
||||
tp_rank = local_rank % server_args.tp_size
|
||||
pp_rank = local_rank // parallel.tp_size
|
||||
tp_rank = local_rank % parallel.tp_size
|
||||
rank_port_args = port_args
|
||||
actual_dp_rank = dp_rank
|
||||
|
||||
@@ -235,7 +237,7 @@ class RayDataParallelController(DataParallelController):
|
||||
_, _, actual_dp_rank, _ = compute_dp_attention_world_info(
|
||||
get_parallel().config.enable_dp_attention,
|
||||
tp_rank,
|
||||
server_args.tp_size,
|
||||
get_parallel().config.tp_size,
|
||||
get_parallel().config.dp_size,
|
||||
get_parallel().config.attn_cp_size,
|
||||
)
|
||||
|
||||
@@ -105,18 +105,17 @@ def _get_bundle_node_ip(placement_group: PlacementGroup, bundle_idx: int) -> str
|
||||
)
|
||||
|
||||
|
||||
def _compute_world_size(server_args: ServerArgs) -> int:
|
||||
def _compute_world_size() -> int:
|
||||
"""Compute world_size (total number of scheduler actors/GPUs needed).
|
||||
|
||||
Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size.
|
||||
Reads the published parallel leaves: the driver is sizing the actors that
|
||||
will hold the process groups, so there is nothing live to ask.
|
||||
"""
|
||||
if get_parallel().config.enable_dp_attention:
|
||||
return server_args.tp_size * get_parallel().config.pp_size
|
||||
return (
|
||||
get_parallel().config.dp_size
|
||||
* server_args.tp_size
|
||||
* get_parallel().config.pp_size
|
||||
)
|
||||
parallel = get_parallel().config
|
||||
if parallel.enable_dp_attention:
|
||||
return parallel.tp_size * parallel.pp_size
|
||||
return parallel.dp_size * parallel.tp_size * parallel.pp_size
|
||||
|
||||
|
||||
def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]:
|
||||
@@ -274,16 +273,13 @@ class RayEngine(Engine):
|
||||
placement_group as create_placement_group,
|
||||
)
|
||||
|
||||
if get_parallel().config.enable_dp_attention:
|
||||
total_gpus = server_args.tp_size * get_parallel().config.pp_size
|
||||
parallel = get_parallel().config
|
||||
if parallel.enable_dp_attention:
|
||||
total_gpus = parallel.tp_size * parallel.pp_size
|
||||
else:
|
||||
total_gpus = (
|
||||
get_parallel().config.dp_size
|
||||
* server_args.tp_size
|
||||
* get_parallel().config.pp_size
|
||||
)
|
||||
total_gpus = parallel.dp_size * parallel.tp_size * parallel.pp_size
|
||||
|
||||
nnodes = server_args.nnodes
|
||||
nnodes = parallel.nnodes
|
||||
gpus_per_node = total_gpus // nnodes
|
||||
strategy = "STRICT_PACK" if nnodes == 1 else "SPREAD"
|
||||
|
||||
@@ -300,8 +296,8 @@ class RayEngine(Engine):
|
||||
ray.get(pg.ready())
|
||||
|
||||
is_custom_pg = placement_group is not None
|
||||
nnodes = server_args.nnodes
|
||||
world_size = _compute_world_size(server_args)
|
||||
nnodes = get_parallel().config.nnodes
|
||||
world_size = _compute_world_size()
|
||||
|
||||
if not is_custom_pg:
|
||||
engine_bundle, engine_ip = _find_engine_bundle(pg, nnodes)
|
||||
@@ -341,7 +337,7 @@ class RayEngine(Engine):
|
||||
_calculate_rank_ranges(
|
||||
nnodes,
|
||||
get_parallel().config.pp_size,
|
||||
server_args.tp_size,
|
||||
get_parallel().config.tp_size,
|
||||
node_rank=node_idx,
|
||||
)
|
||||
)
|
||||
@@ -377,9 +373,10 @@ class RayEngine(Engine):
|
||||
f"bundle_indices={bundle_indices}"
|
||||
)
|
||||
|
||||
tp_size = get_parallel().config.tp_size
|
||||
for rank in range(world_size):
|
||||
pp_rank = rank // server_args.tp_size
|
||||
tp_rank = rank % server_args.tp_size
|
||||
pp_rank = rank // tp_size
|
||||
tp_rank = rank % tp_size
|
||||
bundle_idx = bundle_indices[rank]
|
||||
|
||||
actor = _create_scheduler_actor(
|
||||
@@ -455,21 +452,18 @@ class RayEngine(Engine):
|
||||
RayDataParallelController,
|
||||
)
|
||||
|
||||
if get_parallel().config.enable_dp_attention:
|
||||
parallel = get_parallel().config
|
||||
if parallel.enable_dp_attention:
|
||||
# DP attention folds DP into TP — total GPUs = tp_size * pp_size
|
||||
total_gpus = server_args.tp_size * get_parallel().config.pp_size
|
||||
total_gpus = parallel.tp_size * parallel.pp_size
|
||||
else:
|
||||
total_gpus = (
|
||||
get_parallel().config.dp_size
|
||||
* server_args.tp_size
|
||||
* get_parallel().config.pp_size
|
||||
)
|
||||
gpus_per_node = total_gpus // server_args.nnodes
|
||||
total_gpus = parallel.dp_size * parallel.tp_size * parallel.pp_size
|
||||
gpus_per_node = total_gpus // parallel.nnodes
|
||||
logger.info(
|
||||
f"Ray DP cluster: {server_args.nnodes} nodes, "
|
||||
f"{gpus_per_node} GPUs/node, dp_size={get_parallel().config.dp_size}, "
|
||||
f"tp_size={server_args.tp_size}, pp_size={get_parallel().config.pp_size}, "
|
||||
f"enable_dp_attention={get_parallel().config.enable_dp_attention}"
|
||||
f"Ray DP cluster: {parallel.nnodes} nodes, "
|
||||
f"{gpus_per_node} GPUs/node, dp_size={parallel.dp_size}, "
|
||||
f"tp_size={parallel.tp_size}, pp_size={parallel.pp_size}, "
|
||||
f"enable_dp_attention={parallel.enable_dp_attention}"
|
||||
)
|
||||
|
||||
# Set dist_init_addr on server_args so PortArgs.init_new() can compute
|
||||
|
||||
@@ -7100,7 +7100,11 @@ class ServerArgs:
|
||||
|
||||
from sglang.srt.layers.cp.base import init_cp_strategy
|
||||
|
||||
init_cp_strategy(self)
|
||||
init_cp_strategy(
|
||||
enable_prefill_cp=bool(cfg.enable_prefill_cp),
|
||||
cp_size=cfg.attn_cp_size,
|
||||
cp_strategy=cfg.cp_strategy,
|
||||
)
|
||||
|
||||
def _handle_dwdp(self):
|
||||
cfg = resolving_view(self)
|
||||
|
||||
@@ -9,6 +9,7 @@ import torch
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
|
||||
from sglang.srt.runtime_context import attention_backends, get_spec
|
||||
from sglang.srt.server_args import DRAFT_ATTENTION_BACKEND_CHOICES, ServerArgs
|
||||
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
|
||||
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
||||
@@ -28,12 +29,16 @@ class DraftWorkerBundle(msgspec.Struct, frozen=True):
|
||||
resolved_attention_backend: str
|
||||
|
||||
|
||||
def _resolve_draft_attention_backend_fallback(
|
||||
*, server_args: ServerArgs, algo_label: str
|
||||
) -> str:
|
||||
draft_backend = server_args.speculative_draft_attention_backend
|
||||
def _resolve_draft_attention_backend_fallback(*, algo_label: str) -> str:
|
||||
"""The draft's attention backend, from the published leaves.
|
||||
|
||||
`spec.speculative_draft_attention_backend` when the operator named one,
|
||||
otherwise the process's prefill backend. Both are resolution's answers, so
|
||||
they come from the bags.
|
||||
"""
|
||||
draft_backend = get_spec().speculative_draft_attention_backend
|
||||
if draft_backend is None:
|
||||
draft_backend, _ = server_args.get_attention_backends()
|
||||
draft_backend, _ = attention_backends()
|
||||
if draft_backend is None:
|
||||
return "triton" if torch.version.hip else "flashinfer"
|
||||
if draft_backend not in DRAFT_ATTENTION_BACKEND_CHOICES:
|
||||
@@ -65,9 +70,7 @@ def build_draft_tp_worker(
|
||||
# validated (e.g. a self-drafting architecture); it skips the generic
|
||||
# supported-backend fallback below.
|
||||
draft_backend = attention_backend_override or (
|
||||
_resolve_draft_attention_backend_fallback(
|
||||
server_args=server_args, algo_label=algo_label
|
||||
)
|
||||
_resolve_draft_attention_backend_fallback(algo_label=algo_label)
|
||||
)
|
||||
from sglang.srt.layers.moe.utils import draft_model_build_scope
|
||||
|
||||
|
||||
@@ -4623,11 +4623,15 @@ def cached_triton_kernel(key_fn=None):
|
||||
return decorator
|
||||
|
||||
|
||||
def reserve_rope_cache_for_long_sequences(
|
||||
model, server_args, model_config, logger=None
|
||||
):
|
||||
"""Pre-expand RoPE cache for long sequences and speculative decoding."""
|
||||
def reserve_rope_cache_for_long_sequences(model, model_config, logger=None):
|
||||
"""Pre-expand RoPE cache for long sequences and speculative decoding.
|
||||
|
||||
Runs inside `ModelRunner`, past publish, so the three config inputs come
|
||||
from the bags: the context length and the two speculative counts are
|
||||
resolution's answers.
|
||||
"""
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_model, get_spec
|
||||
|
||||
SAFETY_FACTOR = envs.SGLANG_SPEC_EXPANSION_SAFETY_FACTOR.get()
|
||||
MARGIN = envs.SGLANG_ROPE_CACHE_SAFETY_MARGIN.get()
|
||||
@@ -4635,7 +4639,7 @@ def reserve_rope_cache_for_long_sequences(
|
||||
|
||||
# 1) Estimate base context upper bound
|
||||
base_ctx = (
|
||||
getattr(server_args, "context_length", None)
|
||||
get_model().context_length
|
||||
or getattr(model_config, "context_len", None)
|
||||
or getattr(model_config, "max_model_len", None)
|
||||
or getattr(model_config.hf_text_config, "max_position_embeddings", None)
|
||||
@@ -4643,8 +4647,8 @@ def reserve_rope_cache_for_long_sequences(
|
||||
)
|
||||
|
||||
# 2) Speculative decoding expansion
|
||||
steps = int(getattr(server_args, "speculative_num_steps", 0) or 0)
|
||||
draft = int(getattr(server_args, "speculative_num_draft_tokens", 0) or 0)
|
||||
steps = int(get_spec().speculative_num_steps or 0)
|
||||
draft = int(get_spec().speculative_num_draft_tokens or 0)
|
||||
reserve = base_ctx + steps * draft * SAFETY_FACTOR + MARGIN
|
||||
|
||||
# 3) Align to reduce reallocation frequency
|
||||
|
||||
@@ -54,11 +54,16 @@ from .tokenizer import (
|
||||
_IMAGE_PROCESSOR_BACKENDS = {"auto", "torchvision", "pil"}
|
||||
|
||||
|
||||
def resolve_image_processor_backend(server_args) -> str:
|
||||
"""Resolve the new backend option while honoring the legacy disable flag."""
|
||||
if getattr(server_args, "disable_fast_image_processor", False):
|
||||
def resolve_image_processor_backend(mm_config) -> str:
|
||||
"""Resolve the new backend option while honoring the legacy disable flag.
|
||||
|
||||
Takes the `mm` config bag (`get_mm()`): both leaves are resolved config, and
|
||||
every caller is past publish. `getattr` with a default keeps it working for a
|
||||
stand-in that carries only one of the two.
|
||||
"""
|
||||
if getattr(mm_config, "disable_fast_image_processor", False):
|
||||
return "pil"
|
||||
return getattr(server_args, "image_processor_backend", "auto")
|
||||
return getattr(mm_config, "image_processor_backend", "auto")
|
||||
|
||||
|
||||
def _normalize_image_processor_backend(
|
||||
|
||||
@@ -47,6 +47,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.configs.load_config import LoadConfig
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_parallel, publish
|
||||
@@ -148,27 +149,28 @@ class WeightCacheDaemon:
|
||||
dist_init_method: Optional[str] = None,
|
||||
):
|
||||
self.server_args = server_args
|
||||
self.model_path = server_args.model_path
|
||||
cfg = resolving_view(server_args)
|
||||
self.model_path = cfg.model_path
|
||||
self.gpu_id = gpu_id
|
||||
self.tp_size = server_args.tp_size
|
||||
self.tp_size = cfg.tp_size
|
||||
self.tp_rank = tp_rank
|
||||
self.pp_size = server_args.pp_size
|
||||
self.pp_size = cfg.pp_size
|
||||
self.pp_rank = pp_rank
|
||||
self.dp_size = server_args.dp_size
|
||||
self.ep_size = server_args.ep_size
|
||||
self.moe_dp_size = server_args.moe_dp_size
|
||||
self.enable_dp_attention = server_args.enable_dp_attention
|
||||
self.enable_dp_lm_head = server_args.enable_dp_lm_head
|
||||
self.attn_cp_size = server_args.attn_cp_size
|
||||
self.moe_dense_tp_size = server_args.moe_dense_tp_size
|
||||
self.moe_a2a_backend = server_args.moe_a2a_backend
|
||||
self.deepep_mode = server_args.deepep_mode
|
||||
self.load_format = server_args.load_format
|
||||
self.dtype = server_args.dtype
|
||||
self.quantization = server_args.quantization
|
||||
self.model_loader_extra_config = server_args.model_loader_extra_config
|
||||
self.trust_remote_code = server_args.trust_remote_code
|
||||
self.revision = server_args.revision
|
||||
self.dp_size = cfg.dp_size
|
||||
self.ep_size = cfg.ep_size
|
||||
self.moe_dp_size = cfg.moe_dp_size
|
||||
self.enable_dp_attention = cfg.enable_dp_attention
|
||||
self.enable_dp_lm_head = cfg.enable_dp_lm_head
|
||||
self.attn_cp_size = cfg.attn_cp_size
|
||||
self.moe_dense_tp_size = cfg.moe_dense_tp_size
|
||||
self.moe_a2a_backend = cfg.moe_a2a_backend
|
||||
self.deepep_mode = cfg.deepep_mode
|
||||
self.load_format = cfg.load_format
|
||||
self.dtype = cfg.dtype
|
||||
self.quantization = cfg.quantization
|
||||
self.model_loader_extra_config = cfg.model_loader_extra_config
|
||||
self.trust_remote_code = cfg.trust_remote_code
|
||||
self.revision = cfg.revision
|
||||
self.dist_init_method = dist_init_method
|
||||
|
||||
self.socket_path = get_socket_path(
|
||||
@@ -223,7 +225,7 @@ class WeightCacheDaemon:
|
||||
distributed_init_method=self.dist_init_method,
|
||||
local_rank=self.gpu_id,
|
||||
backend=current_platform.get_torch_distributed_backend_str(),
|
||||
moe_a2a_backend=server_args.moe_a2a_backend,
|
||||
moe_a2a_backend=self.moe_a2a_backend,
|
||||
)
|
||||
|
||||
initialize_model_parallel(
|
||||
@@ -281,7 +283,7 @@ class WeightCacheDaemon:
|
||||
|
||||
from sglang.srt.layers.moe import initialize_moe_config
|
||||
|
||||
initialize_moe_config(server_args)
|
||||
initialize_moe_config()
|
||||
|
||||
# Initialize distributed backend for model loading
|
||||
# (must be done after server_args and model_config are available)
|
||||
@@ -672,23 +674,24 @@ def launch_weight_cache_daemons(
|
||||
--nnodes 2 --node-rank 1 \\
|
||||
--dist-init-method tcp://node0-ip:29500
|
||||
"""
|
||||
cfg = resolving_view(server_args)
|
||||
import socket as sock_mod
|
||||
|
||||
# Replicate _calculate_rank_ranges logic from engine.py
|
||||
pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1)
|
||||
nnodes_per_pp_rank = max(server_args.nnodes // server_args.pp_size, 1)
|
||||
pp_size_per_node = max(cfg.pp_size // cfg.nnodes, 1)
|
||||
nnodes_per_pp_rank = max(cfg.nnodes // cfg.pp_size, 1)
|
||||
pp_rank_range = range(
|
||||
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank),
|
||||
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank + 1),
|
||||
pp_size_per_node * (cfg.node_rank // nnodes_per_pp_rank),
|
||||
pp_size_per_node * (cfg.node_rank // nnodes_per_pp_rank + 1),
|
||||
)
|
||||
nnodes_per_tp_group = nnodes_per_pp_rank
|
||||
tp_size_per_node = server_args.tp_size // nnodes_per_tp_group
|
||||
tp_size_per_node = cfg.tp_size // nnodes_per_tp_group
|
||||
tp_rank_range = range(
|
||||
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group),
|
||||
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1),
|
||||
tp_size_per_node * (cfg.node_rank % nnodes_per_tp_group),
|
||||
tp_size_per_node * (cfg.node_rank % nnodes_per_tp_group + 1),
|
||||
)
|
||||
|
||||
if server_args.nnodes > 1 and dist_init_method is None:
|
||||
if cfg.nnodes > 1 and dist_init_method is None:
|
||||
raise ValueError(
|
||||
"dist_init_method is required for multi-node weight cache daemons. "
|
||||
"Use --dist-init-method tcp://<node0-ip>:<port> to specify the "
|
||||
@@ -705,7 +708,7 @@ def launch_weight_cache_daemons(
|
||||
# Validate and clean up stale .ready/.sock files from prior runs.
|
||||
for pp_rank in pp_rank_range:
|
||||
for tp_rank in tp_rank_range:
|
||||
global_rank = compute_global_rank(server_args.tp_size, pp_rank, tp_rank)
|
||||
global_rank = compute_global_rank(cfg.tp_size, pp_rank, tp_rank)
|
||||
cleanup_stale_daemon_files(global_rank, force=force)
|
||||
|
||||
procs = []
|
||||
@@ -716,8 +719,8 @@ def launch_weight_cache_daemons(
|
||||
tp_rank,
|
||||
pp_size_per_node,
|
||||
tp_size_per_node,
|
||||
base_gpu_id=server_args.base_gpu_id,
|
||||
gpu_id_step=server_args.gpu_id_step,
|
||||
base_gpu_id=cfg.base_gpu_id,
|
||||
gpu_id_step=cfg.gpu_id_step,
|
||||
)
|
||||
proc = spawn_weight_cache_daemon(
|
||||
server_args,
|
||||
@@ -738,7 +741,7 @@ def launch_weight_cache_daemons(
|
||||
start_time = time.time()
|
||||
for pp_rank in pp_rank_range:
|
||||
for tp_rank in tp_rank_range:
|
||||
global_rank = compute_global_rank(server_args.tp_size, pp_rank, tp_rank)
|
||||
global_rank = compute_global_rank(cfg.tp_size, pp_rank, tp_rank)
|
||||
ready_path = get_ready_path(global_rank)
|
||||
while not os.path.exists(ready_path):
|
||||
time.sleep(check_interval)
|
||||
@@ -772,7 +775,7 @@ def launch_weight_cache_daemons(
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"All {num_daemons} weight cache daemons on node {server_args.node_rank} are ready "
|
||||
f"All {num_daemons} weight cache daemons on node {cfg.node_rank} are ready "
|
||||
f"(pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, "
|
||||
f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, "
|
||||
f"dist_init_method={dist_init_method})"
|
||||
|
||||
@@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Generator, List, Optional, Tuple
|
||||
|
||||
import zmq
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
|
||||
from sglang.srt.utils.network import get_zmq_socket
|
||||
@@ -56,8 +57,8 @@ def _drive_engine_through_warmup(ctx: ScriptedContext) -> Generator:
|
||||
"""Run the engine until the server warmup request has been received and
|
||||
fully processed, so scripts never observe foreign warmup traffic."""
|
||||
scheduler = ctx.scheduler
|
||||
server_args = scheduler.server_args
|
||||
if server_args.skip_server_warmup:
|
||||
cfg = resolving_view(scheduler.server_args)
|
||||
if cfg.skip_server_warmup:
|
||||
logger.info("scripted_runtime: skip_server_warmup set, not driving warmup")
|
||||
return
|
||||
|
||||
@@ -67,7 +68,7 @@ def _drive_engine_through_warmup(ctx: ScriptedContext) -> Generator:
|
||||
# is_fully_idle() can transiently report idle while a PP microbatch result
|
||||
# is still in flight, so require it to hold for two full microbatch
|
||||
# rotations after the warmup request was observed on the recv socket.
|
||||
quiesce_iters = 2 * (server_args.pp_size + server_args.pp_async_batch_depth)
|
||||
quiesce_iters = 2 * (cfg.pp_size + cfg.pp_async_batch_depth)
|
||||
proxy = ctx._tokenizer_recv_proxy
|
||||
deadline = start_time + WARMUP_DRIVE_TIMEOUT_S
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from sglang.srt.distributed.parallel_state import (
|
||||
from sglang.srt.layers.dp_attention import set_dp_buffer_len
|
||||
from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatcher
|
||||
from sglang.srt.layers.moe.utils import initialize_moe_config
|
||||
from sglang.srt.runtime_context import publish
|
||||
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -22,7 +23,8 @@ class TestFlashinferDispatcher(CustomTestCase):
|
||||
server_args.moe_runner_backend = "flashinfer_cutlass"
|
||||
server_args.moe_a2a_backend = "flashinfer"
|
||||
set_global_server_args_for_scheduler(server_args)
|
||||
initialize_moe_config(server_args)
|
||||
publish(server_args, role="scheduler")
|
||||
initialize_moe_config()
|
||||
|
||||
init_distributed_environment(
|
||||
world_size=-1, # Auto-detect from environment
|
||||
|
||||
@@ -60,7 +60,7 @@ class _FakeCPGroup:
|
||||
|
||||
class TestCPStrategyUnit(CustomTestCase):
|
||||
def tearDown(self):
|
||||
init_cp_strategy(SimpleNamespace(enable_prefill_cp=False))
|
||||
init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="zigzag")
|
||||
|
||||
def test_strategy_kind_maps_cli_values(self):
|
||||
self.assertEqual(ContextParallelStrategyKind.NONE.value, 0)
|
||||
@@ -77,11 +77,9 @@ class TestCPStrategyUnit(CustomTestCase):
|
||||
|
||||
def test_init_cp_strategy_binds_zigzag_strategy(self):
|
||||
init_cp_strategy(
|
||||
SimpleNamespace(
|
||||
enable_prefill_cp=True,
|
||||
cp_strategy="zigzag",
|
||||
attn_cp_size=4,
|
||||
)
|
||||
enable_prefill_cp=True,
|
||||
cp_size=4,
|
||||
cp_strategy="zigzag",
|
||||
)
|
||||
|
||||
self.assertTrue(is_cp_enabled())
|
||||
@@ -91,11 +89,9 @@ class TestCPStrategyUnit(CustomTestCase):
|
||||
|
||||
def test_get_cp_strategy_is_initialized_under_cp_v2(self):
|
||||
init_cp_strategy(
|
||||
SimpleNamespace(
|
||||
enable_prefill_cp=True,
|
||||
cp_strategy="interleave",
|
||||
attn_cp_size=4,
|
||||
)
|
||||
enable_prefill_cp=True,
|
||||
cp_size=4,
|
||||
cp_strategy="interleave",
|
||||
)
|
||||
|
||||
with patch(
|
||||
@@ -108,7 +104,7 @@ class TestCPStrategyUnit(CustomTestCase):
|
||||
|
||||
class TestPrefillCPBCGReplay(CustomTestCase):
|
||||
def tearDown(self):
|
||||
init_cp_strategy(SimpleNamespace(enable_prefill_cp=False))
|
||||
init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="zigzag")
|
||||
|
||||
def _make_runner(self):
|
||||
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
|
||||
@@ -141,11 +137,9 @@ class TestPrefillCPBCGReplay(CustomTestCase):
|
||||
|
||||
def _enable_zigzag(self):
|
||||
init_cp_strategy(
|
||||
SimpleNamespace(
|
||||
enable_prefill_cp=True,
|
||||
cp_strategy="zigzag",
|
||||
attn_cp_size=4,
|
||||
)
|
||||
enable_prefill_cp=True,
|
||||
cp_size=4,
|
||||
cp_strategy="zigzag",
|
||||
)
|
||||
|
||||
def test_local_capacity_overflow_uses_next_capture_bucket(self):
|
||||
@@ -267,16 +261,13 @@ class TestPrefillCPBCGReplay(CustomTestCase):
|
||||
class TestCPZigzagStrategy(CustomTestCase):
|
||||
def setUp(self):
|
||||
init_cp_strategy(
|
||||
SimpleNamespace(
|
||||
enable_prefill_cp=True,
|
||||
cp_strategy="zigzag",
|
||||
attn_cp_size=4,
|
||||
attention_backend="fa3",
|
||||
)
|
||||
enable_prefill_cp=True,
|
||||
cp_size=4,
|
||||
cp_strategy="zigzag",
|
||||
)
|
||||
|
||||
def tearDown(self):
|
||||
init_cp_strategy(SimpleNamespace(enable_prefill_cp=False))
|
||||
init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="zigzag")
|
||||
|
||||
def _metadata_for_rank(self, rank, *, cp_size, seq_lens, extend_seq_lens):
|
||||
strategy = ZigzagCPStrategy(cp_size=cp_size)
|
||||
@@ -809,16 +800,13 @@ class TestCPZigzagStrategy(CustomTestCase):
|
||||
class TestCPInterleaveStrategy(CustomTestCase):
|
||||
def setUp(self):
|
||||
init_cp_strategy(
|
||||
SimpleNamespace(
|
||||
enable_prefill_cp=True,
|
||||
cp_strategy="interleave",
|
||||
attn_cp_size=4,
|
||||
attention_backend="fa3",
|
||||
)
|
||||
enable_prefill_cp=True,
|
||||
cp_size=4,
|
||||
cp_strategy="interleave",
|
||||
)
|
||||
|
||||
def tearDown(self):
|
||||
init_cp_strategy(SimpleNamespace(enable_prefill_cp=False))
|
||||
init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="zigzag")
|
||||
|
||||
def _metadata_for_rank(self, rank, *, cp_size, seq_lens, extend_seq_lens):
|
||||
strategy = InterleaveCPStrategy(cp_size=cp_size)
|
||||
|
||||
@@ -71,7 +71,7 @@ class TestEmbeddingModelSpec(unittest.TestCase):
|
||||
)
|
||||
plan = resolved_embedding_plan(
|
||||
spec,
|
||||
server_args=SimpleNamespace(
|
||||
config=SimpleNamespace(
|
||||
is_embedding=True,
|
||||
cuda_graph_config=SimpleNamespace(
|
||||
prefill=SimpleNamespace(
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import contextlib
|
||||
import types
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
@@ -90,23 +89,24 @@ def _torch_allreduce_residual_rmsnorm_baseline(
|
||||
|
||||
|
||||
class TestFlashInferCommFusion(CustomTestCase):
|
||||
"""The arch dispatch is `_resolve_backend(backend, is_multi_node)`.
|
||||
|
||||
The public entry above it takes no arguments -- it reads
|
||||
`exec.comm.flashinfer_allreduce_fusion_backend` and `parallel.nnodes` off the
|
||||
published bags -- so the cases here drive the dispatch directly.
|
||||
"""
|
||||
|
||||
def test_auto_backend_resolves_by_arch(self):
|
||||
single_node = types.SimpleNamespace(
|
||||
flashinfer_allreduce_fusion_backend="auto", nnodes=1
|
||||
)
|
||||
multi_node = types.SimpleNamespace(
|
||||
flashinfer_allreduce_fusion_backend="auto", nnodes=2
|
||||
)
|
||||
single_node = ("auto", False)
|
||||
multi_node = ("auto", True)
|
||||
|
||||
# Blackwell: mnnvl on both single-node and multi-node.
|
||||
with patch.object(fusion, "is_sm100_supported", return_value=True):
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node),
|
||||
fusion._resolve_backend(*single_node),
|
||||
"mnnvl",
|
||||
)
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node), "mnnvl"
|
||||
)
|
||||
self.assertEqual(fusion._resolve_backend(*multi_node), "mnnvl")
|
||||
|
||||
# SM90: auto uses trtllm on single-node, multi-node is unsupported.
|
||||
with (
|
||||
@@ -114,11 +114,11 @@ class TestFlashInferCommFusion(CustomTestCase):
|
||||
patch.object(fusion, "is_sm90_supported", return_value=True),
|
||||
):
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node),
|
||||
fusion._resolve_backend(*single_node),
|
||||
"trtllm",
|
||||
)
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node)
|
||||
fusion._resolve_backend(*multi_node)
|
||||
|
||||
# Architectures outside SM90/SM10X are unsupported. Both pre-SM90
|
||||
# and post-SM10X devices (e.g. SM120) must fail closed.
|
||||
@@ -129,48 +129,40 @@ class TestFlashInferCommFusion(CustomTestCase):
|
||||
patch.object(fusion, "is_sm90_supported", return_value=False),
|
||||
):
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node)
|
||||
fusion._resolve_backend(*single_node)
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node)
|
||||
fusion._resolve_backend(*multi_node)
|
||||
|
||||
def test_explicit_backend_validation(self):
|
||||
single_node_mnnvl = types.SimpleNamespace(
|
||||
flashinfer_allreduce_fusion_backend="mnnvl", nnodes=1
|
||||
)
|
||||
multi_node_mnnvl = types.SimpleNamespace(
|
||||
flashinfer_allreduce_fusion_backend="mnnvl", nnodes=2
|
||||
)
|
||||
single_node_trtllm = types.SimpleNamespace(
|
||||
flashinfer_allreduce_fusion_backend="trtllm", nnodes=1
|
||||
)
|
||||
multi_node_trtllm = types.SimpleNamespace(
|
||||
flashinfer_allreduce_fusion_backend="trtllm", nnodes=2
|
||||
)
|
||||
single_node_mnnvl = ("mnnvl", False)
|
||||
multi_node_mnnvl = ("mnnvl", True)
|
||||
single_node_trtllm = ("trtllm", False)
|
||||
multi_node_trtllm = ("trtllm", True)
|
||||
|
||||
with (
|
||||
patch.object(fusion, "is_sm100_supported", return_value=False),
|
||||
patch.object(fusion, "is_sm90_supported", return_value=True),
|
||||
):
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node_mnnvl),
|
||||
fusion._resolve_backend(*single_node_mnnvl),
|
||||
"mnnvl",
|
||||
)
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node_trtllm),
|
||||
fusion._resolve_backend(*single_node_trtllm),
|
||||
"trtllm",
|
||||
)
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_mnnvl)
|
||||
fusion._resolve_backend(*multi_node_mnnvl)
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_trtllm)
|
||||
fusion._resolve_backend(*multi_node_trtllm)
|
||||
|
||||
with patch.object(fusion, "is_sm100_supported", return_value=True):
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_mnnvl),
|
||||
fusion._resolve_backend(*multi_node_mnnvl),
|
||||
"mnnvl",
|
||||
)
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_trtllm)
|
||||
fusion._resolve_backend(*multi_node_trtllm)
|
||||
|
||||
for arch in ("pre_sm90", "post_sm10x"):
|
||||
with (
|
||||
@@ -184,9 +176,9 @@ class TestFlashInferCommFusion(CustomTestCase):
|
||||
single_node_trtllm,
|
||||
multi_node_trtllm,
|
||||
):
|
||||
with self.subTest(backend=args.flashinfer_allreduce_fusion_backend):
|
||||
with self.subTest(backend=args[0], multi_node=args[1]):
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(args)
|
||||
fusion._resolve_backend(*args)
|
||||
|
||||
def test_allreduce_fusion_backends_match_torch_baseline(self):
|
||||
fake_comm = _FakeFlashInferComm()
|
||||
|
||||
@@ -27,6 +27,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Phase,
|
||||
PhaseConfig,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_context, get_serving
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs, prepare_server_args
|
||||
from sglang.srt.server_args_config_parser import ConfigArgumentMerger
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
@@ -2356,9 +2357,13 @@ class TestGrpcServerArgs(CustomTestCase):
|
||||
|
||||
fake_core = SimpleNamespace(start_server=MagicMock(return_value="handle"))
|
||||
fake_bridge = SimpleNamespace(RuntimeHandle=MagicMock(return_value="rt"))
|
||||
server_args = SimpleNamespace(
|
||||
host="127.0.0.1", grpc_port=50051, grpc_worker_threads=4
|
||||
)
|
||||
# The host comes from the `serving` bag; `grpc_worker_threads` is not a
|
||||
# field (resolution sets it from the environment), so it stays on the
|
||||
# stand-in the call site is handed.
|
||||
override = get_context().override_server_args(host="127.0.0.1", grpc_port=50051)
|
||||
override.install()
|
||||
self.addCleanup(override.restore)
|
||||
server_args = SimpleNamespace(grpc_worker_threads=4)
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.rust_extensions.load_rust_extension",
|
||||
@@ -2373,7 +2378,7 @@ class TestGrpcServerArgs(CustomTestCase):
|
||||
tokenizer_manager=MagicMock(),
|
||||
template_manager=MagicMock(),
|
||||
scheduler_info={},
|
||||
grpc_port=resolution_result(server_args, "grpc_port"),
|
||||
grpc_port=get_serving().grpc_port,
|
||||
)
|
||||
|
||||
self.assertEqual(handle, "handle")
|
||||
|
||||
@@ -19,7 +19,7 @@ from sglang.srt.layers.moe.utils import (
|
||||
speculative_moe_backend_context,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.runtime_context import get_context, get_flags, get_model
|
||||
from sglang.srt.runtime_context import get_context, get_flags, get_model, publish
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -145,9 +145,11 @@ class TestFusionDecisionFlag(CustomTestCase):
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
self._seed()
|
||||
initialize_moe_config(
|
||||
ServerArgs(model_path="dummy", disable_shared_experts_fusion=True)
|
||||
publish(
|
||||
ServerArgs(model_path="dummy", disable_shared_experts_fusion=True),
|
||||
role="scheduler",
|
||||
)
|
||||
initialize_moe_config()
|
||||
moe = get_flags().moe
|
||||
self.assertTrue(moe.disable_shared_experts_fusion)
|
||||
self.assertTrue(moe.speculative_disable_shared_experts_fusion)
|
||||
|
||||
@@ -53,6 +53,18 @@ _SLOT_OWNERS = ("srt/runtime_context.py", "srt/server_args.py", "srt/arg_groups/
|
||||
# The test below asserts this map is exactly the set of such reads, so the
|
||||
# reasons cannot drift away from the code.
|
||||
_CONFIGURED_SIZE_CALL_SITES = {
|
||||
("srt/layers/cp/base.py", "attn_cp_size"): (
|
||||
"the lazy strategy bind in a worker: the CP group is what the strategy "
|
||||
"is being built for, and the configured width is what describes it"
|
||||
),
|
||||
("benchmark/one_batch.py", "pp_size"): (
|
||||
"CPU affinity for this rank, computed right after the work function "
|
||||
"publishes and before dist init, so the groups do not exist yet"
|
||||
),
|
||||
("benchmark/one_batch.py", "tp_size"): (
|
||||
"the same affinity computation: the layout is the configured one, and "
|
||||
"the live group is not up at this point in the work function"
|
||||
),
|
||||
("srt/entrypoints/engine.py", "pp_size"): (
|
||||
"the launch path decides how many scheduler processes to spawn; it runs "
|
||||
"before any of them exists, so there is no group to ask"
|
||||
@@ -68,6 +80,13 @@ _CONFIGURED_SIZE_CALL_SITES = {
|
||||
"the Ray driver sizes the actor placement group; the actors it is about "
|
||||
"to create are the ones that will hold the process groups"
|
||||
),
|
||||
("srt/ray/engine.py", "tp_size"): (
|
||||
"the same placement arithmetic as the stage count: the driver sizes "
|
||||
"the actors that will hold the process groups"
|
||||
),
|
||||
("srt/ray/data_parallel_controller.py", "tp_size"): (
|
||||
"the same arithmetic on the DP path, also in the driver"
|
||||
),
|
||||
("srt/ray/data_parallel_controller.py", "pp_size"): (
|
||||
"same placement arithmetic on the DP path -- ranks per TP group, "
|
||||
"computed in the driver before the actors start"
|
||||
@@ -115,6 +134,26 @@ _CONFIGURED_SIZE_CALL_SITES = {
|
||||
("srt/managers/scheduler.py", "dcp_size"): (
|
||||
"same pre-distributed-init arithmetic in configure_scheduler_process"
|
||||
),
|
||||
("srt/model_executor/runner/base_runner.py", "tp_size"): (
|
||||
"the same window as the stage count next to it: a draft runner shares "
|
||||
"the target's groups, so the live property would answer for the wrong "
|
||||
"runner"
|
||||
),
|
||||
("srt/model_executor/cpu_graph_runner.py", "tp_size"): (
|
||||
"the same window, on the CPU graph path"
|
||||
),
|
||||
("srt/entrypoints/v1_loads.py", "tp_size"): (
|
||||
"the accelerator count is arithmetic over the launch shape, reported "
|
||||
"from the tokenizer process, which holds no model groups"
|
||||
),
|
||||
("srt/disaggregation/nixl/conn.py", "tp_size"): (
|
||||
"the NIXL rank arithmetic runs on the transfer path, which the CPU-only "
|
||||
"conn tests exercise without starting torch.distributed"
|
||||
),
|
||||
("srt/managers/tokenizer_control_mixin.py", "tp_size"): (
|
||||
"the tokenizer divides its worker count by the launch width; it holds "
|
||||
"no model groups"
|
||||
),
|
||||
("srt/model_executor/runner/base_runner.py", "pp_size"): (
|
||||
"the runner's layer window is arithmetic over the configured stage "
|
||||
"count; a draft runner shares the target's groups, so the live "
|
||||
@@ -212,6 +251,27 @@ _CONFIGURED_SIZE_CALL_SITES = {
|
||||
"the encode server's launch entry sizes its workers before it has "
|
||||
"spawned any of them"
|
||||
),
|
||||
("srt/disaggregation/encoder/grpc_server.py", "tp_size"): (
|
||||
"the same worker-count arithmetic on the gRPC entry: it spawns the TP "
|
||||
"workers, so their groups do not exist yet"
|
||||
),
|
||||
("srt/disaggregation/encoder/server.py", "tp_size"): (
|
||||
"`MMEncoder` builds its own TP group from this size -- "
|
||||
"`initialize_model_parallel` is the call being handed it, so there is "
|
||||
"nothing live to ask"
|
||||
),
|
||||
("srt/disaggregation/encoder/receiver.py", "tp_size"): (
|
||||
"the receiver labels and shards by the launch width; it runs in the "
|
||||
"tokenizer process, which holds no encoder groups"
|
||||
),
|
||||
("srt/managers/rust_server.py", "tp_size"): (
|
||||
"the rust server decides its transport from the launch width, in the "
|
||||
"tokenizer process, which holds no model groups"
|
||||
),
|
||||
("compile_deep_gemm.py", "tp_size"): (
|
||||
"the warm-up request fans bootstrap rooms across the launch's ranks; it "
|
||||
"runs in the tokenizer process, which holds no model groups"
|
||||
),
|
||||
("srt/utils/common.py", "tp_size"): (
|
||||
"the require_*_tp_gather predicates compared the configured tp_size "
|
||||
"when they read the record; the live property answers a different "
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
"""The Ray driver sizes its actors from the published configuration.
|
||||
|
||||
`RayEngine` publishes as part of `Engine._launch_subprocesses` and *then* lays
|
||||
out the actors, so the placement arithmetic reads the `parallel` bag. That is
|
||||
where a resolution decision lives: a launch that leaves `dp_size` to resolution
|
||||
has it in the `parallel` bag, and the override case below is what tells the two
|
||||
apart.
|
||||
|
||||
There is no CI coverage of the Ray path (`test/manual/test_ray_engine.py` boots a
|
||||
real cluster), so these cases drive the two pure helpers directly against a
|
||||
published config -- including the override direction, which is what tells a bag
|
||||
read from a record read.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
import unittest
|
||||
|
||||
from sglang.srt.runtime_context import get_context, get_parallel
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
||||
|
||||
# `sglang.srt.ray.engine` imports `ray` at module scope, and the CPU runner has
|
||||
# no ray wheel. The file-scoped source scan below is the part that has to run
|
||||
# everywhere; the three arithmetic cases need the import.
|
||||
_HAS_RAY = importlib.util.find_spec("ray") is not None
|
||||
_needs_ray = unittest.skipUnless(_HAS_RAY, "ray is not installed")
|
||||
|
||||
|
||||
class TestRayDriverReadsTheBags(CustomTestCase):
|
||||
def _publish(self, **fields):
|
||||
override = get_context().override_server_args(**fields)
|
||||
override.install()
|
||||
self.addCleanup(override.restore)
|
||||
|
||||
@_needs_ray
|
||||
def test_world_size_multiplies_the_published_sizes(self):
|
||||
from sglang.srt.ray.engine import _compute_world_size
|
||||
|
||||
self._publish(tp_size=2, pp_size=3, dp_size=4, enable_dp_attention=False)
|
||||
self.assertEqual(_compute_world_size(), 24)
|
||||
|
||||
@_needs_ray
|
||||
def test_dp_attention_folds_dp_into_tp(self):
|
||||
from sglang.srt.ray.engine import _compute_world_size
|
||||
|
||||
self._publish(tp_size=4, pp_size=2, dp_size=4, enable_dp_attention=True)
|
||||
# DP attention folds DP into TP, so dp_size drops out of the product.
|
||||
self.assertEqual(_compute_world_size(), 8)
|
||||
|
||||
@_needs_ray
|
||||
def test_the_world_size_follows_a_post_publish_override(self):
|
||||
"""The direction that separates a bag read from a record read.
|
||||
|
||||
`override` writes the bag and never the record, so a driver still
|
||||
reading `server_args.tp_size` would keep answering with the old size.
|
||||
"""
|
||||
from sglang.srt.ray.engine import _compute_world_size
|
||||
|
||||
self._publish(tp_size=2, pp_size=1, dp_size=1, enable_dp_attention=False)
|
||||
self.assertEqual(_compute_world_size(), 2)
|
||||
get_context().override("test.ray_driver", tp_size=8)
|
||||
self.assertEqual(get_parallel().config.tp_size, 8)
|
||||
self.assertEqual(_compute_world_size(), 8)
|
||||
|
||||
def test_the_driver_modules_read_no_field_off_a_record(self):
|
||||
"""File-scoped: neither Ray driver module reads a config field off an
|
||||
instance any more.
|
||||
|
||||
The Ray path has no CI coverage, so this is what keeps a new
|
||||
`server_args.tp_size` from appearing in it -- the placement arithmetic
|
||||
runs after the publish, and the bags are the surface that carries what
|
||||
resolution decided.
|
||||
"""
|
||||
import ast
|
||||
import dataclasses
|
||||
import pathlib
|
||||
|
||||
import sglang
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
fields = {field.name for field in dataclasses.fields(ServerArgs)}
|
||||
srt = pathlib.Path(sglang.__file__).resolve().parent / "srt"
|
||||
offenders = []
|
||||
for rel in ("ray/engine.py", "ray/data_parallel_controller.py"):
|
||||
tree = ast.parse((srt / rel).read_text(encoding="utf-8-sig"))
|
||||
holders = {"server_args", "sa"}
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
|
||||
continue
|
||||
for arg in list(node.args.args) + list(node.args.kwonlyargs):
|
||||
if arg.annotation is not None and "ServerArgs" in ast.dump(
|
||||
arg.annotation
|
||||
):
|
||||
holders.add(arg.arg)
|
||||
for node in ast.walk(tree):
|
||||
if (
|
||||
isinstance(node, ast.Attribute)
|
||||
and node.attr in fields
|
||||
and isinstance(node.ctx, ast.Load)
|
||||
and (
|
||||
(isinstance(node.value, ast.Name) and node.value.id in holders)
|
||||
or (
|
||||
isinstance(node.value, ast.Attribute)
|
||||
and node.value.attr == "server_args"
|
||||
)
|
||||
)
|
||||
):
|
||||
offenders.append(f"{rel}:{node.lineno} reads .{node.attr}")
|
||||
self.assertEqual(
|
||||
offenders,
|
||||
[],
|
||||
"the Ray driver reads a config field off a record; the driver runs "
|
||||
"after the publish, so read `get_parallel().config`:\n "
|
||||
+ "\n ".join(offenders),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -519,8 +519,6 @@ class TestMoeFlagsGroup(_IsolatedServerArgs):
|
||||
swap under the speculative contexts and restore on exit."""
|
||||
|
||||
def _init(self, **kw):
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.layers.moe.utils import initialize_moe_config
|
||||
|
||||
defaults = dict(
|
||||
@@ -538,7 +536,12 @@ class TestMoeFlagsGroup(_IsolatedServerArgs):
|
||||
disable_shared_experts_fusion=False,
|
||||
)
|
||||
defaults.update(kw)
|
||||
initialize_moe_config(SimpleNamespace(**defaults))
|
||||
# The flags are seeded from the bags, so the test publishes a config
|
||||
# carrying these values.
|
||||
override = get_context().override_server_args(**defaults)
|
||||
override.install()
|
||||
self.addCleanup(override.restore)
|
||||
initialize_moe_config()
|
||||
|
||||
def test_lazy_defaults_before_initialize(self):
|
||||
from sglang.srt.layers.moe.utils import (
|
||||
|
||||
@@ -134,51 +134,12 @@ _ENV_MATRIX = (({}, {"SGLANG_IS_IN_CI": "true"}),)
|
||||
# are step-12 exposure like any other pair.
|
||||
_PASSED = frozenset({"model_path", "device", "random_seed"})
|
||||
|
||||
# The reads that still take a value off the supplied instance. `initialize_moe_config`
|
||||
# is handed the record until the replay goes away; the rest are pre-publish launcher
|
||||
# reads.
|
||||
_EXPOSED = {
|
||||
("disaggregation/encoder/server.py", "model_loader_extra_config"),
|
||||
("layers/moe/utils.py", "deepep_mode"),
|
||||
("layers/moe/utils.py", "disable_shared_experts_fusion"),
|
||||
("layers/moe/utils.py", "moe_a2a_backend"),
|
||||
("layers/moe/utils.py", "moe_runner_backend"),
|
||||
("layers/moe/utils.py", "quantization"),
|
||||
("layers/moe/utils.py", "speculative_moe_runner_backend"),
|
||||
("configs/embedding_model_spec.py", "chunked_prefill_size"),
|
||||
("configs/embedding_model_spec.py", "cuda_graph_config"),
|
||||
("configs/embedding_model_spec.py", "disable_radix_cache"),
|
||||
("configs/embedding_model_spec.py", "is_embedding"),
|
||||
("configs/embedding_model_spec.py", "prefill_only_disable_kv_cache"),
|
||||
("entrypoints/engine.py", "enable_symm_mem"),
|
||||
("entrypoints/engine.py", "reasoning_parser"),
|
||||
("entrypoints/engine.py", "tool_call_parser"),
|
||||
("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"),
|
||||
("layers/moe/utils.py", "deepep_mode"),
|
||||
("layers/moe/utils.py", "moe_a2a_backend"),
|
||||
("layers/moe/utils.py", "moe_runner_backend"),
|
||||
("layers/moe/utils.py", "quantization"),
|
||||
("layers/moe/utils.py", "speculative_moe_runner_backend"),
|
||||
("speculative/draft_worker_common.py", "speculative_draft_attention_backend"),
|
||||
("utils/common.py", "speculative_num_draft_tokens"),
|
||||
("utils/common.py", "speculative_num_steps"),
|
||||
("utils/hf_transformers/processor.py", "image_processor_backend"),
|
||||
("weight_cache/daemon.py", "attn_cp_size"),
|
||||
("weight_cache/daemon.py", "deepep_mode"),
|
||||
("weight_cache/daemon.py", "dp_size"),
|
||||
("weight_cache/daemon.py", "dtype"),
|
||||
("weight_cache/daemon.py", "enable_dp_attention"),
|
||||
("weight_cache/daemon.py", "enable_dp_lm_head"),
|
||||
("weight_cache/daemon.py", "ep_size"),
|
||||
("weight_cache/daemon.py", "load_format"),
|
||||
("weight_cache/daemon.py", "model_loader_extra_config"),
|
||||
("weight_cache/daemon.py", "model_path"),
|
||||
("weight_cache/daemon.py", "moe_a2a_backend"),
|
||||
("weight_cache/daemon.py", "moe_dense_tp_size"),
|
||||
("weight_cache/daemon.py", "moe_dp_size"),
|
||||
("weight_cache/daemon.py", "pp_size"),
|
||||
("weight_cache/daemon.py", "quantization"),
|
||||
}
|
||||
# Empty. A pair belongs here when a reader has no bag to read -- it runs before
|
||||
# its process publishes -- and cannot use `resolving_view` either. The launcher's
|
||||
# pre-publish reads (`_set_envs_and_config`, the auto-parser gate) and the
|
||||
# late-resolution detection it calls all read the declarations now, so nothing
|
||||
# qualifies. A new entry needs that kind of reason next to it.
|
||||
_EXPOSED: frozenset = frozenset()
|
||||
|
||||
# Pairs whose resolution write only happens on a CUDA host (capability or
|
||||
# `is_cuda()` gated): asserted on the CUDA registration, invisible to the CPU
|
||||
@@ -191,28 +152,7 @@ _EXPOSED_CUDA_ONLY: frozenset = frozenset()
|
||||
# Axis two: (file, field) pairs where a supplied-instance read names a field that
|
||||
# some code overrides post-publish. Each needs an ordering judgment, not a blanket
|
||||
# conversion; the list exists so a new one is a decision made when it is written.
|
||||
_OVERRIDDEN_AND_READ = {
|
||||
("entrypoints/engine.py", "reasoning_parser"),
|
||||
("entrypoints/engine.py", "tool_call_parser"),
|
||||
("weight_cache/daemon.py", "dp_size"),
|
||||
("weight_cache/daemon.py", "dtype"),
|
||||
("weight_cache/daemon.py", "ep_size"),
|
||||
("weight_cache/daemon.py", "load_format"),
|
||||
("weight_cache/daemon.py", "model_path"),
|
||||
("mem_cache/pool_host/common.py", "hicache_storage_backend"),
|
||||
("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"),
|
||||
("mem_cache/unified_radix_cache.py", "hicache_storage_backend"),
|
||||
("mem_cache/unified_radix_cache.py", "hicache_storage_backend_extra_config"),
|
||||
("mem_cache/unified_radix_cache.py", "hicache_storage_prefetch_policy"),
|
||||
("mem_cache/unified_radix_cache.py", "hicache_write_policy"),
|
||||
("utils/common.py", "speculative_num_draft_tokens"),
|
||||
("utils/common.py", "speculative_num_steps"),
|
||||
("weight_cache/daemon.py", "dp_size"),
|
||||
("weight_cache/daemon.py", "dtype"),
|
||||
("weight_cache/daemon.py", "ep_size"),
|
||||
("weight_cache/daemon.py", "load_format"),
|
||||
("weight_cache/daemon.py", "model_path"),
|
||||
}
|
||||
_OVERRIDDEN_AND_READ: frozenset = frozenset()
|
||||
|
||||
|
||||
def _expanded_override_keys(rel, tree, call, kw) -> set:
|
||||
|
||||
Reference in New Issue
Block a user