[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:
@@ -14,6 +14,7 @@ from sglang.srt.constrained.base_grammar_backend import (
|
||||
from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject
|
||||
from sglang.srt.distributed.communication_tags import P2PTag
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_serving
|
||||
from sglang.srt.sampling.sampling_params import (
|
||||
get_request_reasoning_end_token_ids,
|
||||
)
|
||||
@@ -31,7 +32,7 @@ class GrammarManager:
|
||||
self.scheduler = scheduler
|
||||
self.server_args = scheduler.server_args
|
||||
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.server_args,
|
||||
scheduler.tokenizer,
|
||||
|
||||
@@ -468,11 +468,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
return seq_len
|
||||
|
||||
page_size = self.token_to_kv_pool_allocator.page_size
|
||||
if getattr(
|
||||
self.scheduler.server_args,
|
||||
"disaggregation_decode_enable_radix_cache",
|
||||
False,
|
||||
):
|
||||
if get_disagg().disaggregation_decode_enable_radix_cache:
|
||||
# Keep enough SWA before the page-aligned radix-cache insert
|
||||
# boundary for the cached key to contain a complete window.
|
||||
# `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)
|
||||
prefix_match: Optional[DecodePrefixMatch] = None
|
||||
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
|
||||
)
|
||||
if use_decode_radix_cache:
|
||||
@@ -1648,13 +1644,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
available_size = logical_allocator.available_size()
|
||||
elif self._uses_swa_tail_prealloc():
|
||||
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()
|
||||
else:
|
||||
available_size = self.token_to_kv_pool_allocator.available_size()
|
||||
# Include evictable decode-radix cache entries in the budget -- they
|
||||
# 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()
|
||||
allocatable_tokens = available_size - max(
|
||||
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.
|
||||
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
|
||||
):
|
||||
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()
|
||||
port_args = PortArgs.init_new(server_args)
|
||||
|
||||
if server_args.dist_init_addr:
|
||||
na = NetworkAddress.parse(server_args.dist_init_addr)
|
||||
if get_parallel().dist_init_addr:
|
||||
na = NetworkAddress.parse(get_parallel().dist_init_addr)
|
||||
dist_init_method = na.to_tcp()
|
||||
else:
|
||||
dist_init_method = NetworkAddress(
|
||||
|
||||
@@ -103,7 +103,7 @@ async def _lifespan(app: FastAPI):
|
||||
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.
|
||||
|
||||
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
|
||||
if not host or host in ("0.0.0.0", "::"):
|
||||
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)
|
||||
payload = {"url": encoder_url}
|
||||
bootstrap_urls = list(server_args.encoder_register_urls)
|
||||
bootstrap_urls = list(get_disagg().encoder_register_urls)
|
||||
if not bootstrap_urls:
|
||||
return
|
||||
|
||||
@@ -175,15 +175,15 @@ def _register_encoder_url_with_bootstrap(server_args: ServerArgs):
|
||||
).start()
|
||||
|
||||
|
||||
def _unregister_encoder_url_from_bootstrap(server_args: ServerArgs):
|
||||
def _unregister_encoder_url_from_bootstrap():
|
||||
host = get_serving().host
|
||||
if not host or host in ("0.0.0.0", "::"):
|
||||
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)
|
||||
payload = {"url": encoder_url}
|
||||
|
||||
for bootstrap_url in server_args.encoder_register_urls:
|
||||
for bootstrap_url in get_disagg().encoder_register_urls:
|
||||
try:
|
||||
resp = http_requests.delete(
|
||||
f"{bootstrap_url}/unregister_encoder_url",
|
||||
@@ -229,8 +229,8 @@ def launch_server(server_args: ServerArgs):
|
||||
if get_disagg().encoder_register_urls:
|
||||
import atexit
|
||||
|
||||
_register_encoder_url_with_bootstrap(server_args)
|
||||
atexit.register(_unregister_encoder_url_from_bootstrap, server_args)
|
||||
_register_encoder_url_with_bootstrap()
|
||||
atexit.register(_unregister_encoder_url_from_bootstrap)
|
||||
|
||||
uvicorn.run(app, host=get_serving().host, port=get_serving().port)
|
||||
|
||||
|
||||
@@ -2000,7 +2000,7 @@ class MMReceiverBase(ABC):
|
||||
import_processors("sglang.srt.multimodal.processors")
|
||||
|
||||
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
|
||||
|
||||
_processor = get_processor(
|
||||
|
||||
@@ -52,6 +52,7 @@ from sglang.srt.observability.trace import (
|
||||
trace_set_thread_info,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_observability,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
@@ -1867,7 +1868,7 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher:
|
||||
atexit.register(_kill_workers)
|
||||
|
||||
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
|
||||
# scheduler launcher and DP controller) so the child inherits
|
||||
# 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)
|
||||
|
||||
labels = {"model_name": get_serving().served_model_name}
|
||||
if server_args.extra_metric_labels:
|
||||
labels.update(server_args.extra_metric_labels)
|
||||
if get_observability().extra_metric_labels:
|
||||
labels.update(get_observability().extra_metric_labels)
|
||||
return DPDispatcher(
|
||||
dp_size,
|
||||
dispatch_sockets,
|
||||
|
||||
@@ -548,7 +548,7 @@ class MMEncoder:
|
||||
self.server_args = server_args
|
||||
configure_media_url_security(
|
||||
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.use_mooncake = self.transfer_backend == "mooncake"
|
||||
@@ -563,18 +563,18 @@ class MMEncoder:
|
||||
)
|
||||
self.load_config = LoadConfig(
|
||||
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,
|
||||
remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip,
|
||||
remote_instance_weight_loader_seed_instance_service_port=server_args.remote_instance_weight_loader_seed_instance_service_port,
|
||||
remote_instance_weight_loader_send_weights_group_ports=server_args.remote_instance_weight_loader_send_weights_group_ports,
|
||||
remote_instance_weight_loader_seed_instance_ip=get_model().remote_instance_weight_loader_seed_instance_ip,
|
||||
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=get_model().remote_instance_weight_loader_send_weights_group_ports,
|
||||
)
|
||||
self.model_type = getattr(
|
||||
self.model_config.hf_config, "model_type", "unknown"
|
||||
).lower()
|
||||
|
||||
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(
|
||||
device=self.device,
|
||||
|
||||
@@ -81,14 +81,14 @@ def init_torch_distributed(
|
||||
tic = time.perf_counter()
|
||||
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)
|
||||
if not get_parallel().enable_p2p_check:
|
||||
monkey_patch_p2p_access_check()
|
||||
|
||||
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 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)
|
||||
if device == "cuda" and server_args.elastic_ep_backend == "mooncake":
|
||||
if device == "cuda" and get_exec().moe.elastic_ep_backend == "mooncake":
|
||||
backend = "mooncake"
|
||||
return backend
|
||||
|
||||
@@ -198,9 +198,9 @@ def _resolve_dist_init_method(*, dist_port: int) -> str:
|
||||
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_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_flashinfer_allreduce_only(
|
||||
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_attn_cp_group=(
|
||||
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
|
||||
),
|
||||
enable_symm_mem=get_exec().comm.enable_symm_mem,
|
||||
@@ -308,10 +308,7 @@ def _init_parallel_groups(
|
||||
server_args=server_args,
|
||||
model_config=model_config,
|
||||
)
|
||||
initialize_layernorm_sp(
|
||||
server_args=server_args,
|
||||
model_config=model_config,
|
||||
)
|
||||
initialize_layernorm_sp(model_config=model_config)
|
||||
if is_npu():
|
||||
register_sgl_tp_rank(gpu_id)
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ from http import HTTPStatus
|
||||
from fastapi import APIRouter, Request
|
||||
from fastapi.responses import ORJSONResponse
|
||||
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils.auth import AuthLevel, auth_level
|
||||
|
||||
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.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(
|
||||
{"error": "elastic EP is not enabled (set --elastic-ep-backend)"},
|
||||
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."""
|
||||
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(
|
||||
{"error": "elastic EP is not enabled (set --elastic-ep-backend)"},
|
||||
status_code=HTTPStatus.NOT_FOUND,
|
||||
|
||||
@@ -104,9 +104,11 @@ from sglang.srt.parser.template_manager import TemplateManager
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.plugins import load_plugins
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_model,
|
||||
get_observability,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
publish,
|
||||
@@ -315,7 +317,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
|
||||
# Initialize ZMQ sockets
|
||||
context = zmq.Context(2)
|
||||
if self.server_args.node_rank == 0:
|
||||
if get_parallel().node_rank == 0:
|
||||
self.send_to_rpc = get_zmq_socket(
|
||||
context, zmq.DEALER, self.port_args.rpc_ipc_name, True
|
||||
)
|
||||
@@ -323,11 +325,11 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
self.send_to_rpc = None
|
||||
|
||||
# Enable tracing
|
||||
if server_args.enable_trace:
|
||||
if get_observability().enable_trace:
|
||||
process_tracing_init(
|
||||
server_args.otlp_traces_endpoint,
|
||||
get_observability().otlp_traces_endpoint,
|
||||
"sglang",
|
||||
trace_modules=server_args.trace_modules,
|
||||
trace_modules=get_observability().trace_modules,
|
||||
)
|
||||
thread_label = "Tokenizer"
|
||||
if get_disagg().disaggregation_mode == "prefill":
|
||||
@@ -691,27 +693,27 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
# 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
|
||||
# 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(
|
||||
"Multi-node weight cache daemons (nnodes > 1) require "
|
||||
"--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 = (
|
||||
_calculate_rank_ranges(
|
||||
server_args.nnodes,
|
||||
get_parallel().nnodes,
|
||||
get_parallel().pp_size,
|
||||
tp_size,
|
||||
server_args.node_rank,
|
||||
get_parallel().node_rank,
|
||||
)
|
||||
)
|
||||
|
||||
# Build the distributed init method (multi-node uses the user-provided
|
||||
# dist_init_addr so all nodes reach the same endpoint).
|
||||
if server_args.dist_init_addr:
|
||||
host, port = server_args.dist_init_addr.rsplit(":", 1)
|
||||
if get_parallel().dist_init_addr:
|
||||
host, port = get_parallel().dist_init_addr.rsplit(":", 1)
|
||||
dist_init_method = f"tcp://{host}:{port}"
|
||||
else:
|
||||
# Fresh free port for the daemons' own rendezvous, not the engine's
|
||||
@@ -723,7 +725,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
daemon_procs = []
|
||||
logger.info(
|
||||
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"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, "
|
||||
f"dist_init_method={dist_init_method}"
|
||||
@@ -738,8 +740,8 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
tp_rank,
|
||||
pp_size_per_node,
|
||||
tp_size_per_node,
|
||||
base_gpu_id=server_args.base_gpu_id,
|
||||
gpu_id_step=server_args.gpu_id_step,
|
||||
base_gpu_id=get_device().base_gpu_id,
|
||||
gpu_id_step=get_device().gpu_id_step,
|
||||
)
|
||||
cleanup_stale_daemon_files(current_platform.get_device_uuid(gpu_id))
|
||||
|
||||
@@ -750,8 +752,8 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
tp_rank,
|
||||
pp_size_per_node,
|
||||
tp_size_per_node,
|
||||
base_gpu_id=server_args.base_gpu_id,
|
||||
gpu_id_step=server_args.gpu_id_step,
|
||||
base_gpu_id=get_device().base_gpu_id,
|
||||
gpu_id_step=get_device().gpu_id_step,
|
||||
)
|
||||
proc = spawn_weight_cache_daemon(
|
||||
server_args,
|
||||
@@ -767,7 +769,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
# (readiness timeout or a daemon exiting early) terminate the siblings
|
||||
# we already spawned before propagating, so a partial launch does not
|
||||
# leak GPU-resident daemons.
|
||||
timeout = server_args.weight_cache_timeout
|
||||
timeout = get_model().weight_cache_timeout
|
||||
check_interval = 2
|
||||
start_time = time.time()
|
||||
try:
|
||||
@@ -778,8 +780,8 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
tp_rank,
|
||||
pp_size_per_node,
|
||||
tp_size_per_node,
|
||||
base_gpu_id=server_args.base_gpu_id,
|
||||
gpu_id_step=server_args.gpu_id_step,
|
||||
base_gpu_id=get_device().base_gpu_id,
|
||||
gpu_id_step=get_device().gpu_id_step,
|
||||
)
|
||||
ready_path = get_ready_path(
|
||||
current_platform.get_device_uuid(gpu_id)
|
||||
@@ -809,7 +811,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
|
||||
logger.info(
|
||||
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
|
||||
|
||||
@@ -864,16 +866,16 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
if not use_dp_controller:
|
||||
# Launch tensor parallel scheduler processes
|
||||
memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
enable=server_args.enable_memory_saver
|
||||
enable=get_exec().features.enable_memory_saver
|
||||
)
|
||||
scheduler_pipe_readers = []
|
||||
|
||||
pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = (
|
||||
_calculate_rank_ranges(
|
||||
server_args.nnodes,
|
||||
get_parallel().nnodes,
|
||||
get_parallel().pp_size,
|
||||
server_args.tp_size,
|
||||
server_args.node_rank,
|
||||
get_parallel().tp_size,
|
||||
get_parallel().node_rank,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -881,12 +883,12 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
for tp_rank in tp_rank_range:
|
||||
reader, writer = mp.Pipe(duplex=False)
|
||||
gpu_id = (
|
||||
server_args.base_gpu_id
|
||||
get_device().base_gpu_id
|
||||
+ ((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(
|
||||
server_args, tp_rank
|
||||
tp_rank
|
||||
)
|
||||
|
||||
with maybe_reindex_device_id(gpu_id) as gpu_id:
|
||||
@@ -1109,9 +1111,9 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
engine_info_bootstrap_server = None
|
||||
if (
|
||||
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):
|
||||
raise RuntimeError(
|
||||
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."
|
||||
)
|
||||
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
|
||||
# Engine that spawned them; shutdown() reaps from there.
|
||||
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(
|
||||
server_args
|
||||
)
|
||||
@@ -1147,12 +1149,12 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
)
|
||||
|
||||
if (
|
||||
server_args.enable_elastic_expert_backup
|
||||
and server_args.elastic_ep_backend is not None
|
||||
get_exec().moe.enable_elastic_expert_backup
|
||||
and get_exec().moe.elastic_ep_backend is not None
|
||||
):
|
||||
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.
|
||||
scheduler_init_result.wait_for_ready()
|
||||
|
||||
@@ -1168,7 +1170,9 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
)
|
||||
|
||||
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()
|
||||
@@ -1215,7 +1219,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
scheduler_init_result.all_child_pids.append(p.pid)
|
||||
|
||||
# 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(
|
||||
server_args, port_args
|
||||
)
|
||||
@@ -1534,7 +1538,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
else:
|
||||
return [
|
||||
MultiprocessingSerializer.serialize(tensors)
|
||||
for _ in range(self.server_args.tp_size)
|
||||
for _ in range(get_parallel().tp_size)
|
||||
]
|
||||
|
||||
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
|
||||
|
||||
|
||||
def _compute_parallelism_ranks(
|
||||
server_args: ServerArgs, tp_rank: int
|
||||
) -> Tuple[int, int, int]:
|
||||
def _compute_parallelism_ranks(tp_rank: int) -> Tuple[int, int, int]:
|
||||
"""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
|
||||
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
|
||||
tp_size = server_args.tp_size
|
||||
tp_size = get_parallel().tp_size
|
||||
attn_cp_size = get_parallel().attn_cp_size
|
||||
moe_dp_size = get_parallel().moe_dp_size
|
||||
|
||||
|
||||
@@ -176,9 +176,9 @@ async def serve_grpc(server_args, model_info=None):
|
||||
sidecar_app = web.Application()
|
||||
sidecar_runner = None
|
||||
sidecar_port = (
|
||||
server_args.smg_http_sidecar_port
|
||||
if server_args.smg_http_sidecar_port is not None
|
||||
else server_args.port + 1
|
||||
cfg.smg_http_sidecar_port
|
||||
if cfg.smg_http_sidecar_port is not None
|
||||
else cfg.port + 1
|
||||
)
|
||||
|
||||
# Metrics setup: must set PROMETHEUS_MULTIPROC_DIR before scheduler
|
||||
|
||||
@@ -251,8 +251,8 @@ async def init_multi_tokenizer() -> ServerArgs:
|
||||
template_manager.initialize_templates(
|
||||
tokenizer_manager=tokenizer_manager,
|
||||
model_path=get_model().model_path,
|
||||
chat_template=server_args.chat_template,
|
||||
completion_template=server_args.completion_template,
|
||||
chat_template=get_serving().chat_template,
|
||||
completion_template=get_serving().completion_template,
|
||||
)
|
||||
|
||||
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()
|
||||
|
||||
# Init tracing
|
||||
if server_args.enable_trace:
|
||||
if get_observability().enable_trace:
|
||||
process_tracing_init(
|
||||
server_args.otlp_traces_endpoint,
|
||||
get_observability().otlp_traces_endpoint,
|
||||
"sglang",
|
||||
trace_modules=server_args.trace_modules,
|
||||
trace_modules=get_observability().trace_modules,
|
||||
)
|
||||
if get_disagg().disaggregation_mode == "prefill":
|
||||
thread_label = "Prefill" + thread_label
|
||||
@@ -343,15 +343,15 @@ async def lifespan(fast_api_app: FastAPI):
|
||||
|
||||
# Launch tool server
|
||||
tool_server = None
|
||||
if server_args.tool_server == "demo":
|
||||
if get_serving().tool_server == "demo":
|
||||
from sglang.srt.entrypoints.openai.tool_server import DemoToolServer
|
||||
|
||||
tool_server = DemoToolServer()
|
||||
elif server_args.tool_server:
|
||||
elif get_serving().tool_server:
|
||||
from sglang.srt.entrypoints.openai.tool_server import 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():
|
||||
from sglang.srt.entrypoints.openai.tool_server import NativeToolServer
|
||||
|
||||
@@ -381,10 +381,10 @@ async def lifespan(fast_api_app: FastAPI):
|
||||
)
|
||||
|
||||
# Execute custom warmups
|
||||
if server_args.warmups is not None:
|
||||
if get_serving().warmups is not None:
|
||||
await execute_warmups(
|
||||
get_disagg().disaggregation_mode,
|
||||
server_args.warmups.split(","),
|
||||
get_serving().warmups.split(","),
|
||||
_global_state.tokenizer_manager,
|
||||
)
|
||||
logger.info("Warmup ended")
|
||||
@@ -397,7 +397,7 @@ async def lifespan(fast_api_app: FastAPI):
|
||||
if (
|
||||
getattr(fast_api_app, "is_single_tokenizer_mode", False)
|
||||
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(
|
||||
server_args=server_args,
|
||||
@@ -406,7 +406,7 @@ async def lifespan(fast_api_app: FastAPI):
|
||||
scheduler_info=_global_state.scheduler_info,
|
||||
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
|
||||
|
||||
sidecar = start_sidecar()
|
||||
@@ -1101,7 +1101,7 @@ async def attach_hicache_storage_backend(
|
||||
|
||||
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()
|
||||
|
||||
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.
|
||||
"""
|
||||
if not _global_state.tokenizer_manager.server_args.admin_api_key:
|
||||
if not get_serving().admin_api_key:
|
||||
return _admin_api_key_missing_response()
|
||||
|
||||
ret = await _global_state.tokenizer_manager.detach_hicache_storage()
|
||||
@@ -1157,7 +1157,7 @@ async def detach_hicache_storage_backend():
|
||||
@auth_level(AuthLevel.ADMIN_OPTIONAL)
|
||||
async def hicache_storage_backend_status():
|
||||
"""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 {
|
||||
@@ -2203,8 +2203,8 @@ async def _send_disaggregation_warmup_requests(
|
||||
def _execute_server_warmup(server_args: ServerArgs):
|
||||
headers = {}
|
||||
url = server_args.url()
|
||||
if server_args.api_key:
|
||||
headers["Authorization"] = f"Bearer {server_args.api_key}"
|
||||
if get_serving().api_key:
|
||||
headers["Authorization"] = f"Bearer {get_serving().api_key}"
|
||||
|
||||
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.
|
||||
is_vlm = (
|
||||
bool(model_info.get("has_image_understanding", False))
|
||||
and not server_args.language_only
|
||||
and not server_args.language_model_only
|
||||
and not get_disagg().language_only
|
||||
and not get_disagg().language_model_only
|
||||
and not is_mps()
|
||||
)
|
||||
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"
|
||||
else:
|
||||
request_name = "/generate"
|
||||
@@ -2253,7 +2253,7 @@ def _execute_server_warmup(server_args: ServerArgs):
|
||||
"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)]
|
||||
# TODO Workaround the bug that embedding errors for list of 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]
|
||||
|
||||
# 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["input_ids"] = np.load(
|
||||
server_args.debug_tensor_dump_input_file
|
||||
get_observability().debug_tensor_dump_input_file
|
||||
).tolist()
|
||||
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):
|
||||
# Freeze GC after server warmup so static objects skip future GC gen2 collection.
|
||||
# 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 = {}
|
||||
if freeze_key:
|
||||
freeze_headers["Authorization"] = f"Bearer {freeze_key}"
|
||||
@@ -2394,7 +2394,7 @@ def _wait_and_warmup(
|
||||
launch_callback: Optional[Callable[[], None]] = None,
|
||||
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()
|
||||
|
||||
# Joiner schedulers are served through the primary after adoption.
|
||||
@@ -2416,10 +2416,10 @@ def _wait_and_warmup(
|
||||
# The server is ready for requests
|
||||
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)
|
||||
|
||||
if server_args.debug_tensor_dump_input_file:
|
||||
if get_observability().debug_tensor_dump_input_file:
|
||||
kill_process_tree(os.getpid())
|
||||
|
||||
if launch_callback is not None:
|
||||
@@ -2559,7 +2559,7 @@ def _setup_and_run_http_server(
|
||||
|
||||
# Pass additional arguments to the lifespan function.
|
||||
# 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.
|
||||
app.is_single_tokenizer_mode = True
|
||||
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
|
||||
# admin_api_key is not configured.
|
||||
if (
|
||||
server_args.api_key
|
||||
or server_args.admin_api_key
|
||||
get_serving().api_key
|
||||
or get_serving().admin_api_key
|
||||
or app_has_admin_force_endpoints(app)
|
||||
):
|
||||
from sglang.srt.utils.auth import add_api_key_middleware
|
||||
|
||||
add_api_key_middleware(
|
||||
app,
|
||||
api_key=server_args.api_key,
|
||||
admin_api_key=server_args.admin_api_key,
|
||||
api_key=get_serving().api_key,
|
||||
admin_api_key=get_serving().admin_api_key,
|
||||
)
|
||||
else:
|
||||
# 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
|
||||
set_uvicorn_logging_configs(server_args)
|
||||
|
||||
if server_args.ssl_certfile:
|
||||
if get_serving().ssl_certfile:
|
||||
logger.info(
|
||||
f"SSL enabled: certfile={server_args.ssl_certfile}, "
|
||||
f"keyfile={server_args.ssl_keyfile}"
|
||||
f"SSL enabled: certfile={get_serving().ssl_certfile}, "
|
||||
f"keyfile={get_serving().ssl_keyfile}"
|
||||
)
|
||||
|
||||
# Listen for HTTP requests
|
||||
if server_args.tokenizer_worker_num == 1:
|
||||
if server_args.enable_http2:
|
||||
if get_serving().tokenizer_worker_num == 1:
|
||||
if get_serving().enable_http2:
|
||||
logger.info(
|
||||
f"Starting embedded Granian HTTP/2 server on "
|
||||
f"{get_serving().host}:{get_serving().port}"
|
||||
@@ -2624,32 +2624,32 @@ def _setup_and_run_http_server(
|
||||
log_level=get_observability().log_level_http
|
||||
or get_observability().log_level,
|
||||
http2_max_concurrent_streams=(
|
||||
server_args.http2_max_concurrent_streams
|
||||
get_serving().http2_max_concurrent_streams
|
||||
),
|
||||
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_keyfile=server_args.ssl_keyfile,
|
||||
ssl_ca_certs=server_args.ssl_ca_certs,
|
||||
ssl_keyfile_password=server_args.ssl_keyfile_password,
|
||||
ssl_certfile=get_serving().ssl_certfile,
|
||||
ssl_keyfile=get_serving().ssl_keyfile,
|
||||
ssl_ca_certs=get_serving().ssl_ca_certs,
|
||||
ssl_keyfile_password=get_serving().ssl_keyfile_password,
|
||||
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.
|
||||
config = uvicorn.Config(
|
||||
app,
|
||||
host=get_serving().host,
|
||||
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
|
||||
or get_observability().log_level,
|
||||
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
|
||||
loop="uvloop",
|
||||
ssl_keyfile=server_args.ssl_keyfile,
|
||||
ssl_certfile=server_args.ssl_certfile,
|
||||
ssl_ca_certs=server_args.ssl_ca_certs,
|
||||
ssl_keyfile_password=server_args.ssl_keyfile_password,
|
||||
ssl_keyfile=get_serving().ssl_keyfile,
|
||||
ssl_certfile=get_serving().ssl_certfile,
|
||||
ssl_ca_certs=get_serving().ssl_ca_certs,
|
||||
ssl_keyfile_password=get_serving().ssl_keyfile_password,
|
||||
)
|
||||
config.load() # Creates the SSLContext
|
||||
|
||||
@@ -2660,9 +2660,9 @@ def _setup_and_run_http_server(
|
||||
async def _run_with_ssl_refresh():
|
||||
refresher = SSLCertRefresher(
|
||||
config.ssl,
|
||||
server_args.ssl_keyfile,
|
||||
server_args.ssl_certfile,
|
||||
server_args.ssl_ca_certs,
|
||||
get_serving().ssl_keyfile,
|
||||
get_serving().ssl_certfile,
|
||||
get_serving().ssl_ca_certs,
|
||||
)
|
||||
logger.info("SSL certificate auto-refresh enabled.")
|
||||
try:
|
||||
@@ -2679,15 +2679,15 @@ def _setup_and_run_http_server(
|
||||
app,
|
||||
host=get_serving().host,
|
||||
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
|
||||
or get_observability().log_level,
|
||||
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
|
||||
loop="uvloop",
|
||||
ssl_keyfile=server_args.ssl_keyfile,
|
||||
ssl_certfile=server_args.ssl_certfile,
|
||||
ssl_ca_certs=server_args.ssl_ca_certs,
|
||||
ssl_keyfile_password=server_args.ssl_keyfile_password,
|
||||
ssl_keyfile=get_serving().ssl_keyfile,
|
||||
ssl_certfile=get_serving().ssl_certfile,
|
||||
ssl_ca_certs=get_serving().ssl_ca_certs,
|
||||
ssl_keyfile_password=get_serving().ssl_keyfile_password,
|
||||
)
|
||||
else:
|
||||
# Multiple tokenizer and http processes
|
||||
@@ -2699,14 +2699,14 @@ def _setup_and_run_http_server(
|
||||
"propagate": False,
|
||||
}
|
||||
|
||||
if server_args.enable_ssl_refresh:
|
||||
if get_serving().enable_ssl_refresh:
|
||||
logger.warning(
|
||||
"--enable-ssl-refresh is not supported with multiple "
|
||||
"tokenizer workers (--tokenizer-worker-num > 1). "
|
||||
"SSL refresh will be disabled."
|
||||
)
|
||||
|
||||
if server_args.enable_http2:
|
||||
if get_serving().enable_http2:
|
||||
logger.info(
|
||||
f"Starting embedded Granian HTTP/2 server on "
|
||||
f"{get_serving().host}:{get_serving().port}"
|
||||
@@ -2717,36 +2717,36 @@ def _setup_and_run_http_server(
|
||||
log_level=get_observability().log_level_http
|
||||
or get_observability().log_level,
|
||||
http2_max_concurrent_streams=(
|
||||
server_args.http2_max_concurrent_streams
|
||||
get_serving().http2_max_concurrent_streams
|
||||
),
|
||||
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,
|
||||
ssl_certfile=server_args.ssl_certfile,
|
||||
ssl_keyfile=server_args.ssl_keyfile,
|
||||
ssl_ca_certs=server_args.ssl_ca_certs,
|
||||
ssl_keyfile_password=server_args.ssl_keyfile_password,
|
||||
tokenizer_worker_num=get_serving().tokenizer_worker_num,
|
||||
ssl_certfile=get_serving().ssl_certfile,
|
||||
ssl_keyfile=get_serving().ssl_keyfile,
|
||||
ssl_ca_certs=get_serving().ssl_ca_certs,
|
||||
ssl_keyfile_password=get_serving().ssl_keyfile_password,
|
||||
)
|
||||
else:
|
||||
uvicorn.run(
|
||||
"sglang.srt.entrypoints.http_server:app",
|
||||
host=get_serving().host,
|
||||
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
|
||||
or get_observability().log_level,
|
||||
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
|
||||
timeout_worker_healthcheck=envs.SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT.get(),
|
||||
loop="uvloop",
|
||||
workers=server_args.tokenizer_worker_num,
|
||||
ssl_keyfile=server_args.ssl_keyfile,
|
||||
ssl_certfile=server_args.ssl_certfile,
|
||||
ssl_ca_certs=server_args.ssl_ca_certs,
|
||||
ssl_keyfile_password=server_args.ssl_keyfile_password,
|
||||
workers=get_serving().tokenizer_worker_num,
|
||||
ssl_keyfile=get_serving().ssl_keyfile,
|
||||
ssl_certfile=get_serving().ssl_certfile,
|
||||
ssl_ca_certs=get_serving().ssl_ca_certs,
|
||||
ssl_keyfile_password=get_serving().ssl_keyfile_password,
|
||||
)
|
||||
finally:
|
||||
if server_args.tokenizer_worker_num > 1:
|
||||
if get_serving().tokenizer_worker_num > 1:
|
||||
if multi_tokenizer_args_shm is not None:
|
||||
multi_tokenizer_args_shm.unlink()
|
||||
if _global_state is not None:
|
||||
|
||||
@@ -5,6 +5,7 @@ from typing import List, Optional, Tuple
|
||||
import requests
|
||||
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.entrypoints.EngineBase import EngineBase
|
||||
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()
|
||||
|
||||
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:
|
||||
while time.perf_counter() - start_time < timeout:
|
||||
try:
|
||||
headers = {
|
||||
"Content-Type": "application/json; charset=utf-8",
|
||||
"Authorization": f"Bearer {server_args.api_key}",
|
||||
"Authorization": f"Bearer {cfg.api_key}",
|
||||
}
|
||||
response = session.get(
|
||||
f"{base_url}/health_generate", headers=headers, verify=ssl_verify
|
||||
@@ -60,9 +64,11 @@ class HttpServerEngineAdapter(EngineBase):
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.server_args = ServerArgs(**kwargs)
|
||||
print(
|
||||
f"Launch HttpServerEngineAdapter at: {self.server_args.host}:{self.server_args.port}"
|
||||
)
|
||||
# This process launches the server as a child and never publishes, so
|
||||
# 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)
|
||||
|
||||
def _make_request(self, endpoint: str, payload: Optional[dict] = None):
|
||||
@@ -97,7 +103,7 @@ class HttpServerEngineAdapter(EngineBase):
|
||||
{
|
||||
"serialized_named_tensors": [
|
||||
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,
|
||||
"flush_cache": flush_cache,
|
||||
|
||||
@@ -17,6 +17,7 @@ from sglang.srt.entrypoints.openai.transcription_adapters.base import (
|
||||
TranscriptionAdapter,
|
||||
)
|
||||
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.utils import random_uuid
|
||||
|
||||
@@ -85,7 +86,7 @@ async def handle_realtime_transcription(
|
||||
websocket,
|
||||
"too_many_sessions",
|
||||
f"Maximum concurrent sessions reached "
|
||||
f"({server_args.asr_max_concurrent_sessions}).",
|
||||
f"({get_serving().asr_max_concurrent_sessions}).",
|
||||
error_type="rate_limit_exceeded",
|
||||
)
|
||||
return
|
||||
|
||||
@@ -72,6 +72,7 @@ from sglang.srt.entrypoints.openai.transcription_adapters.base import (
|
||||
TranscriptionAdapter,
|
||||
)
|
||||
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.utils import random_uuid
|
||||
|
||||
@@ -186,7 +187,7 @@ class RealtimeConnection:
|
||||
|
||||
self.model_sample_rate = adapter.model_sample_rate
|
||||
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()
|
||||
|
||||
|
||||
@@ -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.managers.io_struct import EmbeddingReqInput, GenerateReqInput
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -29,11 +30,9 @@ class OpenAIServingBase(ABC):
|
||||
def __init__(self, tokenizer_manager: TokenizerManager):
|
||||
self.tokenizer_manager = tokenizer_manager
|
||||
self.allowed_custom_labels = (
|
||||
set(
|
||||
self.tokenizer_manager.server_args.tokenizer_metrics_allowed_custom_labels
|
||||
)
|
||||
set(get_observability().tokenizer_metrics_allowed_custom_labels)
|
||||
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
|
||||
)
|
||||
|
||||
@@ -230,14 +229,12 @@ class OpenAIServingBase(ABC):
|
||||
def extract_custom_labels(self, raw_request):
|
||||
if (
|
||||
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
|
||||
|
||||
custom_labels = None
|
||||
header = (
|
||||
self.tokenizer_manager.server_args.tokenizer_metrics_custom_labels_header
|
||||
)
|
||||
header = get_observability().tokenizer_metrics_custom_labels_header
|
||||
try:
|
||||
raw_labels = (
|
||||
orjson.loads(raw_request.headers.get(header))
|
||||
|
||||
@@ -10,6 +10,8 @@ from enum import Enum
|
||||
from http import HTTPStatus
|
||||
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):
|
||||
"""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.reasoning_parser = self.tokenizer_manager.config_value("reasoning_parser")
|
||||
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
|
||||
if self.reasoning_parser:
|
||||
@@ -317,7 +319,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
self._dsv4_reasoning_effort_profile = (
|
||||
chat_encoding.resolve_dsv4_reasoning_effort_profile(
|
||||
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(
|
||||
chat_encoding.DSV4_REASONING_EFFORT_PROFILE_OVERRIDE
|
||||
),
|
||||
@@ -727,7 +729,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
def _continuous_usage_cached_details(
|
||||
self, content: Dict[str, Any]
|
||||
) -> Optional[PromptTokensDetails]:
|
||||
if not self.tokenizer_manager.server_args.enable_cache_report:
|
||||
if not get_serving().enable_cache_report:
|
||||
return None
|
||||
return UsageProcessor._details_if_cached(
|
||||
content["meta_info"].get("cached_tokens", 0)
|
||||
@@ -819,7 +821,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""Generate SSE chunks for streaming content."""
|
||||
offset = stream_offsets.get(index, 0)
|
||||
if self.tokenizer_manager.server_args.incremental_streaming_output:
|
||||
if get_serving().incremental_streaming_output:
|
||||
delta = content["text"]
|
||||
else:
|
||||
delta = content["text"][offset:]
|
||||
@@ -1007,12 +1009,12 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
)
|
||||
|
||||
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 (
|
||||
max_output_tokens
|
||||
and 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 (
|
||||
f"max_completion_tokens is too large: {max_output_tokens}."
|
||||
f"This model supports at most {server_context_length} completion tokens."
|
||||
@@ -1714,7 +1716,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
try:
|
||||
include_usage, continuous_usage_stats = should_include_usage(
|
||||
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)
|
||||
@@ -1985,7 +1987,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
completion_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
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,
|
||||
audio_tokens=total_audio_tokens,
|
||||
video_tokens=total_video_tokens,
|
||||
@@ -2209,7 +2211,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
usage = UsageProcessor.calculate_response_usage(
|
||||
ret,
|
||||
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,
|
||||
audio_tokens=audio_tokens,
|
||||
video_tokens=video_tokens,
|
||||
@@ -2432,7 +2434,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
"""Process logprobs for streaming response"""
|
||||
output_token_logprobs = content["meta_info"]["output_token_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[
|
||||
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 (
|
||||
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.utils import convert_json_schema_to_str
|
||||
|
||||
@@ -246,7 +247,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
try:
|
||||
include_usage, continuous_usage_stats = should_include_usage(
|
||||
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(
|
||||
@@ -308,7 +309,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
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[
|
||||
n_prev_token:total_output_logprobs
|
||||
]
|
||||
@@ -327,7 +328,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
chunk_prompt_token_ids = None
|
||||
if request.return_token_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)
|
||||
chunk_token_ids = output_ids[n_prev_token_id:]
|
||||
n_prev_token_ids[index] = len(output_ids)
|
||||
@@ -337,7 +338,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
chunk_prompt_token_ids = content.get("prompt_token_ids")
|
||||
|
||||
# Generate delta
|
||||
if self.tokenizer_manager.server_args.incremental_streaming_output:
|
||||
if get_serving().incremental_streaming_output:
|
||||
delta = text
|
||||
else:
|
||||
delta = text[offset:]
|
||||
@@ -475,7 +476,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
completion_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
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(
|
||||
id=content["meta_info"]["id"],
|
||||
@@ -620,7 +621,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
choices.append(choice_data)
|
||||
|
||||
# Calculate usage
|
||||
cache_report = self.tokenizer_manager.server_args.enable_cache_report
|
||||
cache_report = get_serving().enable_cache_report
|
||||
usage = UsageProcessor.calculate_response_usage(
|
||||
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.managers.io_struct import GenerateReqInput
|
||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
||||
from sglang.srt.runtime_context import get_serving
|
||||
from sglang.srt.sampling.sampling_params import (
|
||||
set_request_reasoning_end_token_ids,
|
||||
)
|
||||
@@ -2031,7 +2032,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
finish_reason: Optional[dict[str, Any]] = None
|
||||
flushed = False
|
||||
stream_offset = 0
|
||||
incremental = self.tokenizer_manager.server_args.incremental_streaming_output
|
||||
incremental = get_serving().incremental_streaming_output
|
||||
|
||||
def _open_reasoning_item() -> str:
|
||||
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.managers.io_struct import GenerateReqInput
|
||||
from sglang.srt.runtime_context import get_serving
|
||||
|
||||
if TYPE_CHECKING:
|
||||
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
|
||||
# event loop on first acquire (uvicorn's loop in normal serving).
|
||||
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:
|
||||
@@ -498,11 +499,7 @@ class OpenAIServingTranscription(OpenAIServingBase):
|
||||
# the cumulative text. Always reconstruct cumulative text locally
|
||||
# so the rest of the loop (prefix parse + visible-buffer slice)
|
||||
# works uniformly under either mode.
|
||||
incremental = getattr(
|
||||
self.tokenizer_manager.server_args,
|
||||
"incremental_streaming_output",
|
||||
False,
|
||||
)
|
||||
incremental = get_serving().incremental_streaming_output
|
||||
cumulative_text = ""
|
||||
|
||||
try:
|
||||
@@ -613,11 +610,7 @@ class OpenAIServingTranscription(OpenAIServingBase):
|
||||
model = request.model
|
||||
fused_mode = getattr(request, "_fused_autodetect", False)
|
||||
ts_variant = getattr(request, "_fused_ts_variant", False)
|
||||
incremental = getattr(
|
||||
self.tokenizer_manager.server_args,
|
||||
"incremental_streaming_output",
|
||||
False,
|
||||
)
|
||||
incremental = get_serving().incremental_streaming_output
|
||||
|
||||
def _frame(delta: Optional[str], finish_reason: Optional[str] = None) -> str:
|
||||
chunk = TranscriptionStreamResponse(
|
||||
|
||||
@@ -8,6 +8,7 @@ from sglang.kernels.ops.kv_canary.consts import (
|
||||
RealKvHashMode,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_observability
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -61,18 +62,18 @@ class CanaryConfig:
|
||||
|
||||
@classmethod
|
||||
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"):
|
||||
raise ValueError(
|
||||
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(
|
||||
mode=CanaryMode(mode_raw),
|
||||
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],
|
||||
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(),
|
||||
|
||||
@@ -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.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
@@ -474,7 +475,7 @@ class DeepseekV4HipRadixBackend(
|
||||
model_runner.model_config.hf_text_config, "index_topk", C4_TOPK
|
||||
)
|
||||
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
|
||||
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
||||
|
||||
@@ -314,7 +314,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
model_runner
|
||||
)
|
||||
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
|
||||
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.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_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
|
||||
# linear_attn_backends stamp first, and that guard test constructs
|
||||
# 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, (
|
||||
f"{self.conv_states_shape[-1]=} should be less than {self.mamba_chunk_size}"
|
||||
)
|
||||
assert (
|
||||
model_runner.server_args.mamba_track_interval >= self.mamba_chunk_size
|
||||
), (
|
||||
f"mamba_track_interval ({model_runner.server_args.mamba_track_interval}) must be >= mamba_chunk_size ({self.mamba_chunk_size})"
|
||||
assert get_exec().mamba.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})"
|
||||
)
|
||||
|
||||
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.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_platform,
|
||||
get_schedule,
|
||||
@@ -189,7 +190,7 @@ class MiniCPMSparseBackend(AttentionBackend):
|
||||
model_runner.token_to_kv_pool_allocator,
|
||||
kernel_size=self.kernel_size,
|
||||
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_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)
|
||||
# Work-Centric (Lean) Attention activation. None => auto-gate from host-side
|
||||
# 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_capture_policy = lean_capture_policy
|
||||
self.extend_attention_fwd = torch.compiler.disable(extend_attention_fwd)
|
||||
@@ -333,9 +333,7 @@ class TritonAttnBackend(AttentionBackend):
|
||||
)
|
||||
self.static_kv_splits = False
|
||||
else:
|
||||
self.split_tile_size = (
|
||||
model_runner.server_args.triton_attention_split_tile_size
|
||||
)
|
||||
self.split_tile_size = get_exec().kernel.triton_attention_split_tile_size
|
||||
|
||||
if self.split_tile_size is not None:
|
||||
self.max_kv_splits = (
|
||||
|
||||
@@ -40,7 +40,11 @@ from typing import Optional
|
||||
import torch
|
||||
|
||||
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
|
||||
|
||||
# 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"})
|
||||
|
||||
|
||||
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
|
||||
setup, alongside ``initialize_dp_attention``."""
|
||||
architectures = model_config.hf_config.architectures
|
||||
get_flags().sp.enabled = bool(
|
||||
server_args.enable_layernorm_sp
|
||||
get_parallel().enable_layernorm_sp
|
||||
and architectures
|
||||
and architectures[0] in SP_SUPPORTED_ARCHITECTURES
|
||||
)
|
||||
|
||||
@@ -33,7 +33,7 @@ _EXPERT_WEIGHT_NAMES = (
|
||||
|
||||
class DwdpManager:
|
||||
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.device_id = torch.cuda.current_device()
|
||||
self.layout: Optional[DwdpExpertLayout] = None
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Optional
|
||||
|
||||
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 (
|
||||
CommunicateWithAllReduceAndLayerNormFn,
|
||||
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.moe import get_moe_a2a_backend
|
||||
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__)
|
||||
|
||||
@@ -33,9 +37,8 @@ def is_supported_forward_mode(forward_mode: ForwardMode) -> bool:
|
||||
|
||||
def resolve_max_m(model_runner) -> int:
|
||||
"""Use framework token bounds as the workspace-capacity source of truth."""
|
||||
server_args = resolving_view(model_runner.server_args)
|
||||
decode_config = server_args.cuda_graph_config.decode
|
||||
prefill_config = server_args.cuda_graph_config.prefill
|
||||
decode_config = get_exec().graph.cuda_graph_config.decode
|
||||
prefill_config = get_exec().graph.cuda_graph_config.prefill
|
||||
candidates = [
|
||||
cutedsl_moe_max_num_tokens(model_runner.server_args),
|
||||
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)
|
||||
if service is None:
|
||||
return
|
||||
if model_runner.server_args.enable_pdmux:
|
||||
if get_disagg().enable_pdmux:
|
||||
raise RuntimeError(
|
||||
"FlashInfer MNNVL CuTe DSL fusion does not support concurrent PDMux "
|
||||
"streams sharing one mutable workspace"
|
||||
|
||||
@@ -56,7 +56,6 @@ if TYPE_CHECKING:
|
||||
DispatchOutput,
|
||||
StandardDispatchOutput,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
from sglang.srt.hardware_backend.npu.quantization.moe_methods import (
|
||||
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()
|
||||
|
||||
|
||||
def initialize_bf16_gemm_config(server_args: ServerArgs) -> None:
|
||||
def initialize_bf16_gemm_config() -> None:
|
||||
global _BF16_GEMM_BACKEND
|
||||
global _cutedsl_bf16_gemm, _use_cutedsl_bf16_gemm
|
||||
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 _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:
|
||||
backend_str = (
|
||||
"torch"
|
||||
|
||||
@@ -17,6 +17,7 @@ from sglang.srt.lora.utils import (
|
||||
merge_and_chunk_segments,
|
||||
)
|
||||
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
|
||||
|
||||
MIN_CHUNK_SIZE = 16
|
||||
@@ -42,7 +43,7 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend):
|
||||
server_args: ServerArgs,
|
||||
):
|
||||
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(
|
||||
self,
|
||||
|
||||
@@ -19,8 +19,8 @@ import multiprocessing as mp
|
||||
import signal
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from enum import Enum, auto
|
||||
from typing import Callable, List, Optional
|
||||
|
||||
import psutil
|
||||
import setproctitle
|
||||
@@ -53,7 +53,9 @@ from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_observability,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
publish,
|
||||
)
|
||||
from sglang.srt.server_args import (
|
||||
@@ -154,7 +156,7 @@ class DataParallelController:
|
||||
|
||||
# Init inter-process communication
|
||||
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.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.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, (
|
||||
f"--max-ep-size ({self.max_dp_size}) must be >= "
|
||||
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
|
||||
)
|
||||
|
||||
@@ -196,9 +198,9 @@ class DataParallelController:
|
||||
|
||||
# Launch data parallel workers
|
||||
self.scheduler_procs = []
|
||||
self.workers: List[Optional[zmq.Socket]] = [None] * self.max_dp_size
|
||||
self.status: List[bool] = list(self.dp_active)
|
||||
self._active_workers: List[int] = list(range(self.launch_dp_size))
|
||||
self.workers: list[zmq.Socket | None] = [None] * self.max_dp_size
|
||||
self.status: list[bool] = list(self.dp_active)
|
||||
self._active_workers: list[int] = list(range(self.launch_dp_size))
|
||||
self._active_count_cache: int = self.launch_dp_size
|
||||
|
||||
if get_parallel().enable_dp_attention:
|
||||
@@ -209,7 +211,7 @@ class DataParallelController:
|
||||
# Otherwise fall back to the original behaviour: send to only the
|
||||
# first leader, which then broadcasts over the full tp_group.
|
||||
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:
|
||||
self.launch_dp_schedulers(server_args, port_args)
|
||||
self.control_message_step = 1
|
||||
@@ -223,7 +225,7 @@ class DataParallelController:
|
||||
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")
|
||||
|
||||
def send_to_all_workers(self, obj):
|
||||
@@ -392,10 +394,12 @@ class DataParallelController:
|
||||
)
|
||||
threads.append(thread)
|
||||
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.context,
|
||||
zmq.PUSH,
|
||||
@@ -430,8 +434,8 @@ class DataParallelController:
|
||||
time.sleep(30 * 24 * 3600)
|
||||
|
||||
def _broadcast_worker_ports(
|
||||
self, server_args: ServerArgs, worker_ports: Optional[List[int]] = None
|
||||
) -> List[int]:
|
||||
self, server_args: ServerArgs, worker_ports: list[int] | None = None
|
||||
) -> list[int]:
|
||||
"""Broadcast worker ports from node 0 to all other nodes.
|
||||
|
||||
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).
|
||||
"""
|
||||
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(
|
||||
server_args.host or "127.0.0.1",
|
||||
server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
||||
get_serving().host or "127.0.0.1",
|
||||
get_serving().port + DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
||||
)
|
||||
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)
|
||||
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
|
||||
return self._broadcast_ports_as_server(
|
||||
endpoint, server_args.nnodes - 1, worker_ports
|
||||
endpoint, get_parallel().nnodes - 1, worker_ports
|
||||
)
|
||||
else:
|
||||
# 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(
|
||||
self, endpoint: str, expected_clients: int, worker_ports: List[int]
|
||||
) -> List[int]:
|
||||
self, endpoint: str, expected_clients: int, worker_ports: list[int]
|
||||
) -> list[int]:
|
||||
"""Broadcast worker ports to all client nodes."""
|
||||
logger.debug(f"Broadcasting worker ports to {expected_clients} client nodes")
|
||||
logger.debug(f"Worker ports: {worker_ports}")
|
||||
@@ -500,7 +504,7 @@ class DataParallelController:
|
||||
daemon=True,
|
||||
).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
|
||||
late-arriving elastic joiners. Publishes port numbers only; the primary
|
||||
keeps ownership of every socket."""
|
||||
@@ -518,9 +522,9 @@ class DataParallelController:
|
||||
sock_send(rep_socket, wrap_as_pickle(worker_ports))
|
||||
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."""
|
||||
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.setsockopt(zmq.RCVTIMEO, 600 * 1000) # 10 minute timeout
|
||||
@@ -543,37 +547,37 @@ class DataParallelController:
|
||||
req_socket.close()
|
||||
|
||||
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:
|
||||
return server_args.ep_join_rank_offset
|
||||
return get_parallel().ep_join_rank_offset
|
||||
|
||||
def launch_dp_attention_schedulers(
|
||||
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"
|
||||
else:
|
||||
bind_host = NetworkAddress.parse(server_args.dist_init_addr).host
|
||||
bind_host = NetworkAddress.parse(get_parallel().dist_init_addr).host
|
||||
|
||||
worker_ports = []
|
||||
if server_args.is_ep_scale_joiner:
|
||||
# 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.host, primary.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA
|
||||
).to_tcp()
|
||||
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)
|
||||
local_tp_span = self._joiner_local_tp_span(server_args)
|
||||
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.
|
||||
bind_count = (
|
||||
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
|
||||
)
|
||||
for slot in range(bind_count):
|
||||
@@ -601,35 +605,35 @@ class DataParallelController:
|
||||
server_args: ServerArgs,
|
||||
port_args: PortArgs,
|
||||
base_gpu_id: int,
|
||||
dp_rank: Optional[int],
|
||||
worker_ports: Optional[List[int]] = None,
|
||||
dp_rank: int | None,
|
||||
worker_ports: list[int] | None = None,
|
||||
):
|
||||
if not get_parallel().enable_dp_attention:
|
||||
logger.info(f"Launch DP{dp_rank} starting at GPU #{base_gpu_id}.")
|
||||
|
||||
memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
enable=server_args.enable_memory_saver
|
||||
enable=get_exec().features.enable_memory_saver
|
||||
)
|
||||
|
||||
scheduler_pipe_readers = []
|
||||
|
||||
pp_size_per_node = max(get_parallel().pp_size // server_args.nnodes, 1)
|
||||
nnodes_per_pp_rank = max(server_args.nnodes // get_parallel().pp_size, 1)
|
||||
pp_size_per_node = max(get_parallel().pp_size // get_parallel().nnodes, 1)
|
||||
nnodes_per_pp_rank = max(get_parallel().nnodes // get_parallel().pp_size, 1)
|
||||
pp_rank_range = range(
|
||||
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank),
|
||||
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank + 1),
|
||||
pp_size_per_node * (get_parallel().node_rank // nnodes_per_pp_rank),
|
||||
pp_size_per_node * (get_parallel().node_rank // nnodes_per_pp_rank + 1),
|
||||
)
|
||||
|
||||
nnodes_per_tp_group = nnodes_per_pp_rank
|
||||
tp_size_per_node = server_args.tp_size // nnodes_per_tp_group
|
||||
tp_size_per_node = get_parallel().tp_size // nnodes_per_tp_group
|
||||
if server_args.is_ep_scale_joiner:
|
||||
# Scale joiners enumerate their full local TP span.
|
||||
tp_rank_range = range(server_args.tp_size)
|
||||
tp_size_per_node = server_args.tp_size
|
||||
tp_rank_range = range(get_parallel().tp_size)
|
||||
tp_size_per_node = get_parallel().tp_size
|
||||
else:
|
||||
tp_rank_range = range(
|
||||
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group),
|
||||
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1),
|
||||
tp_size_per_node * (get_parallel().node_rank % nnodes_per_tp_group),
|
||||
tp_size_per_node * (get_parallel().node_rank % nnodes_per_tp_group + 1),
|
||||
)
|
||||
|
||||
attn_cp_rank = 0
|
||||
@@ -643,7 +647,7 @@ class DataParallelController:
|
||||
_, _, dp_rank, _ = compute_dp_attention_world_info(
|
||||
get_parallel().enable_dp_attention,
|
||||
tp_rank,
|
||||
server_args.tp_size,
|
||||
get_parallel().tp_size,
|
||||
get_parallel().dp_size,
|
||||
get_parallel().attn_cp_size,
|
||||
)
|
||||
@@ -653,7 +657,9 @@ class DataParallelController:
|
||||
)
|
||||
if server_args.is_ep_scale_joiner:
|
||||
# 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
|
||||
rank_port_args.tokenizer_ipc_name = NetworkAddress(
|
||||
primary_addr.host, primary_port_base
|
||||
@@ -668,10 +674,10 @@ class DataParallelController:
|
||||
|
||||
reader, writer = mp.Pipe(duplex=False)
|
||||
gpu_id = (
|
||||
server_args.base_gpu_id
|
||||
get_device().base_gpu_id
|
||||
+ base_gpu_id
|
||||
+ ((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 = (
|
||||
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)
|
||||
# - MoE: Global(TP) -> MOE_DP -> EP -> MOE_TP (innermost)
|
||||
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
|
||||
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 = (
|
||||
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().ep_size
|
||||
)
|
||||
)
|
||||
|
||||
# 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_moe_ep_rank = moe_ep_rank + offset
|
||||
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.
|
||||
publish(server_args, role="dp_controller")
|
||||
configure_logger(server_args)
|
||||
if server_args.enable_trace:
|
||||
if get_observability().enable_trace:
|
||||
process_tracing_init(
|
||||
server_args.otlp_traces_endpoint,
|
||||
get_observability().otlp_traces_endpoint,
|
||||
"sglang",
|
||||
trace_modules=server_args.trace_modules,
|
||||
trace_modules=get_observability().trace_modules,
|
||||
)
|
||||
thread_label = "DP Controller"
|
||||
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.
|
||||
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()
|
||||
for proc in controller.scheduler_procs:
|
||||
proc.join()
|
||||
|
||||
@@ -127,22 +127,22 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
||||
# In multi-tokenizer mode, results are pushed back to each TokenizerWorker
|
||||
# directly via SocketMapping inside multi_http_worker_event_loop, so the
|
||||
# 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(
|
||||
context, zmq.PUSH, port_args.tokenizer_ipc_name, False
|
||||
)
|
||||
|
||||
def init_tokenizer(self, server_args: ServerArgs):
|
||||
if server_args.skip_tokenizer_init:
|
||||
if get_serving().skip_tokenizer_init:
|
||||
self.tokenizer = None
|
||||
self.vocab_size = None
|
||||
else:
|
||||
self.tokenizer = get_tokenizer(
|
||||
get_serving().tokenizer_path,
|
||||
tokenizer_mode=server_args.tokenizer_mode,
|
||||
tokenizer_mode=get_serving().tokenizer_mode,
|
||||
trust_remote_code=get_model().trust_remote_code,
|
||||
revision=server_args.revision,
|
||||
tokenizer_backend=server_args.tokenizer_backend,
|
||||
revision=get_model().revision,
|
||||
tokenizer_backend=get_serving().tokenizer_backend,
|
||||
)
|
||||
try:
|
||||
self.vocab_size = len(self.tokenizer)
|
||||
@@ -151,7 +151,9 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
||||
|
||||
def init_running_status(self, server_args: ServerArgs):
|
||||
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.soft_watchdog = Watchdog.create(
|
||||
@@ -548,7 +550,7 @@ def run_detokenizer_process(
|
||||
manager = None
|
||||
try:
|
||||
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()
|
||||
else:
|
||||
manager.multi_http_worker_event_loop()
|
||||
|
||||
@@ -6,6 +6,7 @@ import pkgutil
|
||||
|
||||
from sglang.srt.configs.model_config import ModelImpl
|
||||
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
|
||||
from sglang.srt.runtime_context import get_model
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -41,10 +42,10 @@ def import_processors(package_name: str, overwrite: bool = False):
|
||||
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
|
||||
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"
|
||||
if model_impl == "auto" and model_config is not None:
|
||||
from sglang.srt.model_loader.utils import get_resolved_model_impl
|
||||
@@ -79,7 +80,7 @@ def get_mm_processor(
|
||||
model_config=None,
|
||||
**kwargs,
|
||||
) -> 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:
|
||||
raise ValueError(
|
||||
f"No processor registered for architecture: {hf_config.architectures}.\n"
|
||||
|
||||
@@ -992,7 +992,7 @@ class Scheduler(
|
||||
# Initialize GEMM-related configuration for FP8 and FP4 backends.
|
||||
initialize_fp8_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
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
|
||||
@@ -248,10 +248,10 @@ class SchedulerMetricsReporter:
|
||||
self.kv_transfer_latency_ms: float = 0.0
|
||||
|
||||
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:
|
||||
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:
|
||||
self._init_estimated_perf_constants()
|
||||
self._mfu_log_flops = 0.0
|
||||
|
||||
@@ -41,6 +41,7 @@ from sglang.srt.managers.io_struct import (
|
||||
UpdateWeightsFromTensorReqInput,
|
||||
UpdateWeightsFromTensorReqOutput,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_model
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -198,7 +199,7 @@ class SchedulerWeightUpdaterManager:
|
||||
freeing them would leave the daemon and every peer pointing at released
|
||||
memory.
|
||||
"""
|
||||
mode = self.tp_worker.model_runner.server_args.weight_cache_mode
|
||||
mode = get_model().weight_cache_mode
|
||||
if mode != "off":
|
||||
raise RuntimeError(
|
||||
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 (
|
||||
get_lora,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.server_args import LoRARef
|
||||
@@ -209,9 +210,7 @@ class TokenizerControlMixin:
|
||||
iter_external_corpus_chunks,
|
||||
)
|
||||
|
||||
max_tokens = (
|
||||
self.server_args.speculative_ngram_external_corpus_max_tokens
|
||||
)
|
||||
max_tokens = get_spec().speculative_ngram_external_corpus_max_tokens
|
||||
obj.token_chunks = list(
|
||||
iter_external_corpus_chunks(
|
||||
obj.file_path, self.tokenizer, max_tokens
|
||||
@@ -222,9 +221,7 @@ class TokenizerControlMixin:
|
||||
SEPARATOR_TOKEN,
|
||||
)
|
||||
|
||||
max_tokens = (
|
||||
self.server_args.speculative_ngram_external_corpus_max_tokens
|
||||
)
|
||||
max_tokens = get_spec().speculative_ngram_external_corpus_max_tokens
|
||||
token_chunks = []
|
||||
total_tokens = 0
|
||||
has_prev = False
|
||||
@@ -642,10 +639,10 @@ class TokenizerControlMixin:
|
||||
await self.lora_registry.register(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 (
|
||||
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(
|
||||
exclude_pinned=True
|
||||
@@ -659,7 +656,7 @@ class TokenizerControlMixin:
|
||||
logger.info(
|
||||
f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
|
||||
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(
|
||||
@@ -718,10 +715,10 @@ class TokenizerControlMixin:
|
||||
if result.success:
|
||||
await self.lora_registry.register(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 (
|
||||
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(
|
||||
exclude_pinned=True
|
||||
@@ -735,7 +732,7 @@ class TokenizerControlMixin:
|
||||
logger.info(
|
||||
f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
|
||||
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(
|
||||
@@ -905,7 +902,7 @@ class TokenizerControlMixin:
|
||||
):
|
||||
self.auto_create_handle_loop()
|
||||
if obj.streaming:
|
||||
if not self.server_args.enable_streaming_session:
|
||||
if not get_serving().enable_streaming_session:
|
||||
raise ValueError(
|
||||
"Streaming sessions are disabled. "
|
||||
"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.managers.embed_types import PositionalEmbeds
|
||||
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -509,7 +510,7 @@ class TokenizerManagerScoreMixin:
|
||||
)
|
||||
|
||||
# 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
|
||||
text_prompts = None
|
||||
|
||||
@@ -108,7 +108,7 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache:
|
||||
logger.info("Using experimental C++ radix tree implementation.")
|
||||
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)
|
||||
|
||||
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(
|
||||
cache.cache_controller.layer_done_counter
|
||||
)
|
||||
elif server_args.enable_unified_cache_external_linker:
|
||||
backend = server_args.unified_cache_external_linker_backend
|
||||
elif get_memory().enable_unified_cache_external_linker:
|
||||
backend = get_memory().unified_cache_external_linker_backend
|
||||
if backend == "mooncake":
|
||||
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_direct_linker import (
|
||||
MooncakeDirectLinker,
|
||||
|
||||
@@ -20,7 +20,11 @@ from sglang.srt.mem_cache.hybrid_cache.linker_pool_assembler import (
|
||||
resolve_hybrid_device_pool_group,
|
||||
)
|
||||
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
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -102,7 +106,7 @@ class MooncakeDirectLinker(UnifiedCacheLinker):
|
||||
self.num_layers = self.pool_group.num_layers
|
||||
|
||||
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
|
||||
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
||||
tp_rank = torch.distributed.get_rank(group=tp_group)
|
||||
|
||||
@@ -596,19 +596,15 @@ class CPUGraphRunner:
|
||||
self.graphs_cross = {}
|
||||
self.output_buffers = {}
|
||||
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.require_gathered_buffer = require_gathered_buffer()
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||
self.enable_two_batch_overlap = (
|
||||
model_runner.server_args.enable_two_batch_overlap
|
||||
)
|
||||
self.enable_two_batch_overlap = get_exec().overlap.enable_two_batch_overlap
|
||||
self.speculative_algorithm = get_spec().speculative_algorithm
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
)
|
||||
self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
|
||||
self.tp_size = get_parallel().tp_size
|
||||
self.dp_size = get_parallel().dp_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 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__)
|
||||
|
||||
@@ -106,9 +109,9 @@ def reuse_hccl_comm():
|
||||
create_group(group_name, group().ranks, group_options)
|
||||
|
||||
|
||||
def init_ms_distributed(world_size, rank, local_rank, server_args, port):
|
||||
if server_args.dist_init_addr:
|
||||
dist_init_method = f"tcp://{server_args.dist_init_addr}"
|
||||
def init_ms_distributed(world_size, rank, local_rank, port):
|
||||
if get_parallel().dist_init_addr:
|
||||
dist_init_method = f"tcp://{get_parallel().dist_init_addr}"
|
||||
else:
|
||||
dist_init_method = f"tcp://{get_serving().host}:{port}"
|
||||
set_ms_parallel_env(rank, local_rank, world_size, dist_init_method)
|
||||
|
||||
@@ -316,7 +316,7 @@ class ModelRunner:
|
||||
|
||||
def supports_sampling_observer(self) -> bool:
|
||||
"""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__(
|
||||
self,
|
||||
@@ -644,7 +644,6 @@ class ModelRunner:
|
||||
world_size=self.ps.tp_size * self.ps.pp_size,
|
||||
rank=self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank,
|
||||
local_rank=self.gpu_id,
|
||||
server_args=self.server_args,
|
||||
port=self.dist_port,
|
||||
)
|
||||
|
||||
|
||||
@@ -42,6 +42,8 @@ from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_model,
|
||||
get_observability,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_spec,
|
||||
@@ -281,10 +283,8 @@ def capture_cuda_graphs(
|
||||
# not traced into any captured graph — capture stays hook-free and hooks
|
||||
# fire only on the eager forward path (capture replay never runs Python
|
||||
# hooks anyway).
|
||||
if model_runner.server_args.forward_hooks:
|
||||
register_forward_hooks(
|
||||
model_runner.model, model_runner.server_args.forward_hooks
|
||||
)
|
||||
if get_observability().forward_hooks:
|
||||
register_forward_hooks(model_runner.model, get_observability().forward_hooks)
|
||||
|
||||
prealloc_symmetric_memory_pool(
|
||||
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:
|
||||
# TODO: Currently, cuda graph only captures decode steps, which only exists for generation models
|
||||
return no_capture
|
||||
if model_runner.server_args.model_impl.lower() == ModelImpl.MINDSPORE:
|
||||
if get_model().model_impl.lower() == ModelImpl.MINDSPORE:
|
||||
return no_capture
|
||||
if model_runner.device != "cpu" and check_cuda_graph_backend(
|
||||
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.weight_utils import default_weight_loader
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_model
|
||||
from sglang.srt.utils import (
|
||||
MultiprocessingSerializer,
|
||||
dynamic_import,
|
||||
@@ -128,7 +129,7 @@ class WeightUpdater:
|
||||
param.data is the daemon's master copy shared with every co-attached
|
||||
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":
|
||||
raise RuntimeError(
|
||||
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.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_memory,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
@@ -887,7 +888,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
|
||||
# keeps the FP8 estimate.
|
||||
self.indexer_bytes_per_token = get_dsv4_indexer_bytes_per_token(
|
||||
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
|
||||
# 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
|
||||
self.dp_size = get_parallel().dp_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 = (
|
||||
CaptureHiddenMode.NULL
|
||||
if model_runner.is_draft_worker
|
||||
|
||||
@@ -229,7 +229,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
|
||||
# --- core state ------------------------------------------------
|
||||
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.require_mlp_tp_gather = (
|
||||
require_mlp_tp_gather() and not self._forward_is_dp_local(model_runner)
|
||||
@@ -244,18 +244,14 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
self.require_mlp_sync = (
|
||||
get_parallel().enable_dp_attention or self.require_gathered_buffer
|
||||
)
|
||||
self.enable_two_batch_overlap = (
|
||||
model_runner.server_args.enable_two_batch_overlap
|
||||
)
|
||||
self.enable_two_batch_overlap = get_exec().overlap.enable_two_batch_overlap
|
||||
self.use_ngram_embedding = model_runner.ngram_embedding_manager.enabled
|
||||
if self.use_ngram_embedding:
|
||||
hf_config = model_runner.model_config.hf_config
|
||||
self.ngram_embedding_n = hf_config.ngram_embedding_n
|
||||
self.ngram_embedding_k = hf_config.ngram_embedding_k
|
||||
self.speculative_algorithm = get_spec().speculative_algorithm
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
)
|
||||
self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
|
||||
|
||||
# --- DSA dense-decode dual-graph -------------------------------
|
||||
# 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.
|
||||
# Capture the large shapes first so that the smaller 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:
|
||||
with (
|
||||
graph_capture(
|
||||
|
||||
@@ -1413,7 +1413,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
# Warm up + autotune kernels once before capture (run-once across the
|
||||
# decode + prefill runners; see BaseRunner.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(
|
||||
stream=get_or_create_global_graph_capture_stream()
|
||||
) as graph_capture_context:
|
||||
|
||||
@@ -8,6 +8,8 @@ from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
|
||||
try:
|
||||
from cuda.bindings import driver as cuda_drv
|
||||
from cuda.bindings import runtime as cuda_rt
|
||||
@@ -314,7 +316,7 @@ class DedupedCudaGraphMixin:
|
||||
server_args = getattr(model_runner, "server_args", None)
|
||||
return bool(
|
||||
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")
|
||||
)
|
||||
|
||||
|
||||
@@ -125,7 +125,7 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
|
||||
config = CompilationConfig(
|
||||
num_tokens,
|
||||
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():
|
||||
|
||||
@@ -43,6 +43,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
|
||||
get_mm_feature_pool_size_per_worker,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_mm,
|
||||
get_serving,
|
||||
)
|
||||
@@ -723,7 +724,7 @@ class BaseMultimodalProcessor(ABC):
|
||||
preprocessing worker there is one more competitor for that device rather
|
||||
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
|
||||
if self.disable_fast_image_processor:
|
||||
return False
|
||||
@@ -759,7 +760,7 @@ class BaseMultimodalProcessor(ABC):
|
||||
tokenizer process each carry their own ``base_gpu_id``.
|
||||
"""
|
||||
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"
|
||||
if _is_xpu:
|
||||
return "xpu"
|
||||
|
||||
@@ -20,6 +20,7 @@ from sglang.srt.multimodal.processors.base_processor import (
|
||||
BaseMultiModalProcessorOutput,
|
||||
MultimodalSpecialTokens,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_model
|
||||
from sglang.srt.utils import get_device
|
||||
from sglang.srt.utils.video_decoder import VideoDecoderWrapper
|
||||
|
||||
@@ -135,7 +136,7 @@ class InternVLProcessor(BaseMultimodalProcessor):
|
||||
).build(_image_processor)
|
||||
|
||||
self.max_context_len = (
|
||||
getattr(server_args, "context_length", None)
|
||||
get_model().context_length
|
||||
or getattr(server_args, "max_context_len", None)
|
||||
or getattr(hf_config, "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 (
|
||||
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_mm
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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
|
||||
# image copy to every rank. The scheduler only honors this marker once
|
||||
# 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:
|
||||
item.model_specific_data[DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY] = (
|
||||
True
|
||||
|
||||
@@ -46,6 +46,7 @@ from sglang.srt.multimodal.processors.base_processor import (
|
||||
BaseMultimodalProcessor,
|
||||
MultimodalSpecialTokens,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_model
|
||||
from sglang.srt.utils.common import sample_video_frames
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -143,7 +144,7 @@ class NanoNemotronVLImageProcessor(BaseMultimodalProcessor):
|
||||
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
|
||||
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 (
|
||||
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.video_decoder import VideoDecoderWrapper
|
||||
from sglang.utils import logger
|
||||
@@ -897,7 +898,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
||||
def _mark_dp_encoder_features_for_deferred_reconstruction(self, mm_items):
|
||||
if not (
|
||||
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
|
||||
in ("qwen3_vl", "qwen3_vl_moe", "qwen3_5", "qwen3_5_moe")
|
||||
):
|
||||
|
||||
@@ -83,7 +83,10 @@ class FileRequestMetricsExporter(RequestMetricsExporter):
|
||||
out_skip_names: Optional[set[str]],
|
||||
):
|
||||
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)
|
||||
|
||||
# File handler state management
|
||||
|
||||
@@ -50,6 +50,7 @@ from sglang.srt.parser.template_detection import (
|
||||
detect_reasoning_pattern,
|
||||
match_rules,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_serving
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -384,7 +385,7 @@ class TemplateManager:
|
||||
logger.info(f"Multiple HuggingFace chat templates available: {available_names}")
|
||||
|
||||
# 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:
|
||||
raise ValueError(
|
||||
f"Specified template '{preferred_name}' not found. "
|
||||
|
||||
@@ -203,9 +203,7 @@ def _create_scheduler_actor(
|
||||
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).
|
||||
"""
|
||||
attn_cp_rank, moe_dp_rank, moe_ep_rank = _compute_parallelism_ranks(
|
||||
server_args, tp_rank
|
||||
)
|
||||
attn_cp_rank, moe_dp_rank, moe_ep_rank = _compute_parallelism_ranks(tp_rank)
|
||||
|
||||
return SchedulerActor.options(
|
||||
num_cpus=0,
|
||||
|
||||
@@ -1347,7 +1347,13 @@ ROLE_NAMESPACE_SETS: dict[str, frozenset[str] | None] = {
|
||||
# The DP controller's static read set, checked against the module: the
|
||||
# elastic-EP gate, the load-balance method, the watchdog timeout, and the
|
||||
# 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 +
|
||||
# /v1/models): reads exactly {"serving"} — the per-instance managers read
|
||||
# self.server_args by design. Still declared full, because that run did not
|
||||
|
||||
@@ -167,7 +167,7 @@ class RustMmProcessor:
|
||||
|
||||
hf_config = self.model_config.hf_config
|
||||
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(
|
||||
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,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
@@ -115,14 +116,12 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
self.attn_dp_size = model_runner.ps.attn_dp_size
|
||||
self.pp_size = get_parallel().pp_size
|
||||
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_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
)
|
||||
self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
|
||||
self.speculative_num_steps = (
|
||||
get_spec().speculative_num_steps
|
||||
if speculative_num_steps is None
|
||||
@@ -196,7 +195,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
(self.max_bs, self.model_runner.model_config.vocab_size),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
if self.model_runner.server_args.speculative_use_rejection_sampling
|
||||
if get_spec().speculative_use_rejection_sampling
|
||||
else None
|
||||
)
|
||||
_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
|
||||
# the copy otherwise to keep the non-RS path free of extra work.
|
||||
if (
|
||||
self.model_runner.server_args.speculative_use_rejection_sampling
|
||||
get_spec().speculative_use_rejection_sampling
|
||||
and forward_batch.sampling_info is not None
|
||||
):
|
||||
self.temperatures[:raw_bs].copy_(
|
||||
|
||||
@@ -36,6 +36,7 @@ from sglang.srt.model_executor.runner_backend_utils import (
|
||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
@@ -114,14 +115,12 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
self.attn_dp_size = model_runner.ps.attn_dp_size
|
||||
self.pp_size = get_parallel().pp_size
|
||||
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_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
)
|
||||
self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
|
||||
self.speculative_num_steps = (
|
||||
get_spec().speculative_num_steps
|
||||
if speculative_num_steps is None
|
||||
|
||||
@@ -33,6 +33,7 @@ from sglang.srt.model_executor.runner_backend_utils import (
|
||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
@@ -92,7 +93,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
self.device = model_runner.device
|
||||
self.device_module = torch.get_device_module(self.device)
|
||||
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_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
@@ -103,9 +104,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||
self.topk = get_spec().speculative_eagle_topk
|
||||
self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
)
|
||||
self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
|
||||
|
||||
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,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
@@ -157,18 +159,16 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
self.dp_size = get_parallel().dp_size
|
||||
self.pp_size = get_parallel().pp_size
|
||||
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_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
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_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||
self.topk = get_spec().speculative_eagle_topk
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
)
|
||||
self.enable_profile_cuda_graph = get_exec().graph.enable_profile_cuda_graph
|
||||
self.attn_backend = self.eagle_worker.draft_extend_attn_backend_list[self.step]
|
||||
self.metadata_captured_in_graph = (
|
||||
self.attn_backend.draft_extend_metadata_captured_in_graph()
|
||||
|
||||
@@ -50,7 +50,11 @@ import torch.distributed as dist
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.configs.load_config import LoadConfig
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_parallel, publish
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
publish,
|
||||
)
|
||||
|
||||
from .protocol import (
|
||||
CacheConfig,
|
||||
@@ -465,7 +469,7 @@ class WeightCacheDaemon:
|
||||
|
||||
def _initialize_eplb_expert_location_metadata(self, model_config) -> None:
|
||||
"""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
|
||||
|
||||
from sglang.srt.eplb.expert_location import (
|
||||
|
||||
@@ -2143,6 +2143,20 @@ def server_args_variant(server_args, **fields):
|
||||
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):
|
||||
def __init_subclass__(cls, **kwargs):
|
||||
super().__init_subclass__(**kwargs)
|
||||
|
||||
Reference in New Issue
Block a user