[Config] Round 6.4: the runtime reads the bags, not the record (#38049)

Last of four; stacked on #38048.

The record is the operator's input; the bags are what is in effect. A reader
that takes the record and reads a field off it gets the input, which is the
wrong one of the two whenever resolution decided something -- and the mistake is
silent, because for most fields and most launches the two agree. Several of
these files already read both ways, sometimes in the same expression:

```python
get_tokenizer(
    get_serving().tokenizer_path,
    tokenizer_mode=server_args.tokenizer_mode,   # the input, not the decision
    ...
)
```

Sixty-odd files convert. Record field reads in runtime code go from 199 to 11.
Nine parameters that the conversion emptied are dropped along with the argument
at every call site -- the dead-parameter ratchet is what names them.

### "Runs after its process publishes" is a per-entry-point claim

Most converted reads sit in the serving and model-executor layers, which only
exist after publication, or in the two subprocess entry points, which publish
first thing. Three places are not like that, and they keep reading the record
they were handed:

- **`HttpServerEngineAdapter`** launches the server as a *child*. The parent
  resolves the record and never publishes, so the adapter's own reads -- the
  launch banner, the API key in its readiness loop, the TP width in
  `update_weights_from_tensor` -- are of `self.server_args`. A bag read here
  fails closed in a bare process, or answers for an unrelated engine in one that
  happens to have published.
- **`serve_grpc`** reads its sidecar port before the integrated servicer builds
  the `Engine` that publishes. The comment above that line already said so and
  already bound `cfg = resolving_view(server_args)` for it; the sidecar port and
  the port it derives from read `cfg`.
- **`initialize_dp_attention`** runs from callers whose publish is not
  guaranteed, so its one predicate stays on the resolution view.

`ROLE_NAMESPACE_SETS["dp_controller"]` gains `observability` and `serving`,
because the controller's metrics gate, tracing setup and worker-port broadcast
now read those namespaces. Under `SGLANG_ROLE_NAMESPACES=enforce` that set is
what the process may read, so a conversion that reaches a new namespace has to
widen it in the same change.

## Three things worth a reviewer's attention

**Eleven reads were `getattr(record, "field", default)`.** An AST scan for
attribute access does not see those, so the census that said "43 readers" was
counting the shape it could match rather than the thing it was after.
`incremental_streaming_output` was read that way twice, and the transcription
tests were the only reason it surfaced.

**Not every record read is a bag read waiting to happen.** A multimodal
processor's `base_gpu_id` is the instance's, not the process's: two engines in
one process keep different ones, and
`test_publishing_another_config_does_not_move_the_device` exists to say so. It
stays on the record while `rl_on_policy_target` beside it moves.
`RequestMetricsExporter` is the same shape -- it is handed the directory it
writes to, and a test builds several with different ones. `configure_logger` is
a third: 17 call sites, one of which passes an `argparse.Namespace`, so it is
not a global-context reader at all. Those eleven remaining reads are the ones
with a reason.

**The fixtures move with the code.** Tests that hung config off a mock manager
now publish a record, which is what the serving layer reads; where a test states
a value it says so with `override_server_args` instead of assigning through the
mock. `test_hisparse_unit` is the last of them: it stubbed a `server_args` onto
a fake scheduler to say the decode radix cache was off, and the value it was
standing in for is the published default, so the stub goes and the class
publishes.

## Two things CI caught that a local sweep could not

**`unittest.TestCase.enterContext` is Python 3.11+.** The converted fixtures used
it at 18 sites; `requires-python` is `>=3.10` and CI runs 3.10, so every one of
them raised `AttributeError` there while passing on a newer local interpreter.
They call `enter_override(self, ...)` now -- a four-line helper in
`sglang/test/test_utils.py` over the override's own `install()` / `restore()`.

**A batched sweep cannot see a missing publish.** Three fixtures needed a
published config and did not have one; each *passed* inside a shard where some
other file had published, and failed when run alone. The affected cases are
`test_serving_completions` (which set `incremental_streaming_output` on the mock
manager's record, where nothing reads it now), `test_qwen3_vl_feature_materialization`
(same shape for `mm_enable_dp_encoder`), and the two Qwen Rust tests -- whose
fixture already carried the comment `# Non-auto: get_resolved_model_impl would
choke on a SimpleNamespace` next to the `model_impl` it sets, which is exactly
what happened once `get_mm_processor_cls` started reading that value from the
bag. Its `publish` mirrors `model_impl` now, like the four fields it already
mirrored.

## Verification

A full registered-unit sweep (648 files) against this stack's merge-base:
19 failures on both sides, the same 19, none of them config. That sweep is what
caught 23 failures the file-scoped runs missed -- and, later, that the narrower
139-file list did not even contain the files this change reaches. It is also
what caught the `test_hisparse_unit` fixture above: the file passes inside a
shard where something else published, and fails when it is run on its own,
which is why every failing file is re-run alone before it is counted.
This commit is contained in:
Cheng Wan
2026-09-06 21:41:46 -07:00
committed by GitHub
parent 98f69ccbf3
commit b99175dc7d
85 changed files with 711 additions and 413 deletions
@@ -14,6 +14,7 @@ from sglang.srt.constrained.base_grammar_backend import (
from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject
from sglang.srt.distributed.communication_tags import P2PTag from sglang.srt.distributed.communication_tags import P2PTag
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_serving
from sglang.srt.sampling.sampling_params import ( from sglang.srt.sampling.sampling_params import (
get_request_reasoning_end_token_ids, get_request_reasoning_end_token_ids,
) )
@@ -31,7 +32,7 @@ class GrammarManager:
self.scheduler = scheduler self.scheduler = scheduler
self.server_args = scheduler.server_args self.server_args = scheduler.server_args
self.grammar_queue: List[Req] = [] self.grammar_queue: List[Req] = []
if not self.server_args.skip_tokenizer_init: if not get_serving().skip_tokenizer_init:
self.grammar_backend = create_grammar_backend( self.grammar_backend = create_grammar_backend(
self.server_args, self.server_args,
scheduler.tokenizer, scheduler.tokenizer,
+5 -9
View File
@@ -468,11 +468,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
return seq_len return seq_len
page_size = self.token_to_kv_pool_allocator.page_size page_size = self.token_to_kv_pool_allocator.page_size
if getattr( if get_disagg().disaggregation_decode_enable_radix_cache:
self.scheduler.server_args,
"disaggregation_decode_enable_radix_cache",
False,
):
# Keep enough SWA before the page-aligned radix-cache insert # Keep enough SWA before the page-aligned radix-cache insert
# boundary for the cached key to contain a complete window. # boundary for the cached key to contain a complete window.
# `seq_len - 1` is the last committed position. # `seq_len - 1` is the last committed position.
@@ -1194,7 +1190,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
origin_input_len = self._rebootstrap_prefill_len(decode_req.req) origin_input_len = self._rebootstrap_prefill_len(decode_req.req)
prefix_match: Optional[DecodePrefixMatch] = None prefix_match: Optional[DecodePrefixMatch] = None
use_decode_radix_cache = ( use_decode_radix_cache = (
self.scheduler.server_args.disaggregation_decode_enable_radix_cache get_disagg().disaggregation_decode_enable_radix_cache
and not decode_req.is_rebootstrap and not decode_req.is_rebootstrap
) )
if use_decode_radix_cache: if use_decode_radix_cache:
@@ -1648,13 +1644,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
available_size = logical_allocator.available_size() available_size = logical_allocator.available_size()
elif self._uses_swa_tail_prealloc(): elif self._uses_swa_tail_prealloc():
available_size = self.token_to_kv_pool_allocator.full_available_size() available_size = self.token_to_kv_pool_allocator.full_available_size()
if self.scheduler.server_args.disaggregation_decode_enable_radix_cache: if get_disagg().disaggregation_decode_enable_radix_cache:
available_size += self._radix_full_evictable() available_size += self._radix_full_evictable()
else: else:
available_size = self.token_to_kv_pool_allocator.available_size() available_size = self.token_to_kv_pool_allocator.available_size()
# Include evictable decode-radix cache entries in the budget -- they # Include evictable decode-radix cache entries in the budget -- they
# can be freed on demand before allocation. # can be freed on demand before allocation.
if self.scheduler.server_args.disaggregation_decode_enable_radix_cache: if get_disagg().disaggregation_decode_enable_radix_cache:
available_size += self._radix_full_evictable() available_size += self._radix_full_evictable()
allocatable_tokens = available_size - max( allocatable_tokens = available_size - max(
reserved_tokens, need_space_for_single_req reserved_tokens, need_space_for_single_req
@@ -1801,7 +1797,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
# Evict cached entries if the pool doesn't have enough free pages. # Evict cached entries if the pool doesn't have enough free pages.
if ( if (
self.scheduler.server_args.disaggregation_decode_enable_radix_cache get_disagg().disaggregation_decode_enable_radix_cache
and self._radix_full_available() < required_alloc_tokens and self._radix_full_available() < required_alloc_tokens
): ):
num_to_evict = required_alloc_tokens - self._radix_full_available() num_to_evict = required_alloc_tokens - self._radix_full_available()
@@ -251,8 +251,8 @@ async def serve_grpc_encoder(server_args: ServerArgs):
ipc_path_prefix = random_uuid() ipc_path_prefix = random_uuid()
port_args = PortArgs.init_new(server_args) port_args = PortArgs.init_new(server_args)
if server_args.dist_init_addr: if get_parallel().dist_init_addr:
na = NetworkAddress.parse(server_args.dist_init_addr) na = NetworkAddress.parse(get_parallel().dist_init_addr)
dist_init_method = na.to_tcp() dist_init_method = na.to_tcp()
else: else:
dist_init_method = NetworkAddress( dist_init_method = NetworkAddress(
@@ -103,7 +103,7 @@ async def _lifespan(app: FastAPI):
app = FastAPI(lifespan=_lifespan) app = FastAPI(lifespan=_lifespan)
def _register_encoder_url_with_bootstrap(server_args: ServerArgs): def _register_encoder_url_with_bootstrap():
"""Asynchronously register this encoder with each bootstrap URL. """Asynchronously register this encoder with each bootstrap URL.
Spawns a daemon thread that retries each URL independently with bounded Spawns a daemon thread that retries each URL independently with bounded
@@ -118,10 +118,10 @@ def _register_encoder_url_with_bootstrap(server_args: ServerArgs):
host = get_serving().host host = get_serving().host
if not host or host in ("0.0.0.0", "::"): if not host or host in ("0.0.0.0", "::"):
host = get_local_ip_auto(get_serving().host) host = get_local_ip_auto(get_serving().host)
scheme = "https" if server_args.ssl_certfile else "http" scheme = "https" if get_serving().ssl_certfile else "http"
encoder_url = NetworkAddress(host, get_serving().port).to_url(scheme) encoder_url = NetworkAddress(host, get_serving().port).to_url(scheme)
payload = {"url": encoder_url} payload = {"url": encoder_url}
bootstrap_urls = list(server_args.encoder_register_urls) bootstrap_urls = list(get_disagg().encoder_register_urls)
if not bootstrap_urls: if not bootstrap_urls:
return return
@@ -175,15 +175,15 @@ def _register_encoder_url_with_bootstrap(server_args: ServerArgs):
).start() ).start()
def _unregister_encoder_url_from_bootstrap(server_args: ServerArgs): def _unregister_encoder_url_from_bootstrap():
host = get_serving().host host = get_serving().host
if not host or host in ("0.0.0.0", "::"): if not host or host in ("0.0.0.0", "::"):
host = get_local_ip_auto(get_serving().host) host = get_local_ip_auto(get_serving().host)
scheme = "https" if server_args.ssl_certfile else "http" scheme = "https" if get_serving().ssl_certfile else "http"
encoder_url = NetworkAddress(host, get_serving().port).to_url(scheme) encoder_url = NetworkAddress(host, get_serving().port).to_url(scheme)
payload = {"url": encoder_url} payload = {"url": encoder_url}
for bootstrap_url in server_args.encoder_register_urls: for bootstrap_url in get_disagg().encoder_register_urls:
try: try:
resp = http_requests.delete( resp = http_requests.delete(
f"{bootstrap_url}/unregister_encoder_url", f"{bootstrap_url}/unregister_encoder_url",
@@ -229,8 +229,8 @@ def launch_server(server_args: ServerArgs):
if get_disagg().encoder_register_urls: if get_disagg().encoder_register_urls:
import atexit import atexit
_register_encoder_url_with_bootstrap(server_args) _register_encoder_url_with_bootstrap()
atexit.register(_unregister_encoder_url_from_bootstrap, server_args) atexit.register(_unregister_encoder_url_from_bootstrap)
uvicorn.run(app, host=get_serving().host, port=get_serving().port) uvicorn.run(app, host=get_serving().host, port=get_serving().port)
@@ -2000,7 +2000,7 @@ class MMReceiverBase(ABC):
import_processors("sglang.srt.multimodal.processors") import_processors("sglang.srt.multimodal.processors")
extra_kwargs = {} extra_kwargs = {}
if getattr(server_args, "tokenizer_backend", None) is not None: if get_serving().tokenizer_backend is not None:
extra_kwargs["tokenizer_backend"] = get_serving().tokenizer_backend extra_kwargs["tokenizer_backend"] = get_serving().tokenizer_backend
_processor = get_processor( _processor = get_processor(
@@ -52,6 +52,7 @@ from sglang.srt.observability.trace import (
trace_set_thread_info, trace_set_thread_info,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_device,
get_observability, get_observability,
get_parallel, get_parallel,
get_serving, get_serving,
@@ -1867,7 +1868,7 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher:
atexit.register(_kill_workers) atexit.register(_kill_workers)
for dp_rank in range(dp_size): for dp_rank in range(dp_size):
gpu_id = server_args.base_gpu_id + dp_rank gpu_id = get_device().base_gpu_id + dp_rank
# Pin the device parent-side around spawn (same convention as the # Pin the device parent-side around spawn (same convention as the
# scheduler launcher and DP controller) so the child inherits # scheduler launcher and DP controller) so the child inherits
# CUDA_VISIBLE_DEVICES from its first instruction, before any import # CUDA_VISIBLE_DEVICES from its first instruction, before any import
@@ -1890,8 +1891,8 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher:
worker_processes.append(process) worker_processes.append(process)
labels = {"model_name": get_serving().served_model_name} labels = {"model_name": get_serving().served_model_name}
if server_args.extra_metric_labels: if get_observability().extra_metric_labels:
labels.update(server_args.extra_metric_labels) labels.update(get_observability().extra_metric_labels)
return DPDispatcher( return DPDispatcher(
dp_size, dp_size,
dispatch_sockets, dispatch_sockets,
@@ -548,7 +548,7 @@ class MMEncoder:
self.server_args = server_args self.server_args = server_args
configure_media_url_security( configure_media_url_security(
get_mm().allowed_media_domains, get_mm().allowed_media_domains,
server_args.media_url_max_file_size_mb, get_mm().media_url_max_file_size_mb,
) )
self.transfer_backend = get_disagg().encoder_transfer_backend self.transfer_backend = get_disagg().encoder_transfer_backend
self.use_mooncake = self.transfer_backend == "mooncake" self.use_mooncake = self.transfer_backend == "mooncake"
@@ -563,18 +563,18 @@ class MMEncoder:
) )
self.load_config = LoadConfig( self.load_config = LoadConfig(
load_format=get_model().load_format, load_format=get_model().load_format,
download_dir=server_args.download_dir, download_dir=get_model().download_dir,
model_loader_extra_config=get_model().model_loader_extra_config, model_loader_extra_config=get_model().model_loader_extra_config,
remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip, remote_instance_weight_loader_seed_instance_ip=get_model().remote_instance_weight_loader_seed_instance_ip,
remote_instance_weight_loader_seed_instance_service_port=server_args.remote_instance_weight_loader_seed_instance_service_port, remote_instance_weight_loader_seed_instance_service_port=get_model().remote_instance_weight_loader_seed_instance_service_port,
remote_instance_weight_loader_send_weights_group_ports=server_args.remote_instance_weight_loader_send_weights_group_ports, remote_instance_weight_loader_send_weights_group_ports=get_model().remote_instance_weight_loader_send_weights_group_ports,
) )
self.model_type = getattr( self.model_type = getattr(
self.model_config.hf_config, "model_type", "unknown" self.model_config.hf_config, "model_type", "unknown"
).lower() ).lower()
self.device = get_device().device self.device = get_device().device
self.gpu_id = server_args.base_gpu_id + rank if gpu_id is None else gpu_id self.gpu_id = get_device().base_gpu_id + rank if gpu_id is None else gpu_id
self.device_config = DeviceConfig( self.device_config = DeviceConfig(
device=self.device, device=self.device,
+8 -11
View File
@@ -81,14 +81,14 @@ def init_torch_distributed(
tic = time.perf_counter() tic = time.perf_counter()
logger.info("Init torch distributed begin.") logger.info("Init torch distributed begin.")
backend = _resolve_backend(device=device, server_args=server_args) backend = _resolve_backend(device=device)
before_avail_memory = get_available_gpu_memory(device, ps.gpu_id) before_avail_memory = get_available_gpu_memory(device, ps.gpu_id)
if not get_parallel().enable_p2p_check: if not get_parallel().enable_p2p_check:
monkey_patch_p2p_access_check() monkey_patch_p2p_access_check()
dist_init_method = _resolve_dist_init_method(dist_port=dist_port) dist_init_method = _resolve_dist_init_method(dist_port=dist_port)
_set_all_reduce_flags(server_args=server_args) _set_all_reduce_flags()
if not is_draft_worker: if not is_draft_worker:
if device == "cpu": if device == "cpu":
@@ -173,9 +173,9 @@ def init_torch_distributed(
) )
def _resolve_backend(*, device: str, server_args: ServerArgs) -> str: def _resolve_backend(*, device: str) -> str:
backend = get_default_distributed_backend(device) backend = get_default_distributed_backend(device)
if device == "cuda" and server_args.elastic_ep_backend == "mooncake": if device == "cuda" and get_exec().moe.elastic_ep_backend == "mooncake":
backend = "mooncake" backend = "mooncake"
return backend return backend
@@ -198,9 +198,9 @@ def _resolve_dist_init_method(*, dist_port: int) -> str:
return dist_init_method return dist_init_method
def _set_all_reduce_flags(*, server_args: ServerArgs) -> None: def _set_all_reduce_flags() -> None:
set_custom_all_reduce(not get_exec().comm.disable_custom_all_reduce) set_custom_all_reduce(not get_exec().comm.disable_custom_all_reduce)
set_mscclpp_all_reduce(server_args.enable_mscclpp) set_mscclpp_all_reduce(get_exec().comm.enable_mscclpp)
set_torch_symm_mem_all_reduce(get_exec().comm.enable_torch_symm_mem) set_torch_symm_mem_all_reduce(get_exec().comm.enable_torch_symm_mem)
set_flashinfer_allreduce_only( set_flashinfer_allreduce_only(
get_exec().comm.flashinfer_allreduce_fusion_backend is not None get_exec().comm.flashinfer_allreduce_fusion_backend is not None
@@ -295,7 +295,7 @@ def _init_parallel_groups(
duplicate_tp_group=get_disagg().enable_pdmux, duplicate_tp_group=get_disagg().enable_pdmux,
duplicate_attn_cp_group=( duplicate_attn_cp_group=(
is_hip() is_hip()
and server_args.enable_two_batch_overlap and get_exec().overlap.enable_two_batch_overlap
and get_parallel().enable_dsa_prefill_context_parallel and get_parallel().enable_dsa_prefill_context_parallel
), ),
enable_symm_mem=get_exec().comm.enable_symm_mem, enable_symm_mem=get_exec().comm.enable_symm_mem,
@@ -308,10 +308,7 @@ def _init_parallel_groups(
server_args=server_args, server_args=server_args,
model_config=model_config, model_config=model_config,
) )
initialize_layernorm_sp( initialize_layernorm_sp(model_config=model_config)
server_args=server_args,
model_config=model_config,
)
if is_npu(): if is_npu():
register_sgl_tp_rank(gpu_id) register_sgl_tp_rank(gpu_id)
+3 -2
View File
@@ -6,6 +6,7 @@ from http import HTTPStatus
from fastapi import APIRouter, Request from fastapi import APIRouter, Request
from fastapi.responses import ORJSONResponse from fastapi.responses import ORJSONResponse
from sglang.srt.runtime_context import get_exec
from sglang.srt.utils.auth import AuthLevel, auth_level from sglang.srt.utils.auth import AuthLevel, auth_level
router = APIRouter() router = APIRouter()
@@ -43,7 +44,7 @@ async def scale_elastic_ep(raw_request: Request):
from sglang.srt.entrypoints.http_server import _global_state from sglang.srt.entrypoints.http_server import _global_state
from sglang.srt.managers.io_struct import ScaleElasticEPReqInput from sglang.srt.managers.io_struct import ScaleElasticEPReqInput
if _global_state.tokenizer_manager.server_args.elastic_ep_backend is None: if get_exec().moe.elastic_ep_backend is None:
return ORJSONResponse( return ORJSONResponse(
{"error": "elastic EP is not enabled (set --elastic-ep-backend)"}, {"error": "elastic EP is not enabled (set --elastic-ep-backend)"},
status_code=HTTPStatus.NOT_FOUND, status_code=HTTPStatus.NOT_FOUND,
@@ -78,7 +79,7 @@ async def is_scaling_elastic_ep(raw_request: Request):
"""Return the tokenizer's mirrored Elastic EP scale state.""" """Return the tokenizer's mirrored Elastic EP scale state."""
from sglang.srt.entrypoints.http_server import _global_state from sglang.srt.entrypoints.http_server import _global_state
if _global_state.tokenizer_manager.server_args.elastic_ep_backend is None: if get_exec().moe.elastic_ep_backend is None:
return ORJSONResponse( return ORJSONResponse(
{"error": "elastic EP is not enabled (set --elastic-ep-backend)"}, {"error": "elastic EP is not enabled (set --elastic-ep-backend)"},
status_code=HTTPStatus.NOT_FOUND, status_code=HTTPStatus.NOT_FOUND,
+42 -40
View File
@@ -104,9 +104,11 @@ from sglang.srt.parser.template_manager import TemplateManager
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.plugins import load_plugins from sglang.srt.plugins import load_plugins
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_device,
get_disagg, get_disagg,
get_exec, get_exec,
get_model, get_model,
get_observability,
get_parallel, get_parallel,
get_serving, get_serving,
publish, publish,
@@ -315,7 +317,7 @@ class Engine(EngineScoreMixin, EngineBase):
# Initialize ZMQ sockets # Initialize ZMQ sockets
context = zmq.Context(2) context = zmq.Context(2)
if self.server_args.node_rank == 0: if get_parallel().node_rank == 0:
self.send_to_rpc = get_zmq_socket( self.send_to_rpc = get_zmq_socket(
context, zmq.DEALER, self.port_args.rpc_ipc_name, True context, zmq.DEALER, self.port_args.rpc_ipc_name, True
) )
@@ -323,11 +325,11 @@ class Engine(EngineScoreMixin, EngineBase):
self.send_to_rpc = None self.send_to_rpc = None
# Enable tracing # Enable tracing
if server_args.enable_trace: if get_observability().enable_trace:
process_tracing_init( process_tracing_init(
server_args.otlp_traces_endpoint, get_observability().otlp_traces_endpoint,
"sglang", "sglang",
trace_modules=server_args.trace_modules, trace_modules=get_observability().trace_modules,
) )
thread_label = "Tokenizer" thread_label = "Tokenizer"
if get_disagg().disaggregation_mode == "prefill": if get_disagg().disaggregation_mode == "prefill":
@@ -691,27 +693,27 @@ class Engine(EngineScoreMixin, EngineBase):
# Multi-node needs an explicit rendezvous address; otherwise each node # Multi-node needs an explicit rendezvous address; otherwise each node
# picks its own local 127.0.0.1 port (below) and the per-node daemons # picks its own local 127.0.0.1 port (below) and the per-node daemons
# can never form the joint process group. # can never form the joint process group.
if server_args.nnodes > 1 and not server_args.dist_init_addr: if get_parallel().nnodes > 1 and not get_parallel().dist_init_addr:
raise ValueError( raise ValueError(
"Multi-node weight cache daemons (nnodes > 1) require " "Multi-node weight cache daemons (nnodes > 1) require "
"--dist-init-addr so all nodes rendezvous at the same endpoint." "--dist-init-addr so all nodes rendezvous at the same endpoint."
) )
tp_size = server_args.tp_size tp_size = get_parallel().tp_size
pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = ( pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = (
_calculate_rank_ranges( _calculate_rank_ranges(
server_args.nnodes, get_parallel().nnodes,
get_parallel().pp_size, get_parallel().pp_size,
tp_size, tp_size,
server_args.node_rank, get_parallel().node_rank,
) )
) )
# Build the distributed init method (multi-node uses the user-provided # Build the distributed init method (multi-node uses the user-provided
# dist_init_addr so all nodes reach the same endpoint). # dist_init_addr so all nodes reach the same endpoint).
if server_args.dist_init_addr: if get_parallel().dist_init_addr:
host, port = server_args.dist_init_addr.rsplit(":", 1) host, port = get_parallel().dist_init_addr.rsplit(":", 1)
dist_init_method = f"tcp://{host}:{port}" dist_init_method = f"tcp://{host}:{port}"
else: else:
# Fresh free port for the daemons' own rendezvous, not the engine's # Fresh free port for the daemons' own rendezvous, not the engine's
@@ -723,7 +725,7 @@ class Engine(EngineScoreMixin, EngineBase):
daemon_procs = [] daemon_procs = []
logger.info( logger.info(
f"Launching {num_daemons} weight cache daemon(s) on node " f"Launching {num_daemons} weight cache daemon(s) on node "
f"{server_args.node_rank} for model={get_model().model_path}, " f"{get_parallel().node_rank} for model={get_model().model_path}, "
f"pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, " f"pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, "
f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, " f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, "
f"dist_init_method={dist_init_method}" f"dist_init_method={dist_init_method}"
@@ -738,8 +740,8 @@ class Engine(EngineScoreMixin, EngineBase):
tp_rank, tp_rank,
pp_size_per_node, pp_size_per_node,
tp_size_per_node, tp_size_per_node,
base_gpu_id=server_args.base_gpu_id, base_gpu_id=get_device().base_gpu_id,
gpu_id_step=server_args.gpu_id_step, gpu_id_step=get_device().gpu_id_step,
) )
cleanup_stale_daemon_files(current_platform.get_device_uuid(gpu_id)) cleanup_stale_daemon_files(current_platform.get_device_uuid(gpu_id))
@@ -750,8 +752,8 @@ class Engine(EngineScoreMixin, EngineBase):
tp_rank, tp_rank,
pp_size_per_node, pp_size_per_node,
tp_size_per_node, tp_size_per_node,
base_gpu_id=server_args.base_gpu_id, base_gpu_id=get_device().base_gpu_id,
gpu_id_step=server_args.gpu_id_step, gpu_id_step=get_device().gpu_id_step,
) )
proc = spawn_weight_cache_daemon( proc = spawn_weight_cache_daemon(
server_args, server_args,
@@ -767,7 +769,7 @@ class Engine(EngineScoreMixin, EngineBase):
# (readiness timeout or a daemon exiting early) terminate the siblings # (readiness timeout or a daemon exiting early) terminate the siblings
# we already spawned before propagating, so a partial launch does not # we already spawned before propagating, so a partial launch does not
# leak GPU-resident daemons. # leak GPU-resident daemons.
timeout = server_args.weight_cache_timeout timeout = get_model().weight_cache_timeout
check_interval = 2 check_interval = 2
start_time = time.time() start_time = time.time()
try: try:
@@ -778,8 +780,8 @@ class Engine(EngineScoreMixin, EngineBase):
tp_rank, tp_rank,
pp_size_per_node, pp_size_per_node,
tp_size_per_node, tp_size_per_node,
base_gpu_id=server_args.base_gpu_id, base_gpu_id=get_device().base_gpu_id,
gpu_id_step=server_args.gpu_id_step, gpu_id_step=get_device().gpu_id_step,
) )
ready_path = get_ready_path( ready_path = get_ready_path(
current_platform.get_device_uuid(gpu_id) current_platform.get_device_uuid(gpu_id)
@@ -809,7 +811,7 @@ class Engine(EngineScoreMixin, EngineBase):
logger.info( logger.info(
f"All {num_daemons} weight cache daemons on node " f"All {num_daemons} weight cache daemons on node "
f"{server_args.node_rank} are ready" f"{get_parallel().node_rank} are ready"
) )
return daemon_procs return daemon_procs
@@ -864,16 +866,16 @@ class Engine(EngineScoreMixin, EngineBase):
if not use_dp_controller: if not use_dp_controller:
# Launch tensor parallel scheduler processes # Launch tensor parallel scheduler processes
memory_saver_adapter = TorchMemorySaverAdapter.create( memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=server_args.enable_memory_saver enable=get_exec().features.enable_memory_saver
) )
scheduler_pipe_readers = [] scheduler_pipe_readers = []
pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = ( pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = (
_calculate_rank_ranges( _calculate_rank_ranges(
server_args.nnodes, get_parallel().nnodes,
get_parallel().pp_size, get_parallel().pp_size,
server_args.tp_size, get_parallel().tp_size,
server_args.node_rank, get_parallel().node_rank,
) )
) )
@@ -881,12 +883,12 @@ class Engine(EngineScoreMixin, EngineBase):
for tp_rank in tp_rank_range: for tp_rank in tp_rank_range:
reader, writer = mp.Pipe(duplex=False) reader, writer = mp.Pipe(duplex=False)
gpu_id = ( gpu_id = (
server_args.base_gpu_id get_device().base_gpu_id
+ ((pp_rank % pp_size_per_node) * tp_size_per_node) + ((pp_rank % pp_size_per_node) * tp_size_per_node)
+ (tp_rank % tp_size_per_node) * server_args.gpu_id_step + (tp_rank % tp_size_per_node) * get_device().gpu_id_step
) )
attn_cp_rank, moe_dp_rank, moe_ep_rank = _compute_parallelism_ranks( attn_cp_rank, moe_dp_rank, moe_ep_rank = _compute_parallelism_ranks(
server_args, tp_rank tp_rank
) )
with maybe_reindex_device_id(gpu_id) as gpu_id: with maybe_reindex_device_id(gpu_id) as gpu_id:
@@ -1109,9 +1111,9 @@ class Engine(EngineScoreMixin, EngineBase):
engine_info_bootstrap_server = None engine_info_bootstrap_server = None
if ( if (
get_model().remote_instance_weight_loader_start_seed_via_transfer_engine get_model().remote_instance_weight_loader_start_seed_via_transfer_engine
and server_args.node_rank == 0 and get_parallel().node_rank == 0
): ):
bootstrap_port = server_args.engine_info_bootstrap_port bootstrap_port = get_model().engine_info_bootstrap_port
if not is_port_available(bootstrap_port): if not is_port_available(bootstrap_port):
raise RuntimeError( raise RuntimeError(
f"engine_info_bootstrap_port {bootstrap_port} is already in use. " f"engine_info_bootstrap_port {bootstrap_port} is already in use. "
@@ -1119,13 +1121,13 @@ class Engine(EngineScoreMixin, EngineBase):
f"different --engine-info-bootstrap-port." f"different --engine-info-bootstrap-port."
) )
engine_info_bootstrap_server = EngineInfoBootstrapServer( engine_info_bootstrap_server = EngineInfoBootstrapServer(
host=server_args.host, port=bootstrap_port host=get_serving().host, port=bootstrap_port
) )
# Launch daemons (daemon mode only). The handles travel back to the # Launch daemons (daemon mode only). The handles travel back to the
# Engine that spawned them; shutdown() reaps from there. # Engine that spawned them; shutdown() reaps from there.
weight_cache_daemon_procs: List = [] weight_cache_daemon_procs: List = []
if server_args.weight_cache_mode == "daemon": if get_model().weight_cache_mode == "daemon":
weight_cache_daemon_procs = cls._launch_weight_cache_daemons( weight_cache_daemon_procs = cls._launch_weight_cache_daemons(
server_args server_args
) )
@@ -1147,12 +1149,12 @@ class Engine(EngineScoreMixin, EngineBase):
) )
if ( if (
server_args.enable_elastic_expert_backup get_exec().moe.enable_elastic_expert_backup
and server_args.elastic_ep_backend is not None and get_exec().moe.elastic_ep_backend is not None
): ):
run_expert_backup_manager(server_args, port_args) run_expert_backup_manager(server_args, port_args)
if server_args.node_rank >= 1: if get_parallel().node_rank >= 1:
# Non-zero-rank nodes do not run tokenizer processes. # Non-zero-rank nodes do not run tokenizer processes.
scheduler_init_result.wait_for_ready() scheduler_init_result.wait_for_ready()
@@ -1168,7 +1170,9 @@ class Engine(EngineScoreMixin, EngineBase):
) )
launch_dummy_health_check_server( launch_dummy_health_check_server(
server_args.host, server_args.port, server_args.enable_metrics get_serving().host,
get_serving().port,
get_observability().enable_metrics,
) )
scheduler_init_result.block_until_scheduler_exits() scheduler_init_result.block_until_scheduler_exits()
@@ -1215,7 +1219,7 @@ class Engine(EngineScoreMixin, EngineBase):
scheduler_init_result.all_child_pids.append(p.pid) scheduler_init_result.all_child_pids.append(p.pid)
# Init tokenizer manager first, as the bootstrap server is initialized here # Init tokenizer manager first, as the bootstrap server is initialized here
if server_args.tokenizer_worker_num == 1: if get_serving().tokenizer_worker_num == 1:
tokenizer_manager, template_manager = init_tokenizer_manager_func( tokenizer_manager, template_manager = init_tokenizer_manager_func(
server_args, port_args server_args, port_args
) )
@@ -1534,7 +1538,7 @@ class Engine(EngineScoreMixin, EngineBase):
else: else:
return [ return [
MultiprocessingSerializer.serialize(tensors) MultiprocessingSerializer.serialize(tensors)
for _ in range(self.server_args.tp_size) for _ in range(get_parallel().tp_size)
] ]
def load_lora_adapter_from_tensors( def load_lora_adapter_from_tensors(
@@ -1866,16 +1870,14 @@ def _calculate_rank_ranges(
return pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node return pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node
def _compute_parallelism_ranks( def _compute_parallelism_ranks(tp_rank: int) -> Tuple[int, int, int]:
server_args: ServerArgs, tp_rank: int
) -> Tuple[int, int, int]:
"""Compute attention-CP, MoE-DP, and MoE-EP ranks for a TP rank. """Compute attention-CP, MoE-DP, and MoE-EP ranks for a TP rank.
Called while the launcher is deciding what to spawn, so the sizes are the Called while the launcher is deciding what to spawn, so the sizes are the
configured ones -- the groups this is laying out do not exist yet. configured ones -- the groups this is laying out do not exist yet.
""" """
attn_dp_size = get_parallel().dp_size if get_parallel().enable_dp_attention else 1 attn_dp_size = get_parallel().dp_size if get_parallel().enable_dp_attention else 1
tp_size = server_args.tp_size tp_size = get_parallel().tp_size
attn_cp_size = get_parallel().attn_cp_size attn_cp_size = get_parallel().attn_cp_size
moe_dp_size = get_parallel().moe_dp_size moe_dp_size = get_parallel().moe_dp_size
+3 -3
View File
@@ -176,9 +176,9 @@ async def serve_grpc(server_args, model_info=None):
sidecar_app = web.Application() sidecar_app = web.Application()
sidecar_runner = None sidecar_runner = None
sidecar_port = ( sidecar_port = (
server_args.smg_http_sidecar_port cfg.smg_http_sidecar_port
if server_args.smg_http_sidecar_port is not None if cfg.smg_http_sidecar_port is not None
else server_args.port + 1 else cfg.port + 1
) )
# Metrics setup: must set PROMETHEUS_MULTIPROC_DIR before scheduler # Metrics setup: must set PROMETHEUS_MULTIPROC_DIR before scheduler
+73 -73
View File
@@ -251,8 +251,8 @@ async def init_multi_tokenizer() -> ServerArgs:
template_manager.initialize_templates( template_manager.initialize_templates(
tokenizer_manager=tokenizer_manager, tokenizer_manager=tokenizer_manager,
model_path=get_model().model_path, model_path=get_model().model_path,
chat_template=server_args.chat_template, chat_template=get_serving().chat_template,
completion_template=server_args.completion_template, completion_template=get_serving().completion_template,
) )
tokenizer_manager.max_req_input_len = scheduler_info["max_req_input_len"] tokenizer_manager.max_req_input_len = scheduler_info["max_req_input_len"]
@@ -290,11 +290,11 @@ async def lifespan(fast_api_app: FastAPI):
enable_func_timer() enable_func_timer()
# Init tracing # Init tracing
if server_args.enable_trace: if get_observability().enable_trace:
process_tracing_init( process_tracing_init(
server_args.otlp_traces_endpoint, get_observability().otlp_traces_endpoint,
"sglang", "sglang",
trace_modules=server_args.trace_modules, trace_modules=get_observability().trace_modules,
) )
if get_disagg().disaggregation_mode == "prefill": if get_disagg().disaggregation_mode == "prefill":
thread_label = "Prefill" + thread_label thread_label = "Prefill" + thread_label
@@ -343,15 +343,15 @@ async def lifespan(fast_api_app: FastAPI):
# Launch tool server # Launch tool server
tool_server = None tool_server = None
if server_args.tool_server == "demo": if get_serving().tool_server == "demo":
from sglang.srt.entrypoints.openai.tool_server import DemoToolServer from sglang.srt.entrypoints.openai.tool_server import DemoToolServer
tool_server = DemoToolServer() tool_server = DemoToolServer()
elif server_args.tool_server: elif get_serving().tool_server:
from sglang.srt.entrypoints.openai.tool_server import MCPToolServer from sglang.srt.entrypoints.openai.tool_server import MCPToolServer
tool_server = MCPToolServer() tool_server = MCPToolServer()
await tool_server.add_tool_server(server_args.tool_server) await tool_server.add_tool_server(get_serving().tool_server)
elif envs.EXA_API_KEY.get(): elif envs.EXA_API_KEY.get():
from sglang.srt.entrypoints.openai.tool_server import NativeToolServer from sglang.srt.entrypoints.openai.tool_server import NativeToolServer
@@ -381,10 +381,10 @@ async def lifespan(fast_api_app: FastAPI):
) )
# Execute custom warmups # Execute custom warmups
if server_args.warmups is not None: if get_serving().warmups is not None:
await execute_warmups( await execute_warmups(
get_disagg().disaggregation_mode, get_disagg().disaggregation_mode,
server_args.warmups.split(","), get_serving().warmups.split(","),
_global_state.tokenizer_manager, _global_state.tokenizer_manager,
) )
logger.info("Warmup ended") logger.info("Warmup ended")
@@ -397,7 +397,7 @@ async def lifespan(fast_api_app: FastAPI):
if ( if (
getattr(fast_api_app, "is_single_tokenizer_mode", False) getattr(fast_api_app, "is_single_tokenizer_mode", False)
and get_serving().grpc_port is not None and get_serving().grpc_port is not None
and not (get_serving().smg_grpc_mode or server_args.grpc_mode) and not (get_serving().smg_grpc_mode or get_serving().grpc_mode)
): ):
grpc_handle = _start_native_grpc_server_for_runtime( grpc_handle = _start_native_grpc_server_for_runtime(
server_args=server_args, server_args=server_args,
@@ -406,7 +406,7 @@ async def lifespan(fast_api_app: FastAPI):
scheduler_info=_global_state.scheduler_info, scheduler_info=_global_state.scheduler_info,
grpc_port=get_serving().grpc_port, grpc_port=get_serving().grpc_port,
) )
if server_args.sidecar is not None: if get_serving().sidecar is not None:
from sglang.srt.entrypoints.sidecar import start_sidecar from sglang.srt.entrypoints.sidecar import start_sidecar
sidecar = start_sidecar() sidecar = start_sidecar()
@@ -1101,7 +1101,7 @@ async def attach_hicache_storage_backend(
Only allowed when there are NO running / queued requests. Only allowed when there are NO running / queued requests.
""" """
if not _global_state.tokenizer_manager.server_args.admin_api_key: if not get_serving().admin_api_key:
return _admin_api_key_missing_response() return _admin_api_key_missing_response()
ret = await _global_state.tokenizer_manager.attach_hicache_storage( ret = await _global_state.tokenizer_manager.attach_hicache_storage(
@@ -1133,7 +1133,7 @@ async def detach_hicache_storage_backend():
Only allowed when there are NO running / queued requests. Only allowed when there are NO running / queued requests.
""" """
if not _global_state.tokenizer_manager.server_args.admin_api_key: if not get_serving().admin_api_key:
return _admin_api_key_missing_response() return _admin_api_key_missing_response()
ret = await _global_state.tokenizer_manager.detach_hicache_storage() ret = await _global_state.tokenizer_manager.detach_hicache_storage()
@@ -1157,7 +1157,7 @@ async def detach_hicache_storage_backend():
@auth_level(AuthLevel.ADMIN_OPTIONAL) @auth_level(AuthLevel.ADMIN_OPTIONAL)
async def hicache_storage_backend_status(): async def hicache_storage_backend_status():
"""Get current HiCache storage backend status (tokenizer-side view).""" """Get current HiCache storage backend status (tokenizer-side view)."""
if not _global_state.tokenizer_manager.server_args.admin_api_key: if not get_serving().admin_api_key:
return _admin_api_key_missing_response() return _admin_api_key_missing_response()
return { return {
@@ -2203,8 +2203,8 @@ async def _send_disaggregation_warmup_requests(
def _execute_server_warmup(server_args: ServerArgs): def _execute_server_warmup(server_args: ServerArgs):
headers = {} headers = {}
url = server_args.url() url = server_args.url()
if server_args.api_key: if get_serving().api_key:
headers["Authorization"] = f"Bearer {server_args.api_key}" headers["Authorization"] = f"Bearer {get_serving().api_key}"
ssl_verify = ssl_verify_of(server_args) ssl_verify = ssl_verify_of(server_args)
@@ -2235,12 +2235,12 @@ def _execute_server_warmup(server_args: ServerArgs):
# disaggregation, but its local warmup must stay on the text path. # disaggregation, but its local warmup must stay on the text path.
is_vlm = ( is_vlm = (
bool(model_info.get("has_image_understanding", False)) bool(model_info.get("has_image_understanding", False))
and not server_args.language_only and not get_disagg().language_only
and not server_args.language_model_only and not get_disagg().language_model_only
and not is_mps() and not is_mps()
) )
if model_info["is_generation"]: if model_info["is_generation"]:
if is_vlm and not server_args.skip_tokenizer_init: if is_vlm and not get_serving().skip_tokenizer_init:
request_name = "/v1/chat/completions" request_name = "/v1/chat/completions"
else: else:
request_name = "/generate" request_name = "/generate"
@@ -2253,7 +2253,7 @@ def _execute_server_warmup(server_args: ServerArgs):
"max_new_tokens": max_new_tokens, "max_new_tokens": max_new_tokens,
}, },
} }
if server_args.skip_tokenizer_init: if get_serving().skip_tokenizer_init:
json_data["input_ids"] = [[10, 11, 12] for _ in range(get_parallel().dp_size)] json_data["input_ids"] = [[10, 11, 12] for _ in range(get_parallel().dp_size)]
# TODO Workaround the bug that embedding errors for list of size 1 # TODO Workaround the bug that embedding errors for list of size 1
if get_parallel().dp_size == 1: if get_parallel().dp_size == 1:
@@ -2306,10 +2306,10 @@ def _execute_server_warmup(server_args: ServerArgs):
json_data["text"] = json_data["text"][0] json_data["text"] = json_data["text"][0]
# Config debug dumping # Config debug dumping
if server_args.debug_tensor_dump_input_file: if get_observability().debug_tensor_dump_input_file:
json_data.pop("text", None) json_data.pop("text", None)
json_data["input_ids"] = np.load( json_data["input_ids"] = np.load(
server_args.debug_tensor_dump_input_file get_observability().debug_tensor_dump_input_file
).tolist() ).tolist()
json_data["sampling_params"]["max_new_tokens"] = 0 json_data["sampling_params"]["max_new_tokens"] = 0
@@ -2373,7 +2373,7 @@ def _execute_server_warmup(server_args: ServerArgs):
def _freeze_gc_after_server_warmup(server_args: ServerArgs): def _freeze_gc_after_server_warmup(server_args: ServerArgs):
# Freeze GC after server warmup so static objects skip future GC gen2 collection. # Freeze GC after server warmup so static objects skip future GC gen2 collection.
# Use /freeze_gc to freeze scheduler and detokenizer as well. # Use /freeze_gc to freeze scheduler and detokenizer as well.
freeze_key = server_args.admin_api_key or server_args.api_key freeze_key = get_serving().admin_api_key or get_serving().api_key
freeze_headers = {} freeze_headers = {}
if freeze_key: if freeze_key:
freeze_headers["Authorization"] = f"Bearer {freeze_key}" freeze_headers["Authorization"] = f"Bearer {freeze_key}"
@@ -2394,7 +2394,7 @@ def _wait_and_warmup(
launch_callback: Optional[Callable[[], None]] = None, launch_callback: Optional[Callable[[], None]] = None,
execute_warmup_func: Callable = _execute_server_warmup, execute_warmup_func: Callable = _execute_server_warmup,
): ):
if server_args.checkpoint_engine_wait_weights_before_ready: if get_model().checkpoint_engine_wait_weights_before_ready:
_wait_weights_ready() _wait_weights_ready()
# Joiner schedulers are served through the primary after adoption. # Joiner schedulers are served through the primary after adoption.
@@ -2416,10 +2416,10 @@ def _wait_and_warmup(
# The server is ready for requests # The server is ready for requests
logger.info("The server is fired up and ready to roll!") logger.info("The server is fired up and ready to roll!")
if server_args.delete_ckpt_after_loading: if get_model().delete_ckpt_after_loading:
delete_directory(get_model().model_path) delete_directory(get_model().model_path)
if server_args.debug_tensor_dump_input_file: if get_observability().debug_tensor_dump_input_file:
kill_process_tree(os.getpid()) kill_process_tree(os.getpid())
if launch_callback is not None: if launch_callback is not None:
@@ -2559,7 +2559,7 @@ def _setup_and_run_http_server(
# Pass additional arguments to the lifespan function. # Pass additional arguments to the lifespan function.
# They will be used for additional initialization setups. # They will be used for additional initialization setups.
if server_args.tokenizer_worker_num == 1: if get_serving().tokenizer_worker_num == 1:
# If it is single tokenizer mode, we can pass the arguments by attributes of the app object. # If it is single tokenizer mode, we can pass the arguments by attributes of the app object.
app.is_single_tokenizer_mode = True app.is_single_tokenizer_mode = True
app.server_args = server_args app.server_args = server_args
@@ -2577,16 +2577,16 @@ def _setup_and_run_http_server(
# - no keys: legacy had no restriction; ADMIN_FORCE endpoints must still be rejected when # - no keys: legacy had no restriction; ADMIN_FORCE endpoints must still be rejected when
# admin_api_key is not configured. # admin_api_key is not configured.
if ( if (
server_args.api_key get_serving().api_key
or server_args.admin_api_key or get_serving().admin_api_key
or app_has_admin_force_endpoints(app) or app_has_admin_force_endpoints(app)
): ):
from sglang.srt.utils.auth import add_api_key_middleware from sglang.srt.utils.auth import add_api_key_middleware
add_api_key_middleware( add_api_key_middleware(
app, app,
api_key=server_args.api_key, api_key=get_serving().api_key,
admin_api_key=server_args.admin_api_key, admin_api_key=get_serving().admin_api_key,
) )
else: else:
# If it is multi-tokenizer mode, we need to write the arguments to shared memory # If it is multi-tokenizer mode, we need to write the arguments to shared memory
@@ -2605,15 +2605,15 @@ def _setup_and_run_http_server(
# Update logging configs # Update logging configs
set_uvicorn_logging_configs(server_args) set_uvicorn_logging_configs(server_args)
if server_args.ssl_certfile: if get_serving().ssl_certfile:
logger.info( logger.info(
f"SSL enabled: certfile={server_args.ssl_certfile}, " f"SSL enabled: certfile={get_serving().ssl_certfile}, "
f"keyfile={server_args.ssl_keyfile}" f"keyfile={get_serving().ssl_keyfile}"
) )
# Listen for HTTP requests # Listen for HTTP requests
if server_args.tokenizer_worker_num == 1: if get_serving().tokenizer_worker_num == 1:
if server_args.enable_http2: if get_serving().enable_http2:
logger.info( logger.info(
f"Starting embedded Granian HTTP/2 server on " f"Starting embedded Granian HTTP/2 server on "
f"{get_serving().host}:{get_serving().port}" f"{get_serving().host}:{get_serving().port}"
@@ -2624,32 +2624,32 @@ def _setup_and_run_http_server(
log_level=get_observability().log_level_http log_level=get_observability().log_level_http
or get_observability().log_level, or get_observability().log_level,
http2_max_concurrent_streams=( http2_max_concurrent_streams=(
server_args.http2_max_concurrent_streams get_serving().http2_max_concurrent_streams
), ),
http2_initial_connection_window_size=( http2_initial_connection_window_size=(
server_args.http2_initial_connection_window_size get_serving().http2_initial_connection_window_size
), ),
ssl_certfile=server_args.ssl_certfile, ssl_certfile=get_serving().ssl_certfile,
ssl_keyfile=server_args.ssl_keyfile, ssl_keyfile=get_serving().ssl_keyfile,
ssl_ca_certs=server_args.ssl_ca_certs, ssl_ca_certs=get_serving().ssl_ca_certs,
ssl_keyfile_password=server_args.ssl_keyfile_password, ssl_keyfile_password=get_serving().ssl_keyfile_password,
ssl_verify=False, # No MTLS supported for now. ssl_verify=False, # No MTLS supported for now.
) )
elif server_args.enable_ssl_refresh: elif get_serving().enable_ssl_refresh:
# Use Config/Server API for access to the SSLContext. # Use Config/Server API for access to the SSLContext.
config = uvicorn.Config( config = uvicorn.Config(
app, app,
host=get_serving().host, host=get_serving().host,
port=get_serving().port, port=get_serving().port,
root_path=server_args.fastapi_root_path, root_path=get_serving().fastapi_root_path,
log_level=get_observability().log_level_http log_level=get_observability().log_level_http
or get_observability().log_level, or get_observability().log_level,
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(), timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
loop="uvloop", loop="uvloop",
ssl_keyfile=server_args.ssl_keyfile, ssl_keyfile=get_serving().ssl_keyfile,
ssl_certfile=server_args.ssl_certfile, ssl_certfile=get_serving().ssl_certfile,
ssl_ca_certs=server_args.ssl_ca_certs, ssl_ca_certs=get_serving().ssl_ca_certs,
ssl_keyfile_password=server_args.ssl_keyfile_password, ssl_keyfile_password=get_serving().ssl_keyfile_password,
) )
config.load() # Creates the SSLContext config.load() # Creates the SSLContext
@@ -2660,9 +2660,9 @@ def _setup_and_run_http_server(
async def _run_with_ssl_refresh(): async def _run_with_ssl_refresh():
refresher = SSLCertRefresher( refresher = SSLCertRefresher(
config.ssl, config.ssl,
server_args.ssl_keyfile, get_serving().ssl_keyfile,
server_args.ssl_certfile, get_serving().ssl_certfile,
server_args.ssl_ca_certs, get_serving().ssl_ca_certs,
) )
logger.info("SSL certificate auto-refresh enabled.") logger.info("SSL certificate auto-refresh enabled.")
try: try:
@@ -2679,15 +2679,15 @@ def _setup_and_run_http_server(
app, app,
host=get_serving().host, host=get_serving().host,
port=get_serving().port, port=get_serving().port,
root_path=server_args.fastapi_root_path, root_path=get_serving().fastapi_root_path,
log_level=get_observability().log_level_http log_level=get_observability().log_level_http
or get_observability().log_level, or get_observability().log_level,
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(), timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
loop="uvloop", loop="uvloop",
ssl_keyfile=server_args.ssl_keyfile, ssl_keyfile=get_serving().ssl_keyfile,
ssl_certfile=server_args.ssl_certfile, ssl_certfile=get_serving().ssl_certfile,
ssl_ca_certs=server_args.ssl_ca_certs, ssl_ca_certs=get_serving().ssl_ca_certs,
ssl_keyfile_password=server_args.ssl_keyfile_password, ssl_keyfile_password=get_serving().ssl_keyfile_password,
) )
else: else:
# Multiple tokenizer and http processes # Multiple tokenizer and http processes
@@ -2699,14 +2699,14 @@ def _setup_and_run_http_server(
"propagate": False, "propagate": False,
} }
if server_args.enable_ssl_refresh: if get_serving().enable_ssl_refresh:
logger.warning( logger.warning(
"--enable-ssl-refresh is not supported with multiple " "--enable-ssl-refresh is not supported with multiple "
"tokenizer workers (--tokenizer-worker-num > 1). " "tokenizer workers (--tokenizer-worker-num > 1). "
"SSL refresh will be disabled." "SSL refresh will be disabled."
) )
if server_args.enable_http2: if get_serving().enable_http2:
logger.info( logger.info(
f"Starting embedded Granian HTTP/2 server on " f"Starting embedded Granian HTTP/2 server on "
f"{get_serving().host}:{get_serving().port}" f"{get_serving().host}:{get_serving().port}"
@@ -2717,36 +2717,36 @@ def _setup_and_run_http_server(
log_level=get_observability().log_level_http log_level=get_observability().log_level_http
or get_observability().log_level, or get_observability().log_level,
http2_max_concurrent_streams=( http2_max_concurrent_streams=(
server_args.http2_max_concurrent_streams get_serving().http2_max_concurrent_streams
), ),
http2_initial_connection_window_size=( http2_initial_connection_window_size=(
server_args.http2_initial_connection_window_size get_serving().http2_initial_connection_window_size
), ),
tokenizer_worker_num=server_args.tokenizer_worker_num, tokenizer_worker_num=get_serving().tokenizer_worker_num,
ssl_certfile=server_args.ssl_certfile, ssl_certfile=get_serving().ssl_certfile,
ssl_keyfile=server_args.ssl_keyfile, ssl_keyfile=get_serving().ssl_keyfile,
ssl_ca_certs=server_args.ssl_ca_certs, ssl_ca_certs=get_serving().ssl_ca_certs,
ssl_keyfile_password=server_args.ssl_keyfile_password, ssl_keyfile_password=get_serving().ssl_keyfile_password,
) )
else: else:
uvicorn.run( uvicorn.run(
"sglang.srt.entrypoints.http_server:app", "sglang.srt.entrypoints.http_server:app",
host=get_serving().host, host=get_serving().host,
port=get_serving().port, port=get_serving().port,
root_path=server_args.fastapi_root_path, root_path=get_serving().fastapi_root_path,
log_level=get_observability().log_level_http log_level=get_observability().log_level_http
or get_observability().log_level, or get_observability().log_level,
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(), timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
timeout_worker_healthcheck=envs.SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT.get(), timeout_worker_healthcheck=envs.SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT.get(),
loop="uvloop", loop="uvloop",
workers=server_args.tokenizer_worker_num, workers=get_serving().tokenizer_worker_num,
ssl_keyfile=server_args.ssl_keyfile, ssl_keyfile=get_serving().ssl_keyfile,
ssl_certfile=server_args.ssl_certfile, ssl_certfile=get_serving().ssl_certfile,
ssl_ca_certs=server_args.ssl_ca_certs, ssl_ca_certs=get_serving().ssl_ca_certs,
ssl_keyfile_password=server_args.ssl_keyfile_password, ssl_keyfile_password=get_serving().ssl_keyfile_password,
) )
finally: finally:
if server_args.tokenizer_worker_num > 1: if get_serving().tokenizer_worker_num > 1:
if multi_tokenizer_args_shm is not None: if multi_tokenizer_args_shm is not None:
multi_tokenizer_args_shm.unlink() multi_tokenizer_args_shm.unlink()
if _global_state is not None: if _global_state is not None:
@@ -5,6 +5,7 @@ from typing import List, Optional, Tuple
import requests import requests
import torch import torch
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.arg_groups.serving_hook import ssl_verify_of from sglang.srt.arg_groups.serving_hook import ssl_verify_of
from sglang.srt.entrypoints.EngineBase import EngineBase from sglang.srt.entrypoints.EngineBase import EngineBase
from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.entrypoints.http_server import launch_server
@@ -26,13 +27,16 @@ def launch_server_process(server_args: ServerArgs) -> multiprocessing.Process:
start_time = time.perf_counter() start_time = time.perf_counter()
ssl_verify = ssl_verify_of(server_args) ssl_verify = ssl_verify_of(server_args)
# The adapter's own configuration, not the bags: this runs in the parent,
# and the record it was handed is published only inside the server child.
cfg = resolving_view(server_args)
with requests.Session() as session: with requests.Session() as session:
while time.perf_counter() - start_time < timeout: while time.perf_counter() - start_time < timeout:
try: try:
headers = { headers = {
"Content-Type": "application/json; charset=utf-8", "Content-Type": "application/json; charset=utf-8",
"Authorization": f"Bearer {server_args.api_key}", "Authorization": f"Bearer {cfg.api_key}",
} }
response = session.get( response = session.get(
f"{base_url}/health_generate", headers=headers, verify=ssl_verify f"{base_url}/health_generate", headers=headers, verify=ssl_verify
@@ -60,9 +64,11 @@ class HttpServerEngineAdapter(EngineBase):
def __init__(self, **kwargs): def __init__(self, **kwargs):
self.server_args = ServerArgs(**kwargs) self.server_args = ServerArgs(**kwargs)
print( # This process launches the server as a child and never publishes, so
f"Launch HttpServerEngineAdapter at: {self.server_args.host}:{self.server_args.port}" # every read here is of the record it just built -- a bag read would
) # either fail closed or answer for an unrelated engine in the process.
cfg = resolving_view(self.server_args)
print(f"Launch HttpServerEngineAdapter at: {cfg.host}:{cfg.port}")
self.process = launch_server_process(self.server_args) self.process = launch_server_process(self.server_args)
def _make_request(self, endpoint: str, payload: Optional[dict] = None): def _make_request(self, endpoint: str, payload: Optional[dict] = None):
@@ -97,7 +103,7 @@ class HttpServerEngineAdapter(EngineBase):
{ {
"serialized_named_tensors": [ "serialized_named_tensors": [
MultiprocessingSerializer.serialize(named_tensors, output_str=True) MultiprocessingSerializer.serialize(named_tensors, output_str=True)
for _ in range(self.server_args.tp_size) for _ in range(resolving_view(self.server_args).tp_size)
], ],
"load_format": load_format, "load_format": load_format,
"flush_cache": flush_cache, "flush_cache": flush_cache,
@@ -17,6 +17,7 @@ from sglang.srt.entrypoints.openai.transcription_adapters.base import (
TranscriptionAdapter, TranscriptionAdapter,
) )
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_serving
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import random_uuid from sglang.srt.utils import random_uuid
@@ -85,7 +86,7 @@ async def handle_realtime_transcription(
websocket, websocket,
"too_many_sessions", "too_many_sessions",
f"Maximum concurrent sessions reached " f"Maximum concurrent sessions reached "
f"({server_args.asr_max_concurrent_sessions}).", f"({get_serving().asr_max_concurrent_sessions}).",
error_type="rate_limit_exceeded", error_type="rate_limit_exceeded",
) )
return return
@@ -72,6 +72,7 @@ from sglang.srt.entrypoints.openai.transcription_adapters.base import (
TranscriptionAdapter, TranscriptionAdapter,
) )
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_serving
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import random_uuid from sglang.srt.utils import random_uuid
@@ -186,7 +187,7 @@ class RealtimeConnection:
self.model_sample_rate = adapter.model_sample_rate self.model_sample_rate = adapter.model_sample_rate
self.bytes_per_second = self.model_sample_rate * _SAMPLE_WIDTH self.bytes_per_second = self.model_sample_rate * _SAMPLE_WIDTH
self.max_buffer_seconds = server_args.asr_max_buffer_seconds self.max_buffer_seconds = get_serving().asr_max_buffer_seconds
self.config = _SessionConfig() self.config = _SessionConfig()
@@ -14,6 +14,7 @@ from sglang.srt.entrypoints.openai.encoding_dsv32 import DS32EncodingError
from sglang.srt.entrypoints.openai.protocol import ErrorResponse, OpenAIServingRequest from sglang.srt.entrypoints.openai.protocol import ErrorResponse, OpenAIServingRequest
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
from sglang.srt.observability.req_time_stats import monotonic_time from sglang.srt.observability.req_time_stats import monotonic_time
from sglang.srt.runtime_context import get_observability
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -29,11 +30,9 @@ class OpenAIServingBase(ABC):
def __init__(self, tokenizer_manager: TokenizerManager): def __init__(self, tokenizer_manager: TokenizerManager):
self.tokenizer_manager = tokenizer_manager self.tokenizer_manager = tokenizer_manager
self.allowed_custom_labels = ( self.allowed_custom_labels = (
set( set(get_observability().tokenizer_metrics_allowed_custom_labels)
self.tokenizer_manager.server_args.tokenizer_metrics_allowed_custom_labels
)
if isinstance(self.tokenizer_manager.server_args, ServerArgs) if isinstance(self.tokenizer_manager.server_args, ServerArgs)
and self.tokenizer_manager.server_args.tokenizer_metrics_allowed_custom_labels and get_observability().tokenizer_metrics_allowed_custom_labels
else None else None
) )
@@ -230,14 +229,12 @@ class OpenAIServingBase(ABC):
def extract_custom_labels(self, raw_request): def extract_custom_labels(self, raw_request):
if ( if (
not self.allowed_custom_labels not self.allowed_custom_labels
or not self.tokenizer_manager.server_args.tokenizer_metrics_custom_labels_header or not get_observability().tokenizer_metrics_custom_labels_header
): ):
return None return None
custom_labels = None custom_labels = None
header = ( header = get_observability().tokenizer_metrics_custom_labels_header
self.tokenizer_manager.server_args.tokenizer_metrics_custom_labels_header
)
try: try:
raw_labels = ( raw_labels = (
orjson.loads(raw_request.headers.get(header)) orjson.loads(raw_request.headers.get(header))
@@ -10,6 +10,8 @@ from enum import Enum
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, Any, AsyncGenerator, Dict, List, Optional, Union from typing import TYPE_CHECKING, Any, AsyncGenerator, Dict, List, Optional, Union
from sglang.srt.runtime_context import get_model, get_serving
class ThinkingMode(str, Enum): class ThinkingMode(str, Enum):
"""Mode for message encoding - chat vs thinking/reasoning.""" """Mode for message encoding - chat vs thinking/reasoning."""
@@ -267,7 +269,7 @@ class OpenAIServingChat(OpenAIServingBase):
self.tool_call_parser = self.tokenizer_manager.config_value("tool_call_parser") self.tool_call_parser = self.tokenizer_manager.config_value("tool_call_parser")
self.reasoning_parser = self.tokenizer_manager.config_value("reasoning_parser") self.reasoning_parser = self.tokenizer_manager.config_value("reasoning_parser")
self.default_chat_template_kwargs = ( self.default_chat_template_kwargs = (
self.tokenizer_manager.server_args.default_chat_template_kwargs or {} get_serving().default_chat_template_kwargs or {}
) )
self._reasoning_detector = None self._reasoning_detector = None
if self.reasoning_parser: if self.reasoning_parser:
@@ -317,7 +319,7 @@ class OpenAIServingChat(OpenAIServingBase):
self._dsv4_reasoning_effort_profile = ( self._dsv4_reasoning_effort_profile = (
chat_encoding.resolve_dsv4_reasoning_effort_profile( chat_encoding.resolve_dsv4_reasoning_effort_profile(
model_path=self.tokenizer_manager.model_path, model_path=self.tokenizer_manager.model_path,
revision=self.tokenizer_manager.server_args.revision, revision=get_model().revision,
override=self.tokenizer_manager.model_config.hf_config.to_dict().get( override=self.tokenizer_manager.model_config.hf_config.to_dict().get(
chat_encoding.DSV4_REASONING_EFFORT_PROFILE_OVERRIDE chat_encoding.DSV4_REASONING_EFFORT_PROFILE_OVERRIDE
), ),
@@ -727,7 +729,7 @@ class OpenAIServingChat(OpenAIServingBase):
def _continuous_usage_cached_details( def _continuous_usage_cached_details(
self, content: Dict[str, Any] self, content: Dict[str, Any]
) -> Optional[PromptTokensDetails]: ) -> Optional[PromptTokensDetails]:
if not self.tokenizer_manager.server_args.enable_cache_report: if not get_serving().enable_cache_report:
return None return None
return UsageProcessor._details_if_cached( return UsageProcessor._details_if_cached(
content["meta_info"].get("cached_tokens", 0) content["meta_info"].get("cached_tokens", 0)
@@ -819,7 +821,7 @@ class OpenAIServingChat(OpenAIServingBase):
) -> AsyncGenerator[str, None]: ) -> AsyncGenerator[str, None]:
"""Generate SSE chunks for streaming content.""" """Generate SSE chunks for streaming content."""
offset = stream_offsets.get(index, 0) offset = stream_offsets.get(index, 0)
if self.tokenizer_manager.server_args.incremental_streaming_output: if get_serving().incremental_streaming_output:
delta = content["text"] delta = content["text"]
else: else:
delta = content["text"][offset:] delta = content["text"][offset:]
@@ -1007,12 +1009,12 @@ class OpenAIServingChat(OpenAIServingBase):
) )
max_output_tokens = request.max_completion_tokens or request.max_tokens max_output_tokens = request.max_completion_tokens or request.max_tokens
server_context_length = self.tokenizer_manager.server_args.context_length server_context_length = get_model().context_length
if ( if (
max_output_tokens max_output_tokens
and server_context_length and server_context_length
and max_output_tokens > server_context_length and max_output_tokens > server_context_length
) and not self.tokenizer_manager.server_args.allow_auto_truncate: ) and not get_serving().allow_auto_truncate:
return ( return (
f"max_completion_tokens is too large: {max_output_tokens}." f"max_completion_tokens is too large: {max_output_tokens}."
f"This model supports at most {server_context_length} completion tokens." f"This model supports at most {server_context_length} completion tokens."
@@ -1714,7 +1716,7 @@ class OpenAIServingChat(OpenAIServingBase):
try: try:
include_usage, continuous_usage_stats = should_include_usage( include_usage, continuous_usage_stats = should_include_usage(
request.stream_options, request.stream_options,
self.tokenizer_manager.server_args.stream_response_default_include_usage, get_serving().stream_response_default_include_usage,
) )
return_input_ids = self._should_return_input_ids(request) return_input_ids = self._should_return_input_ids(request)
@@ -1985,7 +1987,7 @@ class OpenAIServingChat(OpenAIServingBase):
completion_tokens, completion_tokens,
cached_tokens=cached_tokens, cached_tokens=cached_tokens,
n_choices=request.n, n_choices=request.n,
enable_cache_report=self.tokenizer_manager.server_args.enable_cache_report, enable_cache_report=get_serving().enable_cache_report,
image_tokens=total_image_tokens, image_tokens=total_image_tokens,
audio_tokens=total_audio_tokens, audio_tokens=total_audio_tokens,
video_tokens=total_video_tokens, video_tokens=total_video_tokens,
@@ -2209,7 +2211,7 @@ class OpenAIServingChat(OpenAIServingBase):
usage = UsageProcessor.calculate_response_usage( usage = UsageProcessor.calculate_response_usage(
ret, ret,
n_choices=request.n, n_choices=request.n,
enable_cache_report=self.tokenizer_manager.server_args.enable_cache_report, enable_cache_report=get_serving().enable_cache_report,
image_tokens=image_tokens, image_tokens=image_tokens,
audio_tokens=audio_tokens, audio_tokens=audio_tokens,
video_tokens=video_tokens, video_tokens=video_tokens,
@@ -2432,7 +2434,7 @@ class OpenAIServingChat(OpenAIServingBase):
"""Process logprobs for streaming response""" """Process logprobs for streaming response"""
output_token_logprobs = content["meta_info"]["output_token_logprobs"] output_token_logprobs = content["meta_info"]["output_token_logprobs"]
output_top_logprobs = content["meta_info"].get("output_top_logprobs", []) output_top_logprobs = content["meta_info"].get("output_top_logprobs", [])
if not self.tokenizer_manager.server_args.incremental_streaming_output: if not get_serving().incremental_streaming_output:
output_token_logprobs = output_token_logprobs[ output_token_logprobs = output_token_logprobs[
n_prev_token:total_output_logprobs n_prev_token:total_output_logprobs
] ]
@@ -34,6 +34,7 @@ from sglang.srt.managers.io_struct import GenerateReqInput
from sglang.srt.parser.code_completion_parser import ( from sglang.srt.parser.code_completion_parser import (
generate_completion_prompt_from_request, generate_completion_prompt_from_request,
) )
from sglang.srt.runtime_context import get_serving
from sglang.srt.utils.weight_versions import build_endpoint_weight_version_metadata from sglang.srt.utils.weight_versions import build_endpoint_weight_version_metadata
from sglang.utils import convert_json_schema_to_str from sglang.utils import convert_json_schema_to_str
@@ -246,7 +247,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
try: try:
include_usage, continuous_usage_stats = should_include_usage( include_usage, continuous_usage_stats = should_include_usage(
request.stream_options, request.stream_options,
self.tokenizer_manager.server_args.stream_response_default_include_usage, get_serving().stream_response_default_include_usage,
) )
async for content in self.tokenizer_manager.generate_request( async for content in self.tokenizer_manager.generate_request(
@@ -308,7 +309,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
output_top_logprobs = content["meta_info"].get( output_top_logprobs = content["meta_info"].get(
"output_top_logprobs", [] "output_top_logprobs", []
) )
if not self.tokenizer_manager.server_args.incremental_streaming_output: if not get_serving().incremental_streaming_output:
output_token_logprobs = output_token_logprobs[ output_token_logprobs = output_token_logprobs[
n_prev_token:total_output_logprobs n_prev_token:total_output_logprobs
] ]
@@ -327,7 +328,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
chunk_prompt_token_ids = None chunk_prompt_token_ids = None
if request.return_token_ids: if request.return_token_ids:
output_ids = content["output_ids"] output_ids = content["output_ids"]
if not self.tokenizer_manager.server_args.incremental_streaming_output: if not get_serving().incremental_streaming_output:
n_prev_token_id = n_prev_token_ids.get(index, 0) n_prev_token_id = n_prev_token_ids.get(index, 0)
chunk_token_ids = output_ids[n_prev_token_id:] chunk_token_ids = output_ids[n_prev_token_id:]
n_prev_token_ids[index] = len(output_ids) n_prev_token_ids[index] = len(output_ids)
@@ -337,7 +338,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
chunk_prompt_token_ids = content.get("prompt_token_ids") chunk_prompt_token_ids = content.get("prompt_token_ids")
# Generate delta # Generate delta
if self.tokenizer_manager.server_args.incremental_streaming_output: if get_serving().incremental_streaming_output:
delta = text delta = text
else: else:
delta = text[offset:] delta = text[offset:]
@@ -475,7 +476,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
completion_tokens, completion_tokens,
cached_tokens=cached_tokens, cached_tokens=cached_tokens,
n_choices=request.n, n_choices=request.n,
enable_cache_report=self.tokenizer_manager.server_args.enable_cache_report, enable_cache_report=get_serving().enable_cache_report,
) )
final_usage_chunk = CompletionStreamResponse( final_usage_chunk = CompletionStreamResponse(
id=content["meta_info"]["id"], id=content["meta_info"]["id"],
@@ -620,7 +621,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
choices.append(choice_data) choices.append(choice_data)
# Calculate usage # Calculate usage
cache_report = self.tokenizer_manager.server_args.enable_cache_report cache_report = get_serving().enable_cache_report
usage = UsageProcessor.calculate_response_usage( usage = UsageProcessor.calculate_response_usage(
ret, n_choices=request.n, enable_cache_report=cache_report ret, n_choices=request.n, enable_cache_report=cache_report
) )
@@ -73,6 +73,7 @@ from sglang.srt.function_call.function_call_parser import FunctionCallParser
from sglang.srt.function_call.json_array_parser import JsonArrayParser from sglang.srt.function_call.json_array_parser import JsonArrayParser
from sglang.srt.managers.io_struct import GenerateReqInput from sglang.srt.managers.io_struct import GenerateReqInput
from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.runtime_context import get_serving
from sglang.srt.sampling.sampling_params import ( from sglang.srt.sampling.sampling_params import (
set_request_reasoning_end_token_ids, set_request_reasoning_end_token_ids,
) )
@@ -2031,7 +2032,7 @@ class OpenAIServingResponses(OpenAIServingChat):
finish_reason: Optional[dict[str, Any]] = None finish_reason: Optional[dict[str, Any]] = None
flushed = False flushed = False
stream_offset = 0 stream_offset = 0
incremental = self.tokenizer_manager.server_args.incremental_streaming_output incremental = get_serving().incremental_streaming_output
def _open_reasoning_item() -> str: def _open_reasoning_item() -> str:
nonlocal current_output_index nonlocal current_output_index
@@ -55,6 +55,7 @@ from sglang.srt.entrypoints.openai.streaming_asr import (
) )
from sglang.srt.entrypoints.openai.transcription_adapters import resolve_adapter from sglang.srt.entrypoints.openai.transcription_adapters import resolve_adapter
from sglang.srt.managers.io_struct import GenerateReqInput from sglang.srt.managers.io_struct import GenerateReqInput
from sglang.srt.runtime_context import get_serving
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
@@ -74,7 +75,7 @@ class OpenAIServingTranscription(OpenAIServingBase):
# Cap concurrent /v1/realtime sessions. The Semaphore is bound to the # Cap concurrent /v1/realtime sessions. The Semaphore is bound to the
# event loop on first acquire (uvicorn's loop in normal serving). # event loop on first acquire (uvicorn's loop in normal serving).
self._session_semaphore = asyncio.Semaphore( self._session_semaphore = asyncio.Semaphore(
tokenizer_manager.server_args.asr_max_concurrent_sessions get_serving().asr_max_concurrent_sessions
) )
def _request_id_prefix(self) -> str: def _request_id_prefix(self) -> str:
@@ -498,11 +499,7 @@ class OpenAIServingTranscription(OpenAIServingBase):
# the cumulative text. Always reconstruct cumulative text locally # the cumulative text. Always reconstruct cumulative text locally
# so the rest of the loop (prefix parse + visible-buffer slice) # so the rest of the loop (prefix parse + visible-buffer slice)
# works uniformly under either mode. # works uniformly under either mode.
incremental = getattr( incremental = get_serving().incremental_streaming_output
self.tokenizer_manager.server_args,
"incremental_streaming_output",
False,
)
cumulative_text = "" cumulative_text = ""
try: try:
@@ -613,11 +610,7 @@ class OpenAIServingTranscription(OpenAIServingBase):
model = request.model model = request.model
fused_mode = getattr(request, "_fused_autodetect", False) fused_mode = getattr(request, "_fused_autodetect", False)
ts_variant = getattr(request, "_fused_ts_variant", False) ts_variant = getattr(request, "_fused_ts_variant", False)
incremental = getattr( incremental = get_serving().incremental_streaming_output
self.tokenizer_manager.server_args,
"incremental_streaming_output",
False,
)
def _frame(delta: Optional[str], finish_reason: Optional[str] = None) -> str: def _frame(delta: Optional[str], finish_reason: Optional[str] = None) -> str:
chunk = TranscriptionStreamResponse( chunk = TranscriptionStreamResponse(
+4 -3
View File
@@ -8,6 +8,7 @@ from sglang.kernels.ops.kv_canary.consts import (
RealKvHashMode, RealKvHashMode,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_observability
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -61,18 +62,18 @@ class CanaryConfig:
@classmethod @classmethod
def from_env(cls, server_args: ServerArgs) -> CanaryConfig: def from_env(cls, server_args: ServerArgs) -> CanaryConfig:
mode_raw = server_args.kv_canary.strip().lower() mode_raw = get_observability().kv_canary.strip().lower()
if mode_raw not in ("none", "log", "raise"): if mode_raw not in ("none", "log", "raise"):
raise ValueError( raise ValueError(
f"kv-canary: kv_canary must be one of none/log/raise, got {mode_raw!r}" f"kv-canary: kv_canary must be one of none/log/raise, got {mode_raw!r}"
) )
real_kv_raw = server_args.kv_canary_real_data.strip().upper() real_kv_raw = get_observability().kv_canary_real_data.strip().upper()
return cls( return cls(
mode=CanaryMode(mode_raw), mode=CanaryMode(mode_raw),
ring_capacity=envs.SGLANG_KV_CANARY_RING_CAPACITY.get(), ring_capacity=envs.SGLANG_KV_CANARY_RING_CAPACITY.get(),
sweep_interval=server_args.kv_canary_sweep_interval, sweep_interval=get_observability().kv_canary_sweep_interval,
real_kv_hash_mode=RealKvHashMode[real_kv_raw], real_kv_hash_mode=RealKvHashMode[real_kv_raw],
enable_write_input_assert=envs.SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT.get(), enable_write_input_assert=envs.SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT.get(),
enable_verify_token_assert=envs.SGLANG_KV_CANARY_ENABLE_VERIFY_TOKEN_ASSERT.get(), enable_verify_token_assert=envs.SGLANG_KV_CANARY_ENABLE_VERIFY_TOKEN_ASSERT.get(),
@@ -37,6 +37,7 @@ from sglang.srt.layers.attention.dsv4.metadata import (
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec,
get_parallel, get_parallel,
get_spec, get_spec,
) )
@@ -474,7 +475,7 @@ class DeepseekV4HipRadixBackend(
model_runner.model_config.hf_text_config, "index_topk", C4_TOPK model_runner.model_config.hf_text_config, "index_topk", C4_TOPK
) )
self.enable_deepseek_v4_fp4_indexer: bool = ( self.enable_deepseek_v4_fp4_indexer: bool = (
model_runner.server_args.enable_deepseek_v4_fp4_indexer get_exec().kernel.enable_deepseek_v4_fp4_indexer
) )
self.topk = get_spec().speculative_eagle_topk or 0 self.topk = get_spec().speculative_eagle_topk or 0
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4" assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
@@ -314,7 +314,7 @@ class FlashInferAttnBackend(AttentionBackend):
model_runner model_runner
) )
self.use_sliding_window_kv_pool = self._swa_kv_pool is not None self.use_sliding_window_kv_pool = self._swa_kv_pool is not None
self.enable_mis = model_runner.server_args.enable_mis self.enable_mis = get_exec().features.enable_mis
# FIXME: remove dllm workarounds from flashinfer # FIXME: remove dllm workarounds from flashinfer
self.dllm_config = DllmConfig.from_server_args(model_runner.server_args) self.dllm_config = DllmConfig.from_server_args(model_runner.server_args)
@@ -64,7 +64,7 @@ class MambaAttnBackendBase(AttentionBackend):
self.is_draft_worker = model_runner.is_draft_worker self.is_draft_worker = model_runner.is_draft_worker
self.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_pool self.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
self.enable_unified_memory = model_runner.server_args.enable_unified_memory self.enable_unified_memory = get_memory().enable_unified_memory
# model_config must not be touched here: backend selection reads the # model_config must not be touched here: backend selection reads the
# linear_attn_backends stamp first, and that guard test constructs # linear_attn_backends stamp first, and that guard test constructs
# backends on runners without a real model_config. # backends on runners without a real model_config.
@@ -933,10 +933,8 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
assert self.conv_states_shape[-1] < self.mamba_chunk_size, ( assert self.conv_states_shape[-1] < self.mamba_chunk_size, (
f"{self.conv_states_shape[-1]=} should be less than {self.mamba_chunk_size}" f"{self.conv_states_shape[-1]=} should be less than {self.mamba_chunk_size}"
) )
assert ( assert get_exec().mamba.mamba_track_interval >= self.mamba_chunk_size, (
model_runner.server_args.mamba_track_interval >= self.mamba_chunk_size f"mamba_track_interval ({get_exec().mamba.mamba_track_interval}) must be >= mamba_chunk_size ({self.mamba_chunk_size})"
), (
f"mamba_track_interval ({model_runner.server_args.mamba_track_interval}) must be >= mamba_chunk_size ({self.mamba_chunk_size})"
) )
def init_forward_metadata_out_graph( def init_forward_metadata_out_graph(
@@ -18,6 +18,7 @@ from sglang.srt.layers.attention.minicpm.attention_adapter import (
from sglang.srt.layers.attention.minicpm.cache import attach_compressed_cache from sglang.srt.layers.attention.minicpm.cache import attach_compressed_cache
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec,
get_parallel, get_parallel,
get_platform, get_platform,
get_schedule, get_schedule,
@@ -189,7 +190,7 @@ class MiniCPMSparseBackend(AttentionBackend):
model_runner.token_to_kv_pool_allocator, model_runner.token_to_kv_pool_allocator,
kernel_size=self.kernel_size, kernel_size=self.kernel_size,
kernel_stride=self.kernel_stride, kernel_stride=self.kernel_stride,
enable_memory_saver=model_runner.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.enable_memory_saver,
) )
self.req_to_sparse_k1_token = self.req_to_token_pool.req_to_sparse_k1_token self.req_to_sparse_k1_token = self.req_to_token_pool.req_to_sparse_k1_token
self.req_to_sparse_k2_token = self.req_to_token_pool.req_to_sparse_k2_token self.req_to_sparse_k2_token = self.req_to_token_pool.req_to_sparse_k2_token
@@ -175,7 +175,7 @@ class TritonAttnBackend(AttentionBackend):
self.decode_attention_fwd = torch.compiler.disable(decode_attention_fwd) self.decode_attention_fwd = torch.compiler.disable(decode_attention_fwd)
# Work-Centric (Lean) Attention activation. None => auto-gate from host-side # Work-Centric (Lean) Attention activation. None => auto-gate from host-side
# seqlen metadata in forward_decode; True/False => explicit override. # seqlen metadata in forward_decode; True/False => explicit override.
self.enable_lean_attention = model_runner.server_args.enable_lean_attention self.enable_lean_attention = get_exec().kernel.enable_lean_attention
self._lean_decode_seqlen_gate = lean_decode_seqlen_gate self._lean_decode_seqlen_gate = lean_decode_seqlen_gate
self._lean_capture_policy = lean_capture_policy self._lean_capture_policy = lean_capture_policy
self.extend_attention_fwd = torch.compiler.disable(extend_attention_fwd) self.extend_attention_fwd = torch.compiler.disable(extend_attention_fwd)
@@ -333,9 +333,7 @@ class TritonAttnBackend(AttentionBackend):
) )
self.static_kv_splits = False self.static_kv_splits = False
else: else:
self.split_tile_size = ( self.split_tile_size = get_exec().kernel.triton_attention_split_tile_size
model_runner.server_args.triton_attention_split_tile_size
)
if self.split_tile_size is not None: if self.split_tile_size is not None:
self.max_kv_splits = ( self.max_kv_splits = (
+7 -3
View File
@@ -40,7 +40,11 @@ from typing import Optional
import torch import torch
from sglang.srt.distributed import get_tp_group from sglang.srt.distributed import get_tp_group
from sglang.srt.runtime_context import get_flags, get_forward from sglang.srt.runtime_context import (
get_flags,
get_forward,
get_parallel,
)
from sglang.srt.utils.common import ceil_align from sglang.srt.utils.common import ceil_align
# Architectures whose decoder layers route attention/MLP through # Architectures whose decoder layers route attention/MLP through
@@ -50,12 +54,12 @@ from sglang.srt.utils.common import ceil_align
SP_SUPPORTED_ARCHITECTURES = frozenset({"Qwen3ForCausalLM"}) SP_SUPPORTED_ARCHITECTURES = frozenset({"Qwen3ForCausalLM"})
def initialize_layernorm_sp(*, server_args, model_config) -> None: def initialize_layernorm_sp(*, model_config) -> None:
"""Materialize ``flags.sp.enabled``; runs once per worker after distributed """Materialize ``flags.sp.enabled``; runs once per worker after distributed
setup, alongside ``initialize_dp_attention``.""" setup, alongside ``initialize_dp_attention``."""
architectures = model_config.hf_config.architectures architectures = model_config.hf_config.architectures
get_flags().sp.enabled = bool( get_flags().sp.enabled = bool(
server_args.enable_layernorm_sp get_parallel().enable_layernorm_sp
and architectures and architectures
and architectures[0] in SP_SUPPORTED_ARCHITECTURES and architectures[0] in SP_SUPPORTED_ARCHITECTURES
) )
@@ -33,7 +33,7 @@ _EXPERT_WEIGHT_NAMES = (
class DwdpManager: class DwdpManager:
def __init__(self, server_args: ServerArgs): def __init__(self, server_args: ServerArgs):
self.dwdp_size = server_args.dwdp_size self.dwdp_size = get_parallel().dwdp_size
self.dwdp_rank = get_parallel().tp_rank self.dwdp_rank = get_parallel().tp_rank
self.device_id = torch.cuda.current_device() self.device_id = torch.cuda.current_device()
self.layout: Optional[DwdpExpertLayout] = None self.layout: Optional[DwdpExpertLayout] = None
@@ -8,7 +8,7 @@ from typing import Optional
import torch import torch
from sglang.srt.arg_groups.overrides import cutedsl_moe_max_num_tokens, resolving_view from sglang.srt.arg_groups.overrides import cutedsl_moe_max_num_tokens
from sglang.srt.layers.communicator import ( from sglang.srt.layers.communicator import (
CommunicateWithAllReduceAndLayerNormFn, CommunicateWithAllReduceAndLayerNormFn,
LayerCommunicator, LayerCommunicator,
@@ -18,7 +18,11 @@ from sglang.srt.layers.communicator import (
from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.dp_attention import is_dp_attention_enabled
from sglang.srt.layers.moe import get_moe_a2a_backend from sglang.srt.layers.moe import get_moe_a2a_backend
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_parallel,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -33,9 +37,8 @@ def is_supported_forward_mode(forward_mode: ForwardMode) -> bool:
def resolve_max_m(model_runner) -> int: def resolve_max_m(model_runner) -> int:
"""Use framework token bounds as the workspace-capacity source of truth.""" """Use framework token bounds as the workspace-capacity source of truth."""
server_args = resolving_view(model_runner.server_args) decode_config = get_exec().graph.cuda_graph_config.decode
decode_config = server_args.cuda_graph_config.decode prefill_config = get_exec().graph.cuda_graph_config.prefill
prefill_config = server_args.cuda_graph_config.prefill
candidates = [ candidates = [
cutedsl_moe_max_num_tokens(model_runner.server_args), cutedsl_moe_max_num_tokens(model_runner.server_args),
model_runner.max_running_requests, model_runner.max_running_requests,
@@ -335,7 +338,7 @@ def prepare_qwen35_flashinfer_fusion(model, model_runner) -> None:
service = getattr(model, "flashinfer_mnnvl_cutedsl_fusion", None) service = getattr(model, "flashinfer_mnnvl_cutedsl_fusion", None)
if service is None: if service is None:
return return
if model_runner.server_args.enable_pdmux: if get_disagg().enable_pdmux:
raise RuntimeError( raise RuntimeError(
"FlashInfer MNNVL CuTe DSL fusion does not support concurrent PDMux " "FlashInfer MNNVL CuTe DSL fusion does not support concurrent PDMux "
"streams sharing one mutable workspace" "streams sharing one mutable workspace"
@@ -56,7 +56,6 @@ if TYPE_CHECKING:
DispatchOutput, DispatchOutput,
StandardDispatchOutput, StandardDispatchOutput,
) )
from sglang.srt.server_args import ServerArgs
from sglang.srt.hardware_backend.npu.quantization.moe_methods import ( from sglang.srt.hardware_backend.npu.quantization.moe_methods import (
NPUUnquantMoEMethod, NPUUnquantMoEMethod,
@@ -152,7 +151,7 @@ def should_enable_bf16_splitk_gemm(backend: Bf16GemmBackend) -> bool:
return backend.is_optimized() and envs.SGLANG_ENABLE_BF16_SPLITK_GEMM.get() return backend.is_optimized() and envs.SGLANG_ENABLE_BF16_SPLITK_GEMM.get()
def initialize_bf16_gemm_config(server_args: ServerArgs) -> None: def initialize_bf16_gemm_config() -> None:
global _BF16_GEMM_BACKEND global _BF16_GEMM_BACKEND
global _cutedsl_bf16_gemm, _use_cutedsl_bf16_gemm global _cutedsl_bf16_gemm, _use_cutedsl_bf16_gemm
global _flashinfer_pr4266_splitk_tactic global _flashinfer_pr4266_splitk_tactic
@@ -162,7 +161,7 @@ def initialize_bf16_gemm_config(server_args: ServerArgs) -> None:
global _flashinfer_pr4266_run_direct_dense global _flashinfer_pr4266_run_direct_dense
global _enable_bf16_splitk_gemm global _enable_bf16_splitk_gemm
backend_str = server_args.bf16_gemm_backend backend_str = get_exec().kernel.bf16_gemm_backend
if backend_str == "auto" and get_platform().is_sm100: if backend_str == "auto" and get_platform().is_sm100:
backend_str = ( backend_str = (
"torch" "torch"
@@ -17,6 +17,7 @@ from sglang.srt.lora.utils import (
merge_and_chunk_segments, merge_and_chunk_segments,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import get_lora
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
MIN_CHUNK_SIZE = 16 MIN_CHUNK_SIZE = 16
@@ -42,7 +43,7 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend):
server_args: ServerArgs, server_args: ServerArgs,
): ):
super().__init__(max_loras_per_batch, device) super().__init__(max_loras_per_batch, device)
self.max_chunk_size = server_args.max_lora_chunk_size self.max_chunk_size = get_lora().max_lora_chunk_size
def run_lora_a_embedding( def run_lora_a_embedding(
self, self,
@@ -19,8 +19,8 @@ import multiprocessing as mp
import signal import signal
import threading import threading
import time import time
from collections.abc import Callable
from enum import Enum, auto from enum import Enum, auto
from typing import Callable, List, Optional
import psutil import psutil
import setproctitle import setproctitle
@@ -53,7 +53,9 @@ from sglang.srt.runtime_context import (
get_device, get_device,
get_disagg, get_disagg,
get_exec, get_exec,
get_observability,
get_parallel, get_parallel,
get_serving,
publish, publish,
) )
from sglang.srt.server_args import ( from sglang.srt.server_args import (
@@ -154,7 +156,7 @@ class DataParallelController:
# Init inter-process communication # Init inter-process communication
self.context = zmq.Context(1 + get_parallel().dp_size) self.context = zmq.Context(1 + get_parallel().dp_size)
if server_args.node_rank == 0: if get_parallel().node_rank == 0:
self.recv_from_tokenizer = get_zmq_socket( self.recv_from_tokenizer = get_zmq_socket(
self.context, zmq.PULL, port_args.scheduler_input_ipc_name, False self.context, zmq.PULL, port_args.scheduler_input_ipc_name, False
) )
@@ -174,13 +176,13 @@ class DataParallelController:
) )
self.launch_dp_size: int = get_parallel().dp_size self.launch_dp_size: int = get_parallel().dp_size
self.max_dp_size: int = server_args.max_ep_size or get_parallel().dp_size self.max_dp_size: int = get_parallel().max_ep_size or get_parallel().dp_size
assert self.max_dp_size >= self.launch_dp_size, ( assert self.max_dp_size >= self.launch_dp_size, (
f"--max-ep-size ({self.max_dp_size}) must be >= " f"--max-ep-size ({self.max_dp_size}) must be >= "
f"--dp ({self.launch_dp_size})." f"--dp ({self.launch_dp_size})."
) )
self.dp_active: List[bool] = [True] * self.launch_dp_size + [False] * ( self.dp_active: list[bool] = [True] * self.launch_dp_size + [False] * (
self.max_dp_size - self.launch_dp_size self.max_dp_size - self.launch_dp_size
) )
@@ -196,9 +198,9 @@ class DataParallelController:
# Launch data parallel workers # Launch data parallel workers
self.scheduler_procs = [] self.scheduler_procs = []
self.workers: List[Optional[zmq.Socket]] = [None] * self.max_dp_size self.workers: list[zmq.Socket | None] = [None] * self.max_dp_size
self.status: List[bool] = list(self.dp_active) self.status: list[bool] = list(self.dp_active)
self._active_workers: List[int] = list(range(self.launch_dp_size)) self._active_workers: list[int] = list(range(self.launch_dp_size))
self._active_count_cache: int = self.launch_dp_size self._active_count_cache: int = self.launch_dp_size
if get_parallel().enable_dp_attention: if get_parallel().enable_dp_attention:
@@ -209,7 +211,7 @@ class DataParallelController:
# Otherwise fall back to the original behaviour: send to only the # Otherwise fall back to the original behaviour: send to only the
# first leader, which then broadcasts over the full tp_group. # first leader, which then broadcasts over the full tp_group.
local_ctrl = get_parallel().enable_dp_attention_local_control_broadcast local_ctrl = get_parallel().enable_dp_attention_local_control_broadcast
self.control_message_step = 1 if local_ctrl else server_args.tp_size self.control_message_step = 1 if local_ctrl else get_parallel().tp_size
else: else:
self.launch_dp_schedulers(server_args, port_args) self.launch_dp_schedulers(server_args, port_args)
self.control_message_step = 1 self.control_message_step = 1
@@ -223,7 +225,7 @@ class DataParallelController:
test_stuck_time=envs.SGLANG_TEST_STUCK_DP_CONTROLLER.get(), test_stuck_time=envs.SGLANG_TEST_STUCK_DP_CONTROLLER.get(),
) )
if server_args.enable_metrics: if get_observability().enable_metrics:
start_cpu_monitor_thread("data_parallel_controller") start_cpu_monitor_thread("data_parallel_controller")
def send_to_all_workers(self, obj): def send_to_all_workers(self, obj):
@@ -392,10 +394,12 @@ class DataParallelController:
) )
threads.append(thread) threads.append(thread)
base_gpu_id += ( base_gpu_id += (
server_args.tp_size * get_parallel().pp_size * server_args.gpu_id_step get_parallel().tp_size
* get_parallel().pp_size
* get_device().gpu_id_step
) )
if server_args.node_rank == 0: if get_parallel().node_rank == 0:
self.workers[dp_rank] = get_zmq_socket( self.workers[dp_rank] = get_zmq_socket(
self.context, self.context,
zmq.PUSH, zmq.PUSH,
@@ -430,8 +434,8 @@ class DataParallelController:
time.sleep(30 * 24 * 3600) time.sleep(30 * 24 * 3600)
def _broadcast_worker_ports( def _broadcast_worker_ports(
self, server_args: ServerArgs, worker_ports: Optional[List[int]] = None self, server_args: ServerArgs, worker_ports: list[int] | None = None
) -> List[int]: ) -> list[int]:
"""Broadcast worker ports from node 0 to all other nodes. """Broadcast worker ports from node 0 to all other nodes.
Node 0 acts as the server, waiting for all other nodes to connect and Node 0 acts as the server, waiting for all other nodes to connect and
@@ -446,28 +450,28 @@ class DataParallelController:
List of worker ports (same on all nodes after broadcast). List of worker ports (same on all nodes after broadcast).
""" """
is_joiner = server_args.is_ep_scale_joiner is_joiner = server_args.is_ep_scale_joiner
if server_args.dist_init_addr is None or is_joiner: if get_parallel().dist_init_addr is None or is_joiner:
na = NetworkAddress( na = NetworkAddress(
server_args.host or "127.0.0.1", get_serving().host or "127.0.0.1",
server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA, get_serving().port + DP_ATTENTION_HANDSHAKE_PORT_DELTA,
) )
else: else:
na = NetworkAddress.parse(server_args.dist_init_addr) na = NetworkAddress.parse(get_parallel().dist_init_addr)
na = NetworkAddress(na.host, na.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA) na = NetworkAddress(na.host, na.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA)
endpoint = na.to_tcp() endpoint = na.to_tcp()
if server_args.node_rank == 0: if get_parallel().node_rank == 0:
# Node 0: Broadcast worker ports to all other nodes # Node 0: Broadcast worker ports to all other nodes
return self._broadcast_ports_as_server( return self._broadcast_ports_as_server(
endpoint, server_args.nnodes - 1, worker_ports endpoint, get_parallel().nnodes - 1, worker_ports
) )
else: else:
# Other nodes: Receive worker ports from node 0 # Other nodes: Receive worker ports from node 0
return self._receive_ports_as_client(endpoint, server_args.node_rank) return self._receive_ports_as_client(endpoint, get_parallel().node_rank)
def _broadcast_ports_as_server( def _broadcast_ports_as_server(
self, endpoint: str, expected_clients: int, worker_ports: List[int] self, endpoint: str, expected_clients: int, worker_ports: list[int]
) -> List[int]: ) -> list[int]:
"""Broadcast worker ports to all client nodes.""" """Broadcast worker ports to all client nodes."""
logger.debug(f"Broadcasting worker ports to {expected_clients} client nodes") logger.debug(f"Broadcasting worker ports to {expected_clients} client nodes")
logger.debug(f"Worker ports: {worker_ports}") logger.debug(f"Worker ports: {worker_ports}")
@@ -500,7 +504,7 @@ class DataParallelController:
daemon=True, daemon=True,
).start() ).start()
def _reply_ports_as_server(self, rep_socket: zmq.Socket, worker_ports: List[int]): def _reply_ports_as_server(self, rep_socket: zmq.Socket, worker_ports: list[int]):
"""Background thread: serve the pre-bound worker-port list to """Background thread: serve the pre-bound worker-port list to
late-arriving elastic joiners. Publishes port numbers only; the primary late-arriving elastic joiners. Publishes port numbers only; the primary
keeps ownership of every socket.""" keeps ownership of every socket."""
@@ -518,9 +522,9 @@ class DataParallelController:
sock_send(rep_socket, wrap_as_pickle(worker_ports)) sock_send(rep_socket, wrap_as_pickle(worker_ports))
logger.debug(f"Sent worker ports to node {client_rank}") logger.debug(f"Sent worker ports to node {client_rank}")
def _receive_ports_as_client(self, endpoint: str, node_rank: int) -> List[int]: def _receive_ports_as_client(self, endpoint: str, node_rank: int) -> list[int]:
"""Receive worker ports from the server node.""" """Receive worker ports from the server node."""
logger.debug(f"Connecting to node 0 to receive worker ports") logger.debug("Connecting to node 0 to receive worker ports")
req_socket = get_zmq_socket(self.context, zmq.REQ, endpoint, False) req_socket = get_zmq_socket(self.context, zmq.REQ, endpoint, False)
req_socket.setsockopt(zmq.RCVTIMEO, 600 * 1000) # 10 minute timeout req_socket.setsockopt(zmq.RCVTIMEO, 600 * 1000) # 10 minute timeout
@@ -543,37 +547,37 @@ class DataParallelController:
req_socket.close() req_socket.close()
def _joiner_local_tp_span(self, server_args: ServerArgs) -> int: def _joiner_local_tp_span(self, server_args: ServerArgs) -> int:
return server_args.tp_size return get_parallel().tp_size
def _joiner_slot_offset(self, server_args: ServerArgs) -> int: def _joiner_slot_offset(self, server_args: ServerArgs) -> int:
return server_args.ep_join_rank_offset return get_parallel().ep_join_rank_offset
def launch_dp_attention_schedulers( def launch_dp_attention_schedulers(
self, server_args: ServerArgs, port_args: PortArgs self, server_args: ServerArgs, port_args: PortArgs
): ):
if server_args.dist_init_addr is None: if get_parallel().dist_init_addr is None:
bind_host = "127.0.0.1" bind_host = "127.0.0.1"
else: else:
bind_host = NetworkAddress.parse(server_args.dist_init_addr).host bind_host = NetworkAddress.parse(get_parallel().dist_init_addr).host
worker_ports = [] worker_ports = []
if server_args.is_ep_scale_joiner: if server_args.is_ep_scale_joiner:
# Scale joiners connect to their pre-bound primary worker sockets. # Scale joiners connect to their pre-bound primary worker sockets.
primary = NetworkAddress.parse(server_args.dist_init_addr) primary = NetworkAddress.parse(get_parallel().dist_init_addr)
primary_endpoint = NetworkAddress( primary_endpoint = NetworkAddress(
primary.host, primary.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA primary.host, primary.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA
).to_tcp() ).to_tcp()
all_ports = self._receive_ports_as_client( all_ports = self._receive_ports_as_client(
primary_endpoint, server_args.node_rank primary_endpoint, get_parallel().node_rank
) )
offset = self._joiner_slot_offset(server_args) offset = self._joiner_slot_offset(server_args)
local_tp_span = self._joiner_local_tp_span(server_args) local_tp_span = self._joiner_local_tp_span(server_args)
broadcasted_ports = all_ports[offset : offset + local_tp_span] broadcasted_ports = all_ports[offset : offset + local_tp_span]
elif server_args.node_rank == 0: elif get_parallel().node_rank == 0:
# Elastic primaries reserve sockets for the maximum DP size. # Elastic primaries reserve sockets for the maximum DP size.
bind_count = ( bind_count = (
self.max_dp_size self.max_dp_size
if server_args.elastic_ep_backend is not None if get_exec().moe.elastic_ep_backend is not None
else get_parallel().dp_size else get_parallel().dp_size
) )
for slot in range(bind_count): for slot in range(bind_count):
@@ -601,35 +605,35 @@ class DataParallelController:
server_args: ServerArgs, server_args: ServerArgs,
port_args: PortArgs, port_args: PortArgs,
base_gpu_id: int, base_gpu_id: int,
dp_rank: Optional[int], dp_rank: int | None,
worker_ports: Optional[List[int]] = None, worker_ports: list[int] | None = None,
): ):
if not get_parallel().enable_dp_attention: if not get_parallel().enable_dp_attention:
logger.info(f"Launch DP{dp_rank} starting at GPU #{base_gpu_id}.") logger.info(f"Launch DP{dp_rank} starting at GPU #{base_gpu_id}.")
memory_saver_adapter = TorchMemorySaverAdapter.create( memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=server_args.enable_memory_saver enable=get_exec().features.enable_memory_saver
) )
scheduler_pipe_readers = [] scheduler_pipe_readers = []
pp_size_per_node = max(get_parallel().pp_size // server_args.nnodes, 1) pp_size_per_node = max(get_parallel().pp_size // get_parallel().nnodes, 1)
nnodes_per_pp_rank = max(server_args.nnodes // get_parallel().pp_size, 1) nnodes_per_pp_rank = max(get_parallel().nnodes // get_parallel().pp_size, 1)
pp_rank_range = range( pp_rank_range = range(
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank), pp_size_per_node * (get_parallel().node_rank // nnodes_per_pp_rank),
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank + 1), pp_size_per_node * (get_parallel().node_rank // nnodes_per_pp_rank + 1),
) )
nnodes_per_tp_group = nnodes_per_pp_rank nnodes_per_tp_group = nnodes_per_pp_rank
tp_size_per_node = server_args.tp_size // nnodes_per_tp_group tp_size_per_node = get_parallel().tp_size // nnodes_per_tp_group
if server_args.is_ep_scale_joiner: if server_args.is_ep_scale_joiner:
# Scale joiners enumerate their full local TP span. # Scale joiners enumerate their full local TP span.
tp_rank_range = range(server_args.tp_size) tp_rank_range = range(get_parallel().tp_size)
tp_size_per_node = server_args.tp_size tp_size_per_node = get_parallel().tp_size
else: else:
tp_rank_range = range( tp_rank_range = range(
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group), tp_size_per_node * (get_parallel().node_rank % nnodes_per_tp_group),
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1), tp_size_per_node * (get_parallel().node_rank % nnodes_per_tp_group + 1),
) )
attn_cp_rank = 0 attn_cp_rank = 0
@@ -643,7 +647,7 @@ class DataParallelController:
_, _, dp_rank, _ = compute_dp_attention_world_info( _, _, dp_rank, _ = compute_dp_attention_world_info(
get_parallel().enable_dp_attention, get_parallel().enable_dp_attention,
tp_rank, tp_rank,
server_args.tp_size, get_parallel().tp_size,
get_parallel().dp_size, get_parallel().dp_size,
get_parallel().attn_cp_size, get_parallel().attn_cp_size,
) )
@@ -653,7 +657,9 @@ class DataParallelController:
) )
if server_args.is_ep_scale_joiner: if server_args.is_ep_scale_joiner:
# Scale-joiner outputs return through the primary tokenizer. # Scale-joiner outputs return through the primary tokenizer.
primary_addr = NetworkAddress.parse(server_args.dist_init_addr) primary_addr = NetworkAddress.parse(
get_parallel().dist_init_addr
)
primary_port_base = primary_addr.port + 1 primary_port_base = primary_addr.port + 1
rank_port_args.tokenizer_ipc_name = NetworkAddress( rank_port_args.tokenizer_ipc_name = NetworkAddress(
primary_addr.host, primary_port_base primary_addr.host, primary_port_base
@@ -668,10 +674,10 @@ class DataParallelController:
reader, writer = mp.Pipe(duplex=False) reader, writer = mp.Pipe(duplex=False)
gpu_id = ( gpu_id = (
server_args.base_gpu_id get_device().base_gpu_id
+ base_gpu_id + base_gpu_id
+ ((pp_rank % pp_size_per_node) * tp_size_per_node) + ((pp_rank % pp_size_per_node) * tp_size_per_node)
+ (tp_rank % tp_size_per_node) * server_args.gpu_id_step + (tp_rank % tp_size_per_node) * get_device().gpu_id_step
) )
attn_dp_size = ( attn_dp_size = (
get_parallel().dp_size if get_parallel().enable_dp_attention else 1 get_parallel().dp_size if get_parallel().enable_dp_attention else 1
@@ -681,24 +687,26 @@ class DataParallelController:
# - Attention: Global(TP) -> DP -> ATTN_CP -> ATTN_TP (innermost) # - Attention: Global(TP) -> DP -> ATTN_CP -> ATTN_TP (innermost)
# - MoE: Global(TP) -> MOE_DP -> EP -> MOE_TP (innermost) # - MoE: Global(TP) -> MOE_DP -> EP -> MOE_TP (innermost)
attn_tp_size = ( attn_tp_size = (
server_args.tp_size // attn_dp_size // get_parallel().attn_cp_size get_parallel().tp_size
// attn_dp_size
// get_parallel().attn_cp_size
) )
attn_cp_rank = (tp_rank // attn_tp_size) % get_parallel().attn_cp_size attn_cp_rank = (tp_rank // attn_tp_size) % get_parallel().attn_cp_size
moe_dp_rank = tp_rank // ( moe_dp_rank = tp_rank // (
server_args.tp_size // get_parallel().moe_dp_size get_parallel().tp_size // get_parallel().moe_dp_size
) )
moe_ep_rank = ( moe_ep_rank = (
tp_rank tp_rank
% (server_args.tp_size // get_parallel().moe_dp_size) % (get_parallel().tp_size // get_parallel().moe_dp_size)
// ( // (
server_args.tp_size get_parallel().tp_size
// get_parallel().moe_dp_size // get_parallel().moe_dp_size
// get_parallel().ep_size // get_parallel().ep_size
) )
) )
# Scheduler internals use local ranks; logs use global ranks. # Scheduler internals use local ranks; logs use global ranks.
offset = server_args.ep_join_rank_offset offset = get_parallel().ep_join_rank_offset
display_tp_rank = tp_rank + offset display_tp_rank = tp_rank + offset
display_moe_ep_rank = moe_ep_rank + offset display_moe_ep_rank = moe_ep_rank + offset
display_dp_rank = dp_rank + offset if dp_rank is not None else None display_dp_rank = dp_rank + offset if dp_rank is not None else None
@@ -828,11 +836,11 @@ def run_data_parallel_controller_process(
# This process reads the config namespaces before spawning schedulers. # This process reads the config namespaces before spawning schedulers.
publish(server_args, role="dp_controller") publish(server_args, role="dp_controller")
configure_logger(server_args) configure_logger(server_args)
if server_args.enable_trace: if get_observability().enable_trace:
process_tracing_init( process_tracing_init(
server_args.otlp_traces_endpoint, get_observability().otlp_traces_endpoint,
"sglang", "sglang",
trace_modules=server_args.trace_modules, trace_modules=get_observability().trace_modules,
) )
thread_label = "DP Controller" thread_label = "DP Controller"
if get_disagg().disaggregation_mode == "prefill": if get_disagg().disaggregation_mode == "prefill":
@@ -858,7 +866,7 @@ def run_data_parallel_controller_process(
} }
) )
# The primary owns routing for the expanded scheduler set. # The primary owns routing for the expanded scheduler set.
if server_args.node_rank == 0 and not server_args.is_ep_scale_joiner: if get_parallel().node_rank == 0 and not server_args.is_ep_scale_joiner:
controller.event_loop() controller.event_loop()
for proc in controller.scheduler_procs: for proc in controller.scheduler_procs:
proc.join() proc.join()
@@ -127,22 +127,22 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
# In multi-tokenizer mode, results are pushed back to each TokenizerWorker # In multi-tokenizer mode, results are pushed back to each TokenizerWorker
# directly via SocketMapping inside multi_http_worker_event_loop, so the # directly via SocketMapping inside multi_http_worker_event_loop, so the
# single send_to_tokenizer socket is unused. # single send_to_tokenizer socket is unused.
if server_args.tokenizer_worker_num == 1: if get_serving().tokenizer_worker_num == 1:
self.send_to_tokenizer = get_zmq_socket( self.send_to_tokenizer = get_zmq_socket(
context, zmq.PUSH, port_args.tokenizer_ipc_name, False context, zmq.PUSH, port_args.tokenizer_ipc_name, False
) )
def init_tokenizer(self, server_args: ServerArgs): def init_tokenizer(self, server_args: ServerArgs):
if server_args.skip_tokenizer_init: if get_serving().skip_tokenizer_init:
self.tokenizer = None self.tokenizer = None
self.vocab_size = None self.vocab_size = None
else: else:
self.tokenizer = get_tokenizer( self.tokenizer = get_tokenizer(
get_serving().tokenizer_path, get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode, tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=get_model().trust_remote_code, trust_remote_code=get_model().trust_remote_code,
revision=server_args.revision, revision=get_model().revision,
tokenizer_backend=server_args.tokenizer_backend, tokenizer_backend=get_serving().tokenizer_backend,
) )
try: try:
self.vocab_size = len(self.tokenizer) self.vocab_size = len(self.tokenizer)
@@ -151,7 +151,9 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
def init_running_status(self, server_args: ServerArgs): def init_running_status(self, server_args: ServerArgs):
self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES) self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES)
self.disable_tokenizer_batch_decode = server_args.disable_tokenizer_batch_decode self.disable_tokenizer_batch_decode = (
get_serving().disable_tokenizer_batch_decode
)
self.is_tool_call_parser_gpt_oss = get_serving().tool_call_parser == "gpt-oss" self.is_tool_call_parser_gpt_oss = get_serving().tool_call_parser == "gpt-oss"
self.soft_watchdog = Watchdog.create( self.soft_watchdog = Watchdog.create(
@@ -548,7 +550,7 @@ def run_detokenizer_process(
manager = None manager = None
try: try:
manager = detokenizer_manager_class(server_args, port_args) manager = detokenizer_manager_class(server_args, port_args)
if server_args.tokenizer_worker_num == 1: if get_serving().tokenizer_worker_num == 1:
manager.event_loop() manager.event_loop()
else: else:
manager.multi_http_worker_event_loop() manager.multi_http_worker_event_loop()
@@ -6,6 +6,7 @@ import pkgutil
from sglang.srt.configs.model_config import ModelImpl from sglang.srt.configs.model_config import ModelImpl
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
from sglang.srt.runtime_context import get_model
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -41,10 +42,10 @@ def import_processors(package_name: str, overwrite: bool = False):
PROCESSOR_MAPPING[arch] = cls PROCESSOR_MAPPING[arch] = cls
def get_mm_processor_cls(hf_config, server_args: ServerArgs, model_config=None): def get_mm_processor_cls(hf_config, model_config=None):
"""The class :func:`get_mm_processor` would instantiate, or ``None`` when the """The class :func:`get_mm_processor` would instantiate, or ``None`` when the
architecture has no registered processor.""" architecture has no registered processor."""
model_impl = str(getattr(server_args, "model_impl", "auto")).lower() model_impl = str(get_model().model_impl).lower()
uses_transformers_backend = model_impl == "transformers" uses_transformers_backend = model_impl == "transformers"
if model_impl == "auto" and model_config is not None: if model_impl == "auto" and model_config is not None:
from sglang.srt.model_loader.utils import get_resolved_model_impl from sglang.srt.model_loader.utils import get_resolved_model_impl
@@ -79,7 +80,7 @@ def get_mm_processor(
model_config=None, model_config=None,
**kwargs, **kwargs,
) -> BaseMultimodalProcessor: ) -> BaseMultimodalProcessor:
processor_cls = get_mm_processor_cls(hf_config, server_args, model_config) processor_cls = get_mm_processor_cls(hf_config, model_config)
if processor_cls is None: if processor_cls is None:
raise ValueError( raise ValueError(
f"No processor registered for architecture: {hf_config.architectures}.\n" f"No processor registered for architecture: {hf_config.architectures}.\n"
+1 -1
View File
@@ -992,7 +992,7 @@ class Scheduler(
# Initialize GEMM-related configuration for FP8 and FP4 backends. # Initialize GEMM-related configuration for FP8 and FP4 backends.
initialize_fp8_gemm_config() initialize_fp8_gemm_config()
initialize_fp4_gemm_config() initialize_fp4_gemm_config()
initialize_bf16_gemm_config(self.server_args) initialize_bf16_gemm_config()
# This must be called after initialize_moe_config # This must be called after initialize_moe_config
self.require_mlp_sync = require_mlp_sync() self.require_mlp_sync = require_mlp_sync()
@@ -248,10 +248,10 @@ class SchedulerMetricsReporter:
self.kv_transfer_latency_ms: float = 0.0 self.kv_transfer_latency_ms: float = 0.0
self.enable_mfu_metrics = False self.enable_mfu_metrics = False
self.decode_log_interval = self.scheduler.server_args.decode_log_interval self.decode_log_interval = get_observability().decode_log_interval
if self.enable_metrics: if self.enable_metrics:
self.enable_mfu_metrics = self.scheduler.server_args.enable_mfu_metrics self.enable_mfu_metrics = get_observability().enable_mfu_metrics
if self.enable_mfu_metrics: if self.enable_mfu_metrics:
self._init_estimated_perf_constants() self._init_estimated_perf_constants()
self._mfu_log_flops = 0.0 self._mfu_log_flops = 0.0
@@ -41,6 +41,7 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromTensorReqInput, UpdateWeightsFromTensorReqInput,
UpdateWeightsFromTensorReqOutput, UpdateWeightsFromTensorReqOutput,
) )
from sglang.srt.runtime_context import get_model
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -198,7 +199,7 @@ class SchedulerWeightUpdaterManager:
freeing them would leave the daemon and every peer pointing at released freeing them would leave the daemon and every peer pointing at released
memory. memory.
""" """
mode = self.tp_worker.model_runner.server_args.weight_cache_mode mode = get_model().weight_cache_mode
if mode != "off": if mode != "off":
raise RuntimeError( raise RuntimeError(
f"[weight_cache] {op} of model weights is not supported while the " f"[weight_cache] {op} of model weights is not supported while the "
@@ -79,6 +79,7 @@ from sglang.srt.managers.load_snapshot import LoadSnapshot
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_lora, get_lora,
get_parallel, get_parallel,
get_serving,
get_spec, get_spec,
) )
from sglang.srt.server_args import LoRARef from sglang.srt.server_args import LoRARef
@@ -209,9 +210,7 @@ class TokenizerControlMixin:
iter_external_corpus_chunks, iter_external_corpus_chunks,
) )
max_tokens = ( max_tokens = get_spec().speculative_ngram_external_corpus_max_tokens
self.server_args.speculative_ngram_external_corpus_max_tokens
)
obj.token_chunks = list( obj.token_chunks = list(
iter_external_corpus_chunks( iter_external_corpus_chunks(
obj.file_path, self.tokenizer, max_tokens obj.file_path, self.tokenizer, max_tokens
@@ -222,9 +221,7 @@ class TokenizerControlMixin:
SEPARATOR_TOKEN, SEPARATOR_TOKEN,
) )
max_tokens = ( max_tokens = get_spec().speculative_ngram_external_corpus_max_tokens
self.server_args.speculative_ngram_external_corpus_max_tokens
)
token_chunks = [] token_chunks = []
total_tokens = 0 total_tokens = 0
has_prev = False has_prev = False
@@ -642,10 +639,10 @@ class TokenizerControlMixin:
await self.lora_registry.register(new_adapter) await self.lora_registry.register(new_adapter)
self.lora_ref_cache[obj.lora_name] = new_adapter self.lora_ref_cache[obj.lora_name] = new_adapter
if self.server_args.max_loaded_loras is not None: if get_lora().max_loaded_loras is not None:
while ( while (
self.lora_registry.num_registered_loras self.lora_registry.num_registered_loras
> self.server_args.max_loaded_loras > get_lora().max_loaded_loras
): ):
lru_lora_name = await self.lora_registry.lru_lora_name( lru_lora_name = await self.lora_registry.lru_lora_name(
exclude_pinned=True exclude_pinned=True
@@ -659,7 +656,7 @@ class TokenizerControlMixin:
logger.info( logger.info(
f"Unloading least recently used LoRA adapter '{lru_lora_name}' " f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
f"(current number of adapters: {self.lora_registry.num_registered_loras}, " f"(current number of adapters: {self.lora_registry.num_registered_loras}, "
f"max allowed: {self.server_args.max_loaded_loras})" f"max allowed: {get_lora().max_loaded_loras})"
) )
unload_result = await self._unload_lora_adapter_locked( unload_result = await self._unload_lora_adapter_locked(
@@ -718,10 +715,10 @@ class TokenizerControlMixin:
if result.success: if result.success:
await self.lora_registry.register(new_adapter) await self.lora_registry.register(new_adapter)
self.lora_ref_cache[obj.lora_name] = new_adapter self.lora_ref_cache[obj.lora_name] = new_adapter
if self.server_args.max_loaded_loras is not None: if get_lora().max_loaded_loras is not None:
while ( while (
self.lora_registry.num_registered_loras self.lora_registry.num_registered_loras
> self.server_args.max_loaded_loras > get_lora().max_loaded_loras
): ):
lru_lora_name = await self.lora_registry.lru_lora_name( lru_lora_name = await self.lora_registry.lru_lora_name(
exclude_pinned=True exclude_pinned=True
@@ -735,7 +732,7 @@ class TokenizerControlMixin:
logger.info( logger.info(
f"Unloading least recently used LoRA adapter '{lru_lora_name}' " f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
f"(current number of adapters: {self.lora_registry.num_registered_loras}, " f"(current number of adapters: {self.lora_registry.num_registered_loras}, "
f"max allowed: {self.server_args.max_loaded_loras})" f"max allowed: {get_lora().max_loaded_loras})"
) )
unload_result = await self._unload_lora_adapter_locked( unload_result = await self._unload_lora_adapter_locked(
@@ -905,7 +902,7 @@ class TokenizerControlMixin:
): ):
self.auto_create_handle_loop() self.auto_create_handle_loop()
if obj.streaming: if obj.streaming:
if not self.server_args.enable_streaming_session: if not get_serving().enable_streaming_session:
raise ValueError( raise ValueError(
"Streaming sessions are disabled. " "Streaming sessions are disabled. "
"Please relaunch with --enable-streaming-session." "Please relaunch with --enable-streaming-session."
@@ -9,6 +9,7 @@ from sglang.srt.configs.model_config import is_cross_encoding_pooler_model
from sglang.srt.constants import MIS_DELIMITER_TOKEN_ID from sglang.srt.constants import MIS_DELIMITER_TOKEN_ID
from sglang.srt.managers.embed_types import PositionalEmbeds from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
from sglang.srt.runtime_context import get_exec
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -509,7 +510,7 @@ class TokenizerManagerScoreMixin:
) )
# Check if multi-item scoring is enabled # Check if multi-item scoring is enabled
use_multi_item_scoring = self.server_args.enable_mis use_multi_item_scoring = get_exec().features.enable_mis
input_ids = None input_ids = None
text_prompts = None text_prompts = None
+3 -3
View File
@@ -108,7 +108,7 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache:
logger.info("Using experimental C++ radix tree implementation.") logger.info("Using experimental C++ radix tree implementation.")
return RadixCacheCpp(params=params, server_args=server_args) return RadixCacheCpp(params=params, server_args=server_args)
if server_args.enable_unified_cache_external_linker: if get_memory().enable_unified_cache_external_linker:
return _create_unified_radix_cache(ctx, server_args, params) return _create_unified_radix_cache(ctx, server_args, params)
if ctx.is_hybrid_swa and ctx.full_tokens_per_layer == 0: if ctx.is_hybrid_swa and ctx.full_tokens_per_layer == 0:
@@ -196,8 +196,8 @@ def _create_unified_radix_cache(
ctx.tp_worker.register_hicache_layer_transfer_counter( ctx.tp_worker.register_hicache_layer_transfer_counter(
cache.cache_controller.layer_done_counter cache.cache_controller.layer_done_counter
) )
elif server_args.enable_unified_cache_external_linker: elif get_memory().enable_unified_cache_external_linker:
backend = server_args.unified_cache_external_linker_backend backend = get_memory().unified_cache_external_linker_backend
if backend == "mooncake": if backend == "mooncake":
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_direct_linker import ( from sglang.srt.mem_cache.storage.mooncake_store.mooncake_direct_linker import (
MooncakeDirectLinker, MooncakeDirectLinker,
@@ -20,7 +20,11 @@ from sglang.srt.mem_cache.hybrid_cache.linker_pool_assembler import (
resolve_hybrid_device_pool_group, resolve_hybrid_device_pool_group,
) )
from sglang.srt.mem_cache.unified_cache.unified_cache_linker import UnifiedCacheLinker from sglang.srt.mem_cache.unified_cache.unified_cache_linker import UnifiedCacheLinker
from sglang.srt.runtime_context import get_memory, get_model from sglang.srt.runtime_context import (
get_memory,
get_model,
get_parallel,
)
from sglang.srt.utils import freeze_gc, get_device_module from sglang.srt.utils import freeze_gc, get_device_module
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -102,7 +106,7 @@ class MooncakeDirectLinker(UnifiedCacheLinker):
self.num_layers = self.pool_group.num_layers self.num_layers = self.pool_group.num_layers
tp_rank = 0 tp_rank = 0
tp_size = server_args.tp_size tp_size = get_parallel().tp_size
tp_group = params.attn_tp_cache_group or params.tp_cache_group tp_group = params.attn_tp_cache_group or params.tp_cache_group
if torch.distributed.is_available() and torch.distributed.is_initialized(): if torch.distributed.is_available() and torch.distributed.is_initialized():
tp_rank = torch.distributed.get_rank(group=tp_group) tp_rank = torch.distributed.get_rank(group=tp_group)
@@ -596,19 +596,15 @@ class CPUGraphRunner:
self.graphs_cross = {} self.graphs_cross = {}
self.output_buffers = {} self.output_buffers = {}
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = get_exec().graph.disable_cuda_graph_padding
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
self.require_gathered_buffer = require_gathered_buffer() self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather() self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync() self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather() self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_two_batch_overlap = ( self.enable_two_batch_overlap = get_exec().overlap.enable_two_batch_overlap
model_runner.server_args.enable_two_batch_overlap
)
self.speculative_algorithm = get_spec().speculative_algorithm self.speculative_algorithm = get_spec().speculative_algorithm
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
model_runner.server_args.enable_profile_cuda_graph
)
self.tp_size = get_parallel().tp_size self.tp_size = get_parallel().tp_size
self.dp_size = get_parallel().dp_size self.dp_size = get_parallel().dp_size
self.pp_size = get_parallel().pp_size self.pp_size = get_parallel().pp_size
@@ -14,7 +14,10 @@ from mindspore._c_expression import GroupOptions
from mindspore.communication import create_group from mindspore.communication import create_group
from sglang.srt.distributed.parallel_state import _groups from sglang.srt.distributed.parallel_state import _groups
from sglang.srt.runtime_context import get_serving from sglang.srt.runtime_context import (
get_parallel,
get_serving,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -106,9 +109,9 @@ def reuse_hccl_comm():
create_group(group_name, group().ranks, group_options) create_group(group_name, group().ranks, group_options)
def init_ms_distributed(world_size, rank, local_rank, server_args, port): def init_ms_distributed(world_size, rank, local_rank, port):
if server_args.dist_init_addr: if get_parallel().dist_init_addr:
dist_init_method = f"tcp://{server_args.dist_init_addr}" dist_init_method = f"tcp://{get_parallel().dist_init_addr}"
else: else:
dist_init_method = f"tcp://{get_serving().host}:{port}" dist_init_method = f"tcp://{get_serving().host}:{port}"
set_ms_parallel_env(rank, local_rank, world_size, dist_init_method) set_ms_parallel_env(rank, local_rank, world_size, dist_init_method)
@@ -316,7 +316,7 @@ class ModelRunner:
def supports_sampling_observer(self) -> bool: def supports_sampling_observer(self) -> bool:
"""Whether this runner's sampling path publishes observer output.""" """Whether this runner's sampling path publishes observer output."""
return self.server_args.dllm_algorithm is None and self.spec_algorithm.is_none() return get_exec().dllm.dllm_algorithm is None and self.spec_algorithm.is_none()
def __init__( def __init__(
self, self,
@@ -644,7 +644,6 @@ class ModelRunner:
world_size=self.ps.tp_size * self.ps.pp_size, world_size=self.ps.tp_size * self.ps.pp_size,
rank=self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank, rank=self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank,
local_rank=self.gpu_id, local_rank=self.gpu_id,
server_args=self.server_args,
port=self.dist_port, port=self.dist_port,
) )
@@ -42,6 +42,8 @@ from sglang.srt.runtime_context import (
get_disagg, get_disagg,
get_exec, get_exec,
get_flags, get_flags,
get_model,
get_observability,
get_parallel, get_parallel,
get_schedule, get_schedule,
get_spec, get_spec,
@@ -281,10 +283,8 @@ def capture_cuda_graphs(
# not traced into any captured graph — capture stays hook-free and hooks # not traced into any captured graph — capture stays hook-free and hooks
# fire only on the eager forward path (capture replay never runs Python # fire only on the eager forward path (capture replay never runs Python
# hooks anyway). # hooks anyway).
if model_runner.server_args.forward_hooks: if get_observability().forward_hooks:
register_forward_hooks( register_forward_hooks(model_runner.model, get_observability().forward_hooks)
model_runner.model, model_runner.server_args.forward_hooks
)
prealloc_symmetric_memory_pool( prealloc_symmetric_memory_pool(
is_draft_worker=model_runner.is_draft_worker, is_draft_worker=model_runner.is_draft_worker,
@@ -537,7 +537,7 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> GraphCapture:
if not model_runner.is_generation: if not model_runner.is_generation:
# TODO: Currently, cuda graph only captures decode steps, which only exists for generation models # TODO: Currently, cuda graph only captures decode steps, which only exists for generation models
return no_capture return no_capture
if model_runner.server_args.model_impl.lower() == ModelImpl.MINDSPORE: if get_model().model_impl.lower() == ModelImpl.MINDSPORE:
return no_capture return no_capture
if model_runner.device != "cpu" and check_cuda_graph_backend( if model_runner.device != "cpu" and check_cuda_graph_backend(
Phase.DECODE, Backend.DISABLED Phase.DECODE, Backend.DISABLED
@@ -12,6 +12,7 @@ from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
from sglang.srt.model_loader.utils import set_default_torch_dtype from sglang.srt.model_loader.utils import set_default_torch_dtype
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_model
from sglang.srt.utils import ( from sglang.srt.utils import (
MultiprocessingSerializer, MultiprocessingSerializer,
dynamic_import, dynamic_import,
@@ -128,7 +129,7 @@ class WeightUpdater:
param.data is the daemon's master copy shared with every co-attached param.data is the daemon's master copy shared with every co-attached
engine, so an in-place update would silently corrupt them all. engine, so an in-place update would silently corrupt them all.
""" """
mode = self.get_model_runner().server_args.weight_cache_mode mode = get_model().weight_cache_mode
if mode != "off": if mode != "off":
raise RuntimeError( raise RuntimeError(
f"[weight_cache] {op} is not supported while the weight cache is " f"[weight_cache] {op} is not supported while the weight cache is "
@@ -42,6 +42,7 @@ from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_disagg, get_disagg,
get_exec,
get_memory, get_memory,
get_parallel, get_parallel,
get_schedule, get_schedule,
@@ -887,7 +888,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
# keeps the FP8 estimate. # keeps the FP8 estimate.
self.indexer_bytes_per_token = get_dsv4_indexer_bytes_per_token( self.indexer_bytes_per_token = get_dsv4_indexer_bytes_per_token(
self.indexer_head_dim, self.indexer_head_dim,
_is_hip and kvc.server_args.enable_deepseek_v4_fp4_indexer, _is_hip and get_exec().kernel.enable_deepseek_v4_fp4_indexer,
) )
self.context_len = kvc.model_config.context_len self.context_len = kvc.model_config.context_len
# PP-local slice; matches DeepSeekV4TokenToKVPool's stage_ratios. # PP-local slice; matches DeepSeekV4TokenToKVPool's stage_ratios.
@@ -220,7 +220,7 @@ class BaseRunner(ABC):
# elastic-EP scale-up rewrites dp_size on the published config # elastic-EP scale-up rewrites dp_size on the published config
self.dp_size = get_parallel().dp_size self.dp_size = get_parallel().dp_size
self.pp_size = get_parallel().pp_size self.pp_size = get_parallel().pp_size
self.enable_pdmux = model_runner.server_args.enable_pdmux self.enable_pdmux = get_disagg().enable_pdmux
self.return_hidden_states_mode = ( self.return_hidden_states_mode = (
CaptureHiddenMode.NULL CaptureHiddenMode.NULL
if model_runner.is_draft_worker if model_runner.is_draft_worker
@@ -229,7 +229,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# --- core state ------------------------------------------------ # --- core state ------------------------------------------------
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = get_exec().graph.disable_cuda_graph_padding
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
self.require_mlp_tp_gather = ( self.require_mlp_tp_gather = (
require_mlp_tp_gather() and not self._forward_is_dp_local(model_runner) require_mlp_tp_gather() and not self._forward_is_dp_local(model_runner)
@@ -244,18 +244,14 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self.require_mlp_sync = ( self.require_mlp_sync = (
get_parallel().enable_dp_attention or self.require_gathered_buffer get_parallel().enable_dp_attention or self.require_gathered_buffer
) )
self.enable_two_batch_overlap = ( self.enable_two_batch_overlap = get_exec().overlap.enable_two_batch_overlap
model_runner.server_args.enable_two_batch_overlap
)
self.use_ngram_embedding = model_runner.ngram_embedding_manager.enabled self.use_ngram_embedding = model_runner.ngram_embedding_manager.enabled
if self.use_ngram_embedding: if self.use_ngram_embedding:
hf_config = model_runner.model_config.hf_config hf_config = model_runner.model_config.hf_config
self.ngram_embedding_n = hf_config.ngram_embedding_n self.ngram_embedding_n = hf_config.ngram_embedding_n
self.ngram_embedding_k = hf_config.ngram_embedding_k self.ngram_embedding_k = hf_config.ngram_embedding_k
self.speculative_algorithm = get_spec().speculative_algorithm self.speculative_algorithm = get_spec().speculative_algorithm
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
model_runner.server_args.enable_profile_cuda_graph
)
# --- DSA dense-decode dual-graph ------------------------------- # --- DSA dense-decode dual-graph -------------------------------
# Capture a "dense" (k-only, skip-indexer) and a "sparse" (full indexer) # Capture a "dense" (k-only, skip-indexer) and a "sparse" (full indexer)
@@ -1067,7 +1063,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# Trigger CUDA graph capture for specific shapes. # Trigger CUDA graph capture for specific shapes.
# Capture the large shapes first so that the smaller shapes # Capture the large shapes first so that the smaller shapes
# can reuse the memory pool allocated for the large shapes. # can reuse the memory pool allocated for the large shapes.
with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc): with freeze_gc(get_exec().graph.enable_cudagraph_gc):
if not self.enable_pdmux: if not self.enable_pdmux:
with ( with (
graph_capture( graph_capture(
@@ -1413,7 +1413,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
# Warm up + autotune kernels once before capture (run-once across the # Warm up + autotune kernels once before capture (run-once across the
# decode + prefill runners; see BaseRunner.warmup). # decode + prefill runners; see BaseRunner.warmup).
self.warmup() self.warmup()
with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc): with freeze_gc(get_exec().graph.enable_cudagraph_gc):
with graph_capture( with graph_capture(
stream=get_or_create_global_graph_capture_stream() stream=get_or_create_global_graph_capture_stream()
) as graph_capture_context: ) as graph_capture_context:
@@ -8,6 +8,8 @@ from dataclasses import dataclass, field
import torch import torch
from sglang.srt.runtime_context import get_exec
try: try:
from cuda.bindings import driver as cuda_drv from cuda.bindings import driver as cuda_drv
from cuda.bindings import runtime as cuda_rt from cuda.bindings import runtime as cuda_rt
@@ -314,7 +316,7 @@ class DedupedCudaGraphMixin:
server_args = getattr(model_runner, "server_args", None) server_args = getattr(model_runner, "server_args", None)
return bool( return bool(
server_args is not None server_args is not None
and getattr(server_args, "enable_memory_saver", False) and get_exec().features.enable_memory_saver
and get_bool_env_var("SGLANG_MEMORY_SAVER_CUDA_GRAPH") and get_bool_env_var("SGLANG_MEMORY_SAVER_CUDA_GRAPH")
) )
@@ -125,7 +125,7 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
config = CompilationConfig( config = CompilationConfig(
num_tokens, num_tokens,
compiler, compiler,
server_args.enable_torch_compile_debug_mode, get_exec().graph.enable_torch_compile_debug_mode,
) )
if get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake(): if get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake():
@@ -43,6 +43,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
get_mm_feature_pool_size_per_worker, get_mm_feature_pool_size_per_worker,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec,
get_mm, get_mm,
get_serving, get_serving,
) )
@@ -723,7 +724,7 @@ class BaseMultimodalProcessor(ABC):
preprocessing worker there is one more competitor for that device rather preprocessing worker there is one more competitor for that device rather
than added parallelism. than added parallelism.
""" """
if _is_cpu or self.server_args.rl_on_policy_target is not None: if _is_cpu or get_exec().deterministic.rl_on_policy_target is not None:
return False return False
if self.disable_fast_image_processor: if self.disable_fast_image_processor:
return False return False
@@ -759,7 +760,7 @@ class BaseMultimodalProcessor(ABC):
tokenizer process each carry their own ``base_gpu_id``. tokenizer process each carry their own ``base_gpu_id``.
""" """
server_args = self.server_args server_args = self.server_args
if _is_cpu or server_args.rl_on_policy_target is not None: if _is_cpu or get_exec().deterministic.rl_on_policy_target is not None:
return "cpu" return "cpu"
if _is_xpu: if _is_xpu:
return "xpu" return "xpu"
@@ -20,6 +20,7 @@ from sglang.srt.multimodal.processors.base_processor import (
BaseMultiModalProcessorOutput, BaseMultiModalProcessorOutput,
MultimodalSpecialTokens, MultimodalSpecialTokens,
) )
from sglang.srt.runtime_context import get_model
from sglang.srt.utils import get_device from sglang.srt.utils import get_device
from sglang.srt.utils.video_decoder import VideoDecoderWrapper from sglang.srt.utils.video_decoder import VideoDecoderWrapper
@@ -135,7 +136,7 @@ class InternVLProcessor(BaseMultimodalProcessor):
).build(_image_processor) ).build(_image_processor)
self.max_context_len = ( self.max_context_len = (
getattr(server_args, "context_length", None) get_model().context_length
or getattr(server_args, "max_context_len", None) or getattr(server_args, "max_context_len", None)
or getattr(hf_config, "max_position_embeddings", None) or getattr(hf_config, "max_position_embeddings", None)
or getattr(text_cfg, "max_position_embeddings", None) or getattr(text_cfg, "max_position_embeddings", None)
@@ -23,6 +23,7 @@ from sglang.srt.multimodal.processors.kimi_common import KimiGridMMDataMixin
from sglang.srt.multimodal.transport.cuda_ipc import ( from sglang.srt.multimodal.transport.cuda_ipc import (
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY, DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
) )
from sglang.srt.runtime_context import get_mm
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# GPU image preprocessing utilities (resize, pad, normalize, patchify on CUDA) # GPU image preprocessing utilities (resize, pad, normalize, patchify on CUDA)
@@ -595,7 +596,7 @@ class KimiK2_5VLImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
# its GPU transport proxy lazy until that assignment is known, avoiding a full # its GPU transport proxy lazy until that assignment is known, avoiding a full
# image copy to every rank. The scheduler only honors this marker once # image copy to every rank. The scheduler only honors this marker once
# the processor has already set the item's hash and pad value. # the processor has already set the item's hash and pad value.
if self.keep_mm_features_on_device and self.server_args.mm_enable_dp_encoder: if self.keep_mm_features_on_device and get_mm().mm_enable_dp_encoder:
for item in mm_items: for item in mm_items:
item.model_specific_data[DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY] = ( item.model_specific_data[DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY] = (
True True
@@ -46,6 +46,7 @@ from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor, BaseMultimodalProcessor,
MultimodalSpecialTokens, MultimodalSpecialTokens,
) )
from sglang.srt.runtime_context import get_model
from sglang.srt.utils.common import sample_video_frames from sglang.srt.utils.common import sample_video_frames
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -143,7 +144,7 @@ class NanoNemotronVLImageProcessor(BaseMultimodalProcessor):
hf_config, "video_maintain_aspect_ratio", True hf_config, "video_maintain_aspect_ratio", True
) )
self.max_model_len = getattr(server_args, "context_length", None) or 8192 self.max_model_len = get_model().context_length or 8192
self.PLACEHOLDER = self.tokenizer.unk_token self.PLACEHOLDER = self.tokenizer.unk_token
assert isinstance(self.PLACEHOLDER, str) assert isinstance(self.PLACEHOLDER, str)
@@ -41,6 +41,7 @@ from sglang.srt.multimodal.processors.base_processor import (
from sglang.srt.multimodal.transport.cuda_ipc import ( from sglang.srt.multimodal.transport.cuda_ipc import (
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY, DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
) )
from sglang.srt.runtime_context import get_mm
from sglang.srt.utils import cpu_has_amx_support, is_cpu from sglang.srt.utils import cpu_has_amx_support, is_cpu
from sglang.srt.utils.video_decoder import VideoDecoderWrapper from sglang.srt.utils.video_decoder import VideoDecoderWrapper
from sglang.utils import logger from sglang.utils import logger
@@ -897,7 +898,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
def _mark_dp_encoder_features_for_deferred_reconstruction(self, mm_items): def _mark_dp_encoder_features_for_deferred_reconstruction(self, mm_items):
if not ( if not (
self.keep_mm_features_on_device self.keep_mm_features_on_device
and self.server_args.mm_enable_dp_encoder and get_mm().mm_enable_dp_encoder
and self.model_type and self.model_type
in ("qwen3_vl", "qwen3_vl_moe", "qwen3_5", "qwen3_5_moe") in ("qwen3_vl", "qwen3_vl_moe", "qwen3_5", "qwen3_5_moe")
): ):
@@ -83,7 +83,10 @@ class FileRequestMetricsExporter(RequestMetricsExporter):
out_skip_names: Optional[set[str]], out_skip_names: Optional[set[str]],
): ):
super().__init__(server_args, obj_skip_names, out_skip_names) super().__init__(server_args, obj_skip_names, out_skip_names)
self.export_dir = getattr(server_args, "export_metrics_to_file_dir") # Given at construction, not read from the process: the exporter is
# handed the directory it writes to, and a test builds several with
# different ones.
self.export_dir = server_args.export_metrics_to_file_dir
os.makedirs(self.export_dir, exist_ok=True) os.makedirs(self.export_dir, exist_ok=True)
# File handler state management # File handler state management
+2 -1
View File
@@ -50,6 +50,7 @@ from sglang.srt.parser.template_detection import (
detect_reasoning_pattern, detect_reasoning_pattern,
match_rules, match_rules,
) )
from sglang.srt.runtime_context import get_serving
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -384,7 +385,7 @@ class TemplateManager:
logger.info(f"Multiple HuggingFace chat templates available: {available_names}") logger.info(f"Multiple HuggingFace chat templates available: {available_names}")
# Use specified template if provided # Use specified template if provided
if preferred_name := tokenizer_manager.server_args.hf_chat_template_name: if preferred_name := get_serving().hf_chat_template_name:
if preferred_name not in templates: if preferred_name not in templates:
raise ValueError( raise ValueError(
f"Specified template '{preferred_name}' not found. " f"Specified template '{preferred_name}' not found. "
+1 -3
View File
@@ -203,9 +203,7 @@ def _create_scheduler_actor(
rank0_node_ip: IP of rank-0's node, used for NCCL rendezvous. rank0_node_ip: IP of rank-0's node, used for NCCL rendezvous.
dist_init_addr: Distributed init address (tcp://rank0_node_ip:nccl_port). dist_init_addr: Distributed init address (tcp://rank0_node_ip:nccl_port).
""" """
attn_cp_rank, moe_dp_rank, moe_ep_rank = _compute_parallelism_ranks( attn_cp_rank, moe_dp_rank, moe_ep_rank = _compute_parallelism_ranks(tp_rank)
server_args, tp_rank
)
return SchedulerActor.options( return SchedulerActor.options(
num_cpus=0, num_cpus=0,
+7 -1
View File
@@ -1347,7 +1347,13 @@ ROLE_NAMESPACE_SETS: dict[str, frozenset[str] | None] = {
# The DP controller's static read set, checked against the module: the # The DP controller's static read set, checked against the module: the
# elastic-EP gate, the load-balance method, the watchdog timeout, and the # elastic-EP gate, the load-balance method, the watchdog timeout, and the
# disaggregation mode. # disaggregation mode.
"dp_controller": frozenset({"exec", "parallel", "device", "disagg"}), # `observability` and `serving` were added when the controller's metrics
# gate, tracing setup and worker-port broadcast stopped reading the record:
# under `enforce` the set is what the process may read, so a conversion
# that reaches a new namespace has to widen it in the same commit.
"dp_controller": frozenset(
{"exec", "parallel", "device", "disagg", "observability", "serving"}
),
# Record-mode audit (2026-08-06, text model, /generate + /get_server_info + # Record-mode audit (2026-08-06, text model, /generate + /get_server_info +
# /v1/models): reads exactly {"serving"} — the per-instance managers read # /v1/models): reads exactly {"serving"} — the per-instance managers read
# self.server_args by design. Still declared full, because that run did not # self.server_args by design. Still declared full, because that run did not
+1 -1
View File
@@ -167,7 +167,7 @@ class RustMmProcessor:
hf_config = self.model_config.hf_config hf_config = self.model_config.hf_config
mm_processor_cls = get_mm_processor_cls( mm_processor_cls = get_mm_processor_cls(
hf_config, self.server_args, model_config=self.model_config hf_config, model_config=self.model_config
) )
family = rust_mm_family_for( family = rust_mm_family_for(
mm_processor_cls, getattr(hf_config, "model_type", None) mm_processor_cls, getattr(hf_config, "model_type", None)
@@ -35,6 +35,7 @@ from sglang.srt.model_executor.runner_backend_utils import (
CUDA_GRAPH_CAPTURE_FAILED_MSG, CUDA_GRAPH_CAPTURE_FAILED_MSG,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec,
get_flags, get_flags,
get_parallel, get_parallel,
get_spec, get_spec,
@@ -115,14 +116,12 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
self.attn_dp_size = model_runner.ps.attn_dp_size self.attn_dp_size = model_runner.ps.attn_dp_size
self.pp_size = get_parallel().pp_size self.pp_size = get_parallel().pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = get_exec().graph.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer() self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather() self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync() self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather() self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
model_runner.server_args.enable_profile_cuda_graph
)
self.speculative_num_steps = ( self.speculative_num_steps = (
get_spec().speculative_num_steps get_spec().speculative_num_steps
if speculative_num_steps is None if speculative_num_steps is None
@@ -196,7 +195,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
(self.max_bs, self.model_runner.model_config.vocab_size), (self.max_bs, self.model_runner.model_config.vocab_size),
dtype=torch.float32, dtype=torch.float32,
) )
if self.model_runner.server_args.speculative_use_rejection_sampling if get_spec().speculative_use_rejection_sampling
else None else None
) )
_hidden_size, _hidden_dtype = get_draft_recurrent_hidden_state_spec( _hidden_size, _hidden_dtype = get_draft_recurrent_hidden_state_spec(
@@ -616,7 +615,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# Only rejection sampling reads temperatures (renorm_draft_probs); skip # Only rejection sampling reads temperatures (renorm_draft_probs); skip
# the copy otherwise to keep the non-RS path free of extra work. # the copy otherwise to keep the non-RS path free of extra work.
if ( if (
self.model_runner.server_args.speculative_use_rejection_sampling get_spec().speculative_use_rejection_sampling
and forward_batch.sampling_info is not None and forward_batch.sampling_info is not None
): ):
self.temperatures[:raw_bs].copy_( self.temperatures[:raw_bs].copy_(
@@ -36,6 +36,7 @@ from sglang.srt.model_executor.runner_backend_utils import (
CUDA_GRAPH_CAPTURE_FAILED_MSG, CUDA_GRAPH_CAPTURE_FAILED_MSG,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec,
get_flags, get_flags,
get_parallel, get_parallel,
get_spec, get_spec,
@@ -114,14 +115,12 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.attn_dp_size = model_runner.ps.attn_dp_size self.attn_dp_size = model_runner.ps.attn_dp_size
self.pp_size = get_parallel().pp_size self.pp_size = get_parallel().pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = get_exec().graph.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer() self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather() self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync() self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather() self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
model_runner.server_args.enable_profile_cuda_graph
)
self.speculative_num_steps = ( self.speculative_num_steps = (
get_spec().speculative_num_steps get_spec().speculative_num_steps
if speculative_num_steps is None if speculative_num_steps is None
@@ -33,6 +33,7 @@ from sglang.srt.model_executor.runner_backend_utils import (
CUDA_GRAPH_CAPTURE_FAILED_MSG, CUDA_GRAPH_CAPTURE_FAILED_MSG,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec,
get_flags, get_flags,
get_parallel, get_parallel,
get_spec, get_spec,
@@ -92,7 +93,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.device = model_runner.device self.device = model_runner.device
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = get_exec().graph.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer() self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather() self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync() self.require_mlp_sync = require_mlp_sync()
@@ -103,9 +104,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.speculative_num_steps = get_spec().speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.topk = get_spec().speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
model_runner.server_args.enable_profile_cuda_graph
)
self.attn_backend = self.draft_attn_backend self.attn_backend = self.draft_attn_backend
@@ -60,6 +60,8 @@ from sglang.srt.model_executor.runner_backend_utils import (
CUDA_GRAPH_CAPTURE_FAILED_MSG, CUDA_GRAPH_CAPTURE_FAILED_MSG,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_flags, get_flags,
get_parallel, get_parallel,
get_spec, get_spec,
@@ -157,18 +159,16 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.dp_size = get_parallel().dp_size self.dp_size = get_parallel().dp_size
self.pp_size = get_parallel().pp_size self.pp_size = get_parallel().pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = get_exec().graph.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer() self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather() self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync() self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather() self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_pdmux = model_runner.server_args.enable_pdmux self.enable_pdmux = get_disagg().enable_pdmux
self.speculative_num_steps = get_spec().speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.topk = get_spec().speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
model_runner.server_args.enable_profile_cuda_graph
)
self.attn_backend = self.eagle_worker.draft_extend_attn_backend_list[self.step] self.attn_backend = self.eagle_worker.draft_extend_attn_backend_list[self.step]
self.metadata_captured_in_graph = ( self.metadata_captured_in_graph = (
self.attn_backend.draft_extend_metadata_captured_in_graph() self.attn_backend.draft_extend_metadata_captured_in_graph()
+6 -2
View File
@@ -50,7 +50,11 @@ import torch.distributed as dist
from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.load_config import LoadConfig
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_parallel, publish from sglang.srt.runtime_context import (
get_exec,
get_parallel,
publish,
)
from .protocol import ( from .protocol import (
CacheConfig, CacheConfig,
@@ -465,7 +469,7 @@ class WeightCacheDaemon:
def _initialize_eplb_expert_location_metadata(self, model_config) -> None: def _initialize_eplb_expert_location_metadata(self, model_config) -> None:
"""Build the same initial physical expert layout as the engine.""" """Build the same initial physical expert layout as the engine."""
if not self.server_args.enable_eplb: if not get_exec().moe.enable_eplb:
return return
from sglang.srt.eplb.expert_location import ( from sglang.srt.eplb.expert_location import (
+14
View File
@@ -2143,6 +2143,20 @@ def server_args_variant(server_args, **fields):
return variant return variant
def enter_override(test_case, override):
"""Install a scoped context override for the length of one test.
`unittest.TestCase.enterContext` does exactly this in one call, but it is
Python 3.11+ and this package supports 3.10 (`requires-python = ">=3.10"`).
On 3.10 it raises `AttributeError: ... has no attribute 'enterContext'` --
and only there, so a developer on a newer interpreter sees every test pass
while CI does not.
"""
installed = override.install()
test_case.addCleanup(override.restore)
return installed
class CustomTestCase(unittest.TestCase): class CustomTestCase(unittest.TestCase):
def __init_subclass__(cls, **kwargs): def __init_subclass__(cls, **kwargs):
super().__init_subclass__(**kwargs) super().__init_subclass__(**kwargs)
+15 -1
View File
@@ -9,6 +9,7 @@ import unittest
from unittest.mock import MagicMock from unittest.mock import MagicMock
from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_amd_ci, register_cpu_ci from sglang.test.ci.ci_register import register_amd_ci, register_cpu_ci
@@ -16,13 +17,22 @@ register_amd_ci(est_time=30, suite="nightly-amd-1-gpu", nightly=True)
register_cpu_ci(est_time=7, suite="stage-b-test-cpu-intel") register_cpu_ci(est_time=7, suite="stage-b-test-cpu-intel")
def publish_config(case):
"""`OpenAIServingBase.__init__` reads `get_observability()`, so a case that
builds one needs a published config. The mock manager cannot stand in for
it: `MagicMock(spec=ServerArgs)` passes the `isinstance` guard, so the read
happens and there is no bag to answer from."""
reset_context()
case.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
class MockTokenizerManager: class MockTokenizerManager:
"""Mock TokenizerManager for testing.""" """Mock TokenizerManager for testing."""
def __init__(self, enable_lora=False): def __init__(self, enable_lora=False):
self.server_args = MagicMock(spec=ServerArgs) self.server_args = MagicMock(spec=ServerArgs)
self.server_args.enable_lora = enable_lora self.server_args.enable_lora = enable_lora
self.server_args.tokenizer_metrics_allowed_custom_labels = None
class ConcreteServingBase(OpenAIServingBase): class ConcreteServingBase(OpenAIServingBase):
@@ -42,6 +52,7 @@ class TestParseModelParameter(unittest.TestCase):
"""Test _parse_model_parameter method.""" """Test _parse_model_parameter method."""
def setUp(self): def setUp(self):
publish_config(self)
self.tokenizer_manager = MockTokenizerManager(enable_lora=True) self.tokenizer_manager = MockTokenizerManager(enable_lora=True)
self.serving = ConcreteServingBase(self.tokenizer_manager) self.serving = ConcreteServingBase(self.tokenizer_manager)
@@ -98,6 +109,7 @@ class TestResolveLoraPath(unittest.TestCase):
"""Test _resolve_lora_path method.""" """Test _resolve_lora_path method."""
def setUp(self): def setUp(self):
publish_config(self)
self.tokenizer_manager = MockTokenizerManager(enable_lora=True) self.tokenizer_manager = MockTokenizerManager(enable_lora=True)
self.serving = ConcreteServingBase(self.tokenizer_manager) self.serving = ConcreteServingBase(self.tokenizer_manager)
@@ -146,6 +158,7 @@ class TestIntegrationScenarios(unittest.TestCase):
"""Integration tests for common usage scenarios.""" """Integration tests for common usage scenarios."""
def setUp(self): def setUp(self):
publish_config(self)
self.tokenizer_manager = MockTokenizerManager(enable_lora=True) self.tokenizer_manager = MockTokenizerManager(enable_lora=True)
self.serving = ConcreteServingBase(self.tokenizer_manager) self.serving = ConcreteServingBase(self.tokenizer_manager)
@@ -197,6 +210,7 @@ class TestEdgeCases(unittest.TestCase):
"""Test edge cases and error conditions.""" """Test edge cases and error conditions."""
def setUp(self): def setUp(self):
publish_config(self)
self.tokenizer_manager = MockTokenizerManager(enable_lora=True) self.tokenizer_manager = MockTokenizerManager(enable_lora=True)
self.serving = ConcreteServingBase(self.tokenizer_manager) self.serving = ConcreteServingBase(self.tokenizer_manager)
@@ -14,7 +14,8 @@ from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import FINISH_ABORT from sglang.srt.managers.schedule_batch import FINISH_ABORT
from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.runtime_context import get_context from sglang.srt.runtime_context import get_context, publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -34,6 +35,12 @@ class FakeReceiver:
class TestDecodeQueueCleanup(CustomTestCase): class TestDecodeQueueCleanup(CustomTestCase):
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def test_paged_swa_retraction_resume_uses_physical_page_budget(self): def test_paged_swa_retraction_resume_uses_physical_page_budget(self):
# resume_retracted_reqs reads the retraction backend off the disagg # resume_retracted_reqs reads the retraction backend off the disagg
# bag, so the case publishes a config instead of injecting one. # bag, so the case publishes a config instead of injecting one.
@@ -43,9 +43,11 @@ from sglang.srt.parser.jinja_template_utils import (
jinja_template_may_reorder_tool_results, jinja_template_may_reorder_tool_results,
) )
from sglang.srt.parser.template_detection import ReasoningToggleConfig from sglang.srt.parser.template_detection import ReasoningToggleConfig
from sglang.srt.runtime_context import get_context, publish, reset_context
from sglang.srt.sampling.sampling_params import ( from sglang.srt.sampling.sampling_params import (
REQUEST_REASONING_END_TOKEN_IDS_KEY, REQUEST_REASONING_END_TOKEN_IDS_KEY,
) )
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import get_or_create_event_loop from sglang.srt.utils import get_or_create_event_loop
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
@@ -196,6 +198,22 @@ class _MockTemplateManager:
class ServingChatTestCase(unittest.TestCase): class ServingChatTestCase(unittest.TestCase):
# ------------- common fixtures ------------- # ------------- common fixtures -------------
def setUp(self): def setUp(self):
# The serving layer reads its config from the bags, so the fixture has
# to publish one rather than hang the values off a mock manager.
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(
model_path="dummy",
revision=None,
enable_cache_report=False,
tool_call_parser="hermes",
reasoning_parser=None,
stream_response_default_include_usage=False,
default_chat_template_kwargs=None,
),
role="tokenizer",
)
self.tm = _MockTokenizerManager() self.tm = _MockTokenizerManager()
self.template_manager = _MockTemplateManager() self.template_manager = _MockTemplateManager()
self.chat = OpenAIServingChat(self.tm, self.template_manager) self.chat = OpenAIServingChat(self.tm, self.template_manager)
@@ -3643,14 +3661,14 @@ class ServingChatTestCase(unittest.TestCase):
def test_continuous_usage_reports_cached_tokens(self): def test_continuous_usage_reports_cached_tokens(self):
"""continuous_usage_stats chunks include cached tokens when cache reporting is on.""" """continuous_usage_stats chunks include cached tokens when cache reporting is on."""
self.tm.server_args.enable_cache_report = True self.enterContext(get_context().override_server_args(enable_cache_report=True))
usages = self._collect_continuous_usage(cached_tokens=6) usages = self._collect_continuous_usage(cached_tokens=6)
self.assertTrue(usages, "continuous_usage_stats attached no usage") self.assertTrue(usages, "continuous_usage_stats attached no usage")
self.assertEqual(usages[0]["prompt_tokens_details"]["cached_tokens"], 6) self.assertEqual(usages[0]["prompt_tokens_details"]["cached_tokens"], 6)
def test_continuous_usage_omits_cached_tokens_when_report_disabled(self): def test_continuous_usage_omits_cached_tokens_when_report_disabled(self):
"""With cache reporting off, continuous_usage_stats must not leak cached tokens.""" """With cache reporting off, continuous_usage_stats must not leak cached tokens."""
self.tm.server_args.enable_cache_report = False self.enterContext(get_context().override_server_args(enable_cache_report=False))
usages = self._collect_continuous_usage(cached_tokens=6) usages = self._collect_continuous_usage(cached_tokens=6)
self.assertTrue(usages, "continuous_usage_stats attached no usage") self.assertTrue(usages, "continuous_usage_stats attached no usage")
self.assertIsNone(usages[0].get("prompt_tokens_details")) self.assertIsNone(usages[0].get("prompt_tokens_details"))
@@ -3666,7 +3684,9 @@ class ServingChatTestCase(unittest.TestCase):
Regression test for https://github.com/sgl-project/sglang/issues/22510. Regression test for https://github.com/sgl-project/sglang/issues/22510.
""" """
# Enable incremental_streaming_output on the mock # Enable incremental_streaming_output on the mock
self.tm.server_args.incremental_streaming_output = True self.enterContext(
get_context().override_server_args(incremental_streaming_output=True)
)
# Simulate incremental streaming: each yield has ONLY the new text (delta), # Simulate incremental streaming: each yield has ONLY the new text (delta),
# NOT the full accumulated text. # NOT the full accumulated text.
@@ -4146,6 +4166,9 @@ class TestProcessToolCallsWithRequiredToolChoice(unittest.TestCase):
"""Test _process_tool_calls with tool_choice='required' uses model-specific parser.""" """Test _process_tool_calls with tool_choice='required' uses model-specific parser."""
def setUp(self): def setUp(self):
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
tm = _MockTokenizerManager() tm = _MockTokenizerManager()
tm.server_args.tool_call_parser = "kimi_k2" tm.server_args.tool_call_parser = "kimi_k2"
self.chat = OpenAIServingChat(tm, _MockTemplateManager()) self.chat = OpenAIServingChat(tm, _MockTemplateManager())
@@ -19,6 +19,8 @@ from fastapi import Request
from sglang.srt.entrypoints.openai.protocol import CompletionRequest from sglang.srt.entrypoints.openai.protocol import CompletionRequest
from sglang.srt.entrypoints.openai.serving_completions import OpenAIServingCompletion from sglang.srt.entrypoints.openai.serving_completions import OpenAIServingCompletion
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_context, publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import get_or_create_event_loop from sglang.srt.utils import get_or_create_event_loop
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
@@ -66,6 +68,9 @@ class ServingCompletionTestCase(unittest.TestCase):
# ---------- shared test fixtures ---------- # ---------- shared test fixtures ----------
def setUp(self): def setUp(self):
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
# build the mock TokenizerManager once for every test # build the mock TokenizerManager once for every test
tm = Mock(spec=TokenizerManager) tm = Mock(spec=TokenizerManager)
@@ -322,15 +327,18 @@ class ServingCompletionTestCase(unittest.TestCase):
return_token_ids=True, return_token_ids=True,
) )
adapted_request, _ = self.sc._convert_to_internal_request(req) adapted_request, _ = self.sc._convert_to_internal_request(req)
self.sc.tokenizer_manager.server_args.stream_response_default_include_usage = (
False
)
for incremental in (False, True): for incremental in (False, True):
with self.subTest(incremental_streaming_output=incremental): # Both of these are read through `get_serving()` now, so assigning
self.sc.tokenizer_manager.server_args.incremental_streaming_output = ( # them on the mock manager's record has no effect on what the code
incremental # under test sees. State them where the code reads them.
) with (
self.subTest(incremental_streaming_output=incremental),
get_context().override_server_args(
stream_response_default_include_usage=False,
incremental_streaming_output=incremental,
),
):
texts = ("a", "b", "c") if incremental else ("a", "ab", "abc") texts = ("a", "b", "c") if incremental else ("a", "ab", "abc")
output_ids = ( output_ids = (
([5], [6], [7]) if incremental else ([5], [5, 6], [5, 6, 7]) ([5], [6], [7]) if incremental else ([5], [5, 6], [5, 6, 7])
@@ -23,9 +23,11 @@ from sglang.srt.entrypoints.openai.serving_responses import (
) )
from sglang.srt.function_call.core_types import ToolCallItem from sglang.srt.function_call.core_types import ToolCallItem
from sglang.srt.parser.template_detection import ReasoningToggleConfig from sglang.srt.parser.template_detection import ReasoningToggleConfig
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.sampling.sampling_params import ( from sglang.srt.sampling.sampling_params import (
REQUEST_REASONING_END_TOKEN_IDS_KEY, REQUEST_REASONING_END_TOKEN_IDS_KEY,
) )
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -625,6 +627,9 @@ class OutputItemsTestCase(CustomTestCase):
def setUp(self): def setUp(self):
# qwen3_coder is the default for this class; the one no-native-parser # qwen3_coder is the default for this class; the one no-native-parser
# case overrides it. # case overrides it.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
self.serving = make_serving() self.serving = make_serving()
self.serving.tool_call_parser = "qwen3_coder" self.serving.tool_call_parser = "qwen3_coder"
@@ -11,6 +11,8 @@ from utils import (
) )
from sglang.srt.entrypoints.openai.protocol import ResponsesRequest from sglang.srt.entrypoints.openai.protocol import ResponsesRequest
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -246,6 +248,10 @@ class MultiToolCallStreamingOrderTestCase(CustomTestCase):
def setUp(self): def setUp(self):
from sglang.srt.function_call.qwen3_coder_detector import Qwen3CoderDetector from sglang.srt.function_call.qwen3_coder_detector import Qwen3CoderDetector
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
self.serving = make_serving() self.serving = make_serving()
self.serving.tool_call_parser = "qwen3_coder" self.serving.tool_call_parser = "qwen3_coder"
self.serving.reasoning_parser = None self.serving.reasoning_parser = None
@@ -10,7 +10,7 @@ The tests mock ``TokenizerManager.generate_request`` to yield synthetic
``text`` chunks for each of the happy, abort, and boundary cases. ``text`` chunks for each of the happy, abort, and boundary cases.
""" """
from sglang.test.test_utils import maybe_stub_sgl_kernel from sglang.test.test_utils import enter_override, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel() # must precede any import that pulls in sgl_kernel maybe_stub_sgl_kernel() # must precede any import that pulls in sgl_kernel
@@ -32,6 +32,8 @@ from sglang.srt.entrypoints.openai.serving_transcription import (
OpenAIServingTranscription, OpenAIServingTranscription,
) )
from sglang.srt.managers.io_struct import GenerateReqInput from sglang.srt.managers.io_struct import GenerateReqInput
from sglang.srt.runtime_context import get_context, publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import get_or_create_event_loop from sglang.srt.utils import get_or_create_event_loop
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -102,6 +104,16 @@ def _deltas_from_sse(sse_lines: List[str]) -> List[str]:
class TestStreamingFusedAutodetect(CustomTestCase): class TestStreamingFusedAutodetect(CustomTestCase):
"""_generate_transcription_stream with _fused_autodetect=True.""" """_generate_transcription_stream with _fused_autodetect=True."""
def setUp(self):
# The transcription serving layer reads its config from the bags, so
# the fixture publishes one instead of hanging values off a mock.
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(model_path="dummy", asr_max_concurrent_sessions=32),
role="tokenizer",
)
def _run_stream( def _run_stream(
self, chunks: List[dict], fused: bool = True, ts_variant: bool = False self, chunks: List[dict], fused: bool = True, ts_variant: bool = False
): ):
@@ -318,6 +330,16 @@ class TestLongAudioChunkedNonStreaming(CustomTestCase):
requests and the transcripts stitched in order — without chunking the requests and the transcripts stitched in order — without chunking the
feature extractor silently truncates everything past 30 s.""" feature extractor silently truncates everything past 30 s."""
def setUp(self):
# The transcription serving layer reads its config from the bags, so
# the fixture publishes one instead of hanging values off a mock.
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(model_path="dummy", asr_max_concurrent_sessions=32),
role="tokenizer",
)
def _create_transcription(self, tm, audio_bytes, language="en", **kwargs): def _create_transcription(self, tm, audio_bytes, language="en", **kwargs):
serving = OpenAIServingTranscription(tm) serving = OpenAIServingTranscription(tm)
loop = get_or_create_event_loop() loop = get_or_create_event_loop()
@@ -510,6 +532,16 @@ class TestLongAudioChunkedStreaming(CustomTestCase):
"""_generate_long_audio_stream: chunks transcribed sequentially, deltas """_generate_long_audio_stream: chunks transcribed sequentially, deltas
emitted in audio order, exactly one finish frame.""" emitted in audio order, exactly one finish frame."""
def setUp(self):
# The transcription serving layer reads its config from the bags, so
# the fixture publishes one instead of hanging values off a mock.
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(model_path="dummy", asr_max_concurrent_sessions=32),
role="tokenizer",
)
def _run_stream(self, results_per_request, fused=False, n_chunks=2): def _run_stream(self, results_per_request, fused=False, n_chunks=2):
tm = _MockChunkTokenizerManager(results_per_request) tm = _MockChunkTokenizerManager(results_per_request)
serving = OpenAIServingTranscription(tm) serving = OpenAIServingTranscription(tm)
@@ -679,6 +711,16 @@ class TestStreamingIncrementalOutputMode(CustomTestCase):
server already sent as a delta. server already sent as a delta.
""" """
def setUp(self):
# The transcription serving layer reads its config from the bags, so
# the fixture publishes one instead of hanging values off a mock.
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(model_path="dummy", asr_max_concurrent_sessions=32),
role="tokenizer",
)
def _run_incremental_stream(self, chunk_deltas, fused=False): def _run_incremental_stream(self, chunk_deltas, fused=False):
"""Server in incremental mode: yield per-chunk delta, not cumulative.""" """Server in incremental mode: yield per-chunk delta, not cumulative."""
chunks = [ chunks = [
@@ -686,9 +728,8 @@ class TestStreamingIncrementalOutputMode(CustomTestCase):
for i, d in enumerate(chunk_deltas) for i, d in enumerate(chunk_deltas)
] ]
tm = _MockTokenizerManager(chunks) tm = _MockTokenizerManager(chunks)
tm.server_args = Mock( enter_override(
incremental_streaming_output=True, self, get_context().override_server_args(incremental_streaming_output=True)
asr_max_concurrent_sessions=32,
) )
serving = OpenAIServingTranscription(tm) serving = OpenAIServingTranscription(tm)
@@ -27,6 +27,8 @@ from unittest.mock import Mock
from sglang.srt.entrypoints.openai.protocol import RequestResponseMetadata from sglang.srt.entrypoints.openai.protocol import RequestResponseMetadata
from sglang.srt.entrypoints.openai.serving_responses import OpenAIServingResponses from sglang.srt.entrypoints.openai.serving_responses import OpenAIServingResponses
from sglang.srt.runtime_context import get_context, publish
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci( register_cpu_ci(
@@ -83,6 +85,11 @@ class MockTemplateManager:
def make_serving(*, is_multimodal: bool = False) -> OpenAIServingResponses: def make_serving(*, is_multimodal: bool = False) -> OpenAIServingResponses:
"""The serving layer reads its config from the bags, so the fixture
publishes one. Idempotent: a caller that already published keeps its own,
which is how a test states a value the default record does not carry."""
if not get_context().is_config_namespace_published("serving"):
publish(ServerArgs(model_path="dummy"), role="tokenizer")
return OpenAIServingResponses( return OpenAIServingResponses(
MockTokenizerManager(is_multimodal=is_multimodal), MockTemplateManager() MockTokenizerManager(is_multimodal=is_multimodal), MockTemplateManager()
) )
@@ -11,7 +11,13 @@ from types import SimpleNamespace
from sglang.srt.arg_groups.layernorm_sp_hook import validate_layernorm_sp from sglang.srt.arg_groups.layernorm_sp_hook import validate_layernorm_sp
from sglang.srt.layers import layernorm_sp from sglang.srt.layers import layernorm_sp
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.runtime_context import get_flags, get_forward, reset_context from sglang.srt.runtime_context import (
get_flags,
get_forward,
publish,
reset_context,
)
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -19,8 +25,10 @@ register_cpu_ci(est_time=9, suite="base-a-test-cpu")
def _initialize(*, enable=True, arch="Qwen3ForCausalLM"): def _initialize(*, enable=True, arch="Qwen3ForCausalLM"):
publish(
ServerArgs(model_path="dummy", enable_layernorm_sp=enable), role="tokenizer"
)
layernorm_sp.initialize_layernorm_sp( layernorm_sp.initialize_layernorm_sp(
server_args=SimpleNamespace(enable_layernorm_sp=enable),
model_config=SimpleNamespace( model_config=SimpleNamespace(
hf_config=SimpleNamespace(architectures=[arch] if arch else []) hf_config=SimpleNamespace(architectures=[arch] if arch else [])
), ),
@@ -21,6 +21,8 @@ from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.managers.tokenizer_manager_score_mixin import ( from sglang.srt.managers.tokenizer_manager_score_mixin import (
TokenizerManagerScoreMixin, TokenizerManagerScoreMixin,
) )
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -45,6 +47,12 @@ def _vec2d(val: float = 1.0) -> torch.Tensor:
class TestPositionalEmbeds(CustomTestCase): class TestPositionalEmbeds(CustomTestCase):
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def test_from_list_of_1d_tensors(self): def test_from_list_of_1d_tensors(self):
pe = PositionalEmbeds(embeds=[_vec(1), _vec(2)], positions=[0, 5]) pe = PositionalEmbeds(embeds=[_vec(1), _vec(2)], positions=[0, 5])
self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM)) self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM))
@@ -75,6 +83,12 @@ class TestPositionalEmbeds(CustomTestCase):
class TestConvertEmbedsToTensors(CustomTestCase): class TestConvertEmbedsToTensors(CustomTestCase):
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def test_none_returns_none(self): def test_none_returns_none(self):
self.assertIsNone(convert_embeds_to_tensors(None)) self.assertIsNone(convert_embeds_to_tensors(None))
@@ -110,6 +124,12 @@ class TestConvertEmbedsToTensors(CustomTestCase):
class TestResolveEmbedOverrides(CustomTestCase): class TestResolveEmbedOverrides(CustomTestCase):
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def test_basic_resolution(self): def test_basic_resolution(self):
embeds = [_vec(1), _vec(2)] embeds = [_vec(1), _vec(2)]
pe = TokenizerManager._resolve_embed_overrides( pe = TokenizerManager._resolve_embed_overrides(
@@ -144,6 +164,12 @@ class TestResolveEmbedOverrides(CustomTestCase):
class TestGenerateReqInputEmbedOverride(CustomTestCase): class TestGenerateReqInputEmbedOverride(CustomTestCase):
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def test_single_override_in_getitem(self): def test_single_override_in_getitem(self):
"""Single PositionalEmbeds is shared across all items in __getitem__.""" """Single PositionalEmbeds is shared across all items in __getitem__."""
pe = PositionalEmbeds(embeds=[_vec()], positions=[0]) pe = PositionalEmbeds(embeds=[_vec()], positions=[0])
@@ -176,6 +202,12 @@ class TestGenerateReqInputEmbedOverride(CustomTestCase):
class TestEmbeddingReqInputEmbedOverride(CustomTestCase): class TestEmbeddingReqInputEmbedOverride(CustomTestCase):
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def test_override_fields_in_getitem(self): def test_override_fields_in_getitem(self):
"""embed_override_token_id, embed_overrides, and positional_embed_overrides """embed_override_token_id, embed_overrides, and positional_embed_overrides
are correctly sliced in __getitem__.""" are correctly sliced in __getitem__."""
@@ -220,6 +252,9 @@ class _FakeMixin(TokenizerManagerScoreMixin):
class TestResolveOverridesForSequence(CustomTestCase): class TestResolveOverridesForSequence(CustomTestCase):
def setUp(self): def setUp(self):
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
self.mixin = _FakeMixin() self.mixin = _FakeMixin()
def test_none_embeds_returns_empty(self): def test_none_embeds_returns_empty(self):
@@ -276,6 +311,9 @@ class TestResolveOverridesForSequence(CustomTestCase):
class TestResolveEmbedOverridesForRequest(CustomTestCase): class TestResolveEmbedOverridesForRequest(CustomTestCase):
def setUp(self): def setUp(self):
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
self.mixin = _FakeMixin() self.mixin = _FakeMixin()
def test_no_overrides_returns_none(self): def test_no_overrides_returns_none(self):
@@ -339,6 +377,9 @@ DELIM_TOKEN = MIS_DELIMITER_TOKEN_ID
class TestBuildTokenIdInputs(CustomTestCase): class TestBuildTokenIdInputs(CustomTestCase):
def setUp(self): def setUp(self):
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
self.mixin = _FakeMixin(enable_mis=True) self.mixin = _FakeMixin(enable_mis=True)
# --- single-item mode, no embeds --- # --- single-item mode, no embeds ---
@@ -505,6 +546,9 @@ class TestScoreRequestValidation(CustomTestCase):
"""Test validation guards in score_request without running full pipeline.""" """Test validation guards in score_request without running full pipeline."""
def setUp(self): def setUp(self):
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
self.mixin = _FakeMixin() self.mixin = _FakeMixin()
def _call(self, **kwargs): def _call(self, **kwargs):
@@ -15,6 +15,8 @@ from types import SimpleNamespace
import torch import torch
from sglang.srt.managers.schedule_batch import ReqKvInfo from sglang.srt.managers.schedule_batch import ReqKvInfo
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
from sglang.srt.utils.common import Range from sglang.srt.utils.common import Range
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
@@ -165,7 +167,14 @@ class TestHiSparseUnit(unittest.TestCase):
Without this, a mid-test assertion failure skips cleanup and leaks Without this, a mid-test assertion failure skips cleanup and leaks
resources, causing unrelated failures in later tests. resources, causing unrelated failures in later tests.
The code under test reads its configuration from the bags -- the PD
decode prealloc path asks whether the decode radix cache is on -- so a
case here needs a published config, the way a real process has one.
""" """
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="scheduler")
self.allocator.clear() self.allocator.clear()
self.req_to_token_pool.clear() self.req_to_token_pool.clear()
self.coordinator.mem_pool_host.clear() self.coordinator.mem_pool_host.clear()
@@ -758,7 +767,6 @@ class TestHiSparseUnit(unittest.TestCase):
queue.scheduler = SimpleNamespace( queue.scheduler = SimpleNamespace(
enable_hisparse=True, enable_hisparse=True,
hisparse_coordinator=self.coordinator, hisparse_coordinator=self.coordinator,
server_args=SimpleNamespace(disaggregation_decode_enable_radix_cache=False),
) )
host_indices = queue._pre_alloc(req) host_indices = queue._pre_alloc(req)
@@ -19,13 +19,20 @@ from sglang.srt.managers.schedule_batch import ( # noqa: E402
ReqKvInfo, ReqKvInfo,
) )
from sglang.srt.managers.scheduler import Scheduler # noqa: E402 from sglang.srt.managers.scheduler import Scheduler # noqa: E402
from sglang.srt.runtime_context import get_context # noqa: E402 from sglang.srt.runtime_context import get_context, publish, reset_context # noqa: E402
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=12, suite="base-a-test-cpu") register_cpu_ci(est_time=12, suite="base-a-test-cpu")
class TestDisaggregationPriorityQueueing(unittest.TestCase): class TestDisaggregationPriorityQueueing(unittest.TestCase):
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def _new_scheduler(self, disaggregation_mode: DisaggregationMode) -> Scheduler: def _new_scheduler(self, disaggregation_mode: DisaggregationMode) -> Scheduler:
scheduler = Scheduler.__new__(Scheduler) scheduler = Scheduler.__new__(Scheduler)
scheduler.disaggregation_mode = disaggregation_mode scheduler.disaggregation_mode = disaggregation_mode
@@ -91,6 +98,12 @@ class TestDisaggregationPriorityQueueing(unittest.TestCase):
class TestDecodePreallocQueuePriority(unittest.TestCase): class TestDecodePreallocQueuePriority(unittest.TestCase):
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def _new_decode_req(self, rid: str, priority: int, *, failed: bool = False): def _new_decode_req(self, rid: str, priority: int, *, failed: bool = False):
req = SimpleNamespace( req = SimpleNamespace(
rid=rid, rid=rid,
@@ -227,6 +240,12 @@ class TestDecodePreallocQueueRebootstrapPayload(unittest.TestCase):
dispatch itself now lives on the kv manager (see dispatch itself now lives on the kv manager (see
``TestCommonKVManagerPrefillRecompute``).""" ``TestCommonKVManagerPrefillRecompute``)."""
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def _sampling_params(self): def _sampling_params(self):
return SimpleNamespace( return SimpleNamespace(
temperature=0.0, temperature=0.0,
@@ -286,6 +305,12 @@ class TestCommonKVManagerPrefillRecompute(unittest.TestCase):
``KVPoll.Failed`` so the scheduler's normal transfer-failure streaming runs. ``KVPoll.Failed`` so the scheduler's normal transfer-failure streaming runs.
""" """
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def _new_manager(self): def _new_manager(self):
from sglang.srt.disaggregation.common.conn import CommonKVManager from sglang.srt.disaggregation.common.conn import CommonKVManager
@@ -423,6 +448,12 @@ class TestCommonKVManagerPrefillRecompute(unittest.TestCase):
class TestDecodePrebuilt(unittest.TestCase): class TestDecodePrebuilt(unittest.TestCase):
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def _new_scheduler(self, *, enable_overlap: bool) -> Scheduler: def _new_scheduler(self, *, enable_overlap: bool) -> Scheduler:
scheduler = Scheduler.__new__(Scheduler) scheduler = Scheduler.__new__(Scheduler)
scheduler.grammar_manager = MagicMock() scheduler.grammar_manager = MagicMock()
@@ -10,7 +10,8 @@ from sglang.srt.managers.schedule_batch import ReqKvInfo
from sglang.srt.mem_cache.allocator.hisparse import ( from sglang.srt.mem_cache.allocator.hisparse import (
DeepSeekV4HiSparseTokenToKVPoolAllocator, DeepSeekV4HiSparseTokenToKVPoolAllocator,
) )
from sglang.srt.runtime_context import get_context from sglang.srt.runtime_context import get_context, publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -18,6 +19,12 @@ register_cpu_ci(est_time=10, suite="base-a-test-cpu")
class TestDeepSeekV4HiSparseAllocator(CustomTestCase): class TestDeepSeekV4HiSparseAllocator(CustomTestCase):
def setUp(self):
# The code under test reads its config from the bags.
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="tokenizer")
def test_forwards_swa_tail_allocation_to_logical_allocator(self): def test_forwards_swa_tail_allocation_to_logical_allocator(self):
allocator = object.__new__(DeepSeekV4HiSparseTokenToKVPoolAllocator) allocator = object.__new__(DeepSeekV4HiSparseTokenToKVPoolAllocator)
logical_allocator = MagicMock(spec=["alloc_extend_swa_tail"]) logical_allocator = MagicMock(spec=["alloc_extend_swa_tail"])
@@ -12,6 +12,7 @@ from sglang.srt.multimodal.processors.qwen_vl import QwenVLImageProcessor
from sglang.srt.multimodal.transport.cuda_ipc import ( from sglang.srt.multimodal.transport.cuda_ipc import (
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY, DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
) )
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -33,6 +34,15 @@ class _RecordingVisual:
class TestQwen3VLFeatureMaterialization(CustomTestCase): class TestQwen3VLFeatureMaterialization(CustomTestCase):
def setUp(self):
# The transport decision is read from the `mm` bag.
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="test")
@staticmethod @staticmethod
def _model(visual, *, use_data_parallel): def _model(visual, *, use_data_parallel):
model = Qwen3VLForConditionalGeneration.__new__(Qwen3VLForConditionalGeneration) model = Qwen3VLForConditionalGeneration.__new__(Qwen3VLForConditionalGeneration)
@@ -43,10 +53,15 @@ class TestQwen3VLFeatureMaterialization(CustomTestCase):
def test_processor_defers_gpu_transport_for_encoder_dp(self): def test_processor_defers_gpu_transport_for_encoder_dp(self):
for transport in ("cuda_ipc", "cuda_vmm"): for transport in ("cuda_ipc", "cuda_vmm"):
with self.subTest(transport=transport): # `mm_enable_dp_encoder` is read through `get_mm()` now, so stating
# it on the processor's own `server_args` no longer reaches the
# code under test.
with (
self.subTest(transport=transport),
get_context().override_server_args(mm_enable_dp_encoder=True),
):
processor = QwenVLImageProcessor.__new__(QwenVLImageProcessor) processor = QwenVLImageProcessor.__new__(QwenVLImageProcessor)
processor.mm_feature_transport = transport processor.mm_feature_transport = transport
processor.server_args = SimpleNamespace(mm_enable_dp_encoder=True)
processor.model_type = "qwen3_vl" processor.model_type = "qwen3_vl"
items = [ items = [
MultimodalDataItem(modality=Modality.IMAGE), MultimodalDataItem(modality=Modality.IMAGE),
@@ -94,6 +94,11 @@ def make_processor(case, config, image_processor_cls=None):
publish( publish(
ServerArgs( ServerArgs(
model_path="dummy", model_path="dummy",
# Mirrored for the same reason the stub sets it: `get_mm_processor_cls`
# reads `model_impl` from this bag now, and "auto" would send it into
# `get_resolved_model_impl`, which chokes on the SimpleNamespace
# `model_config` these tests hand it.
model_impl=server_args.model_impl,
mm_feature_transport=server_args.mm_feature_transport, mm_feature_transport=server_args.mm_feature_transport,
mm_process_config=server_args.mm_process_config, mm_process_config=server_args.mm_process_config,
allowed_media_domains=server_args.allowed_media_domains, allowed_media_domains=server_args.allowed_media_domains,
@@ -20,7 +20,9 @@ from sglang.srt.managers.multimodal_processor import ( # noqa: E402
get_mm_processor_cls, get_mm_processor_cls,
import_processors, import_processors,
) )
from sglang.srt.runtime_context import publish
from sglang.srt.rust_server.multimodal import rust_mm_family_for # noqa: E402 from sglang.srt.rust_server.multimodal import rust_mm_family_for # noqa: E402
from sglang.srt.server_args import ServerArgs
register_cpu_ci(est_time=14, suite="base-a-test-cpu") register_cpu_ci(est_time=14, suite="base-a-test-cpu")
@@ -28,7 +30,10 @@ register_cpu_ci(est_time=14, suite="base-a-test-cpu")
def processor_cls_for(architecture, model_type): def processor_cls_for(architecture, model_type):
"""Through the production selection, as `resolve_spec` calls it.""" """Through the production selection, as `resolve_spec` calls it."""
hf_config = SimpleNamespace(architectures=[architecture], model_type=model_type) hf_config = SimpleNamespace(architectures=[architecture], model_type=model_type)
return get_mm_processor_cls(hf_config, SimpleNamespace(model_impl="sglang")) # `model_impl` is read from the bags now, so it has to be published rather
# than handed over on a stand-in.
publish(ServerArgs(model_path="dummy", model_impl="sglang"), role="tokenizer")
return get_mm_processor_cls(hf_config)
class TestRustMmGate(CustomTestCase): class TestRustMmGate(CustomTestCase):
@@ -12,6 +12,7 @@ from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import patch
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -31,12 +32,25 @@ class _StubProcessor(BaseMultimodalProcessor):
def _make(**fields): def _make(**fields):
"""Both surfaces, because the device decision reads both.
`base_gpu_id` is the instance's own -- two engines in one process keep
different ones, which `test_publishing_another_config_does_not_move_the_device`
pins -- so it stays on the record the processor holds. `rl_on_policy_target`
is the process's, so it is published.
"""
server_args = ServerArgs(model_path="dummy", **fields)
publish(server_args, role="tokenizer")
processor = _StubProcessor.__new__(_StubProcessor) processor = _StubProcessor.__new__(_StubProcessor)
processor.server_args = ServerArgs(model_path="dummy", **fields) processor.server_args = server_args
return processor return processor
class TestFastImageProcessorDevice(CustomTestCase): class TestFastImageProcessorDevice(CustomTestCase):
def setUp(self):
reset_context()
self.addCleanup(reset_context)
def _device(self, processor, **platform): def _device(self, processor, **platform):
flags = {"_is_cpu": False, "_is_xpu": False, "_is_npu": False} flags = {"_is_cpu": False, "_is_xpu": False, "_is_npu": False}
flags.update(platform) flags.update(platform)
@@ -80,6 +94,10 @@ class TestFastImageProcessorDevice(CustomTestCase):
class TestFastImageProcessorMemoryPool(CustomTestCase): class TestFastImageProcessorMemoryPool(CustomTestCase):
def setUp(self):
reset_context()
self.addCleanup(reset_context)
def _processor(self, *, transport="cpu", precompute_hash=False): def _processor(self, *, transport="cpu", precompute_hash=False):
processor = _make(base_gpu_id=0) processor = _make(base_gpu_id=0)
processor.mm_feature_transport = transport processor.mm_feature_transport = transport