config: read resolved config via namespace accessors (#33013)

This commit is contained in:
Cheng Wan
2026-07-31 15:06:59 -07:00
committed by GitHub
parent 4862edc85f
commit 55b6769b0e
187 changed files with 1110 additions and 923 deletions
@@ -48,7 +48,7 @@ from sglang.srt.managers.scheduler import run_scheduler_process
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
from sglang.srt.runtime_context import publish
from sglang.srt.runtime_context import get_exec, publish
from sglang.srt.server_args import (
DP_ATTENTION_HANDSHAKE_PORT_DELTA,
PortArgs,
@@ -232,7 +232,7 @@ class DataParallelController:
sock_send(worker, obj)
def update_active_ranks(self, ranks: ActiveRanksOutput):
if self.server_args.elastic_ep_backend is not None:
if get_exec().moe.elastic_ep_backend is not None:
if len(ranks.status) != self.max_dp_size:
logger.warning(
"[Elastic EP][DPC] active rank status len=%d != max_dp_size=%d; "
@@ -485,7 +485,7 @@ class DataParallelController:
logger.debug("Worker port broadcast completed")
return worker_ports
finally:
if self.server_args.elastic_ep_backend is None:
if get_exec().moe.elastic_ep_backend is None:
rep_socket.close()
else:
threading.Thread(
+9 -4
View File
@@ -33,7 +33,12 @@ from sglang.srt.managers.schedule_batch import (
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.multimodal.evs import EVSEmbeddingResult
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.runtime_context import (
get_disagg,
get_parallel,
get_schedule,
get_server_args,
)
from sglang.srt.utils import flatten_nested_list, is_hip, is_npu, print_warning_once
from sglang.srt.utils.stale_shm_cleanup import make_shm_name
from sglang.utils import logger
@@ -931,7 +936,7 @@ def _adjust_embedding_length(
f"tokens from multimodal embeddings."
)
if num_mm_tokens_in_input_ids < num_mm_tokens_in_embedding:
chunked_prefill_size = get_server_args().chunked_prefill_size
chunked_prefill_size = get_schedule().chunked_prefill_size
if chunked_prefill_size != -1:
logger.warning(
"You may want to avoid this issue by raising `chunked_prefill_size`, or disabling chunked prefill"
@@ -1295,7 +1300,7 @@ def general_mm_embed_routine(
# encoder/ViT execution and multimodal feature placement, while
# the language model range below excludes both.
with torch.profiler.record_function("sglang.vlm.mm_embedding"):
if server_args and server_args.enable_adaptive_dispatch_to_encoder:
if server_args and get_disagg().enable_adaptive_dispatch_to_encoder:
# Split by precomputed vs non-precomputed so get_embedding_and_mask only sees uniform batches
input_embeds, other_info = _embed_mm_inputs_with_split(
mm_inputs_list=mm_inputs_list,
@@ -1340,7 +1345,7 @@ def general_mm_embed_routine(
feature = getattr(mm_item, "feature", None)
if isinstance(feature, torch.Tensor) and feature.is_cuda:
mm_item.feature = feature.to("cpu", non_blocking=True)
if get_server_args().language_only:
if get_disagg().language_only:
precomputed_embeddings = getattr(
mm_item, "precomputed_embeddings", None
)
+6 -5
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
from sglang.srt.dllm.config import DllmConfig
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import get_exec, get_schedule, get_serving, get_spec
from sglang.srt.utils.common import (
Range,
ceil_align,
@@ -1097,7 +1098,7 @@ class Req(ReqDllmMixin):
"""Check if this request is prefill-only (no token generation needed)."""
# NOTE: when spec is enabled, prefill_only optimizations are disabled
spec_alg = get_server_args().speculative_algorithm
spec_alg = get_spec().speculative_algorithm
return self.sampling_params.max_new_tokens == 0 and spec_alg is None
@property
@@ -1118,7 +1119,7 @@ class Req(ReqDllmMixin):
def effective_kv_committed_len(self) -> int:
# Report only the prompt prefix so thinking + answer fall into the
# overallocated range and are reclaimed by release_kv_cache. #22373.
if get_server_args().strip_thinking_cache and self.reasoning_tokens > 0:
if get_serving().strip_thinking_cache and self.reasoning_tokens > 0:
return min(self.kv_committed_len, len(self.origin_input_ids))
return self.kv_committed_len
@@ -2922,7 +2923,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
)
if server_args.enable_mamba_extra_buffer():
mamba_track_interval = server_args.mamba_track_interval
mamba_track_interval = get_exec().mamba.mamba_track_interval
if len(self.reqs) == 0:
self.mamba_track_indices = torch.empty(
@@ -3168,8 +3169,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
continue
else:
pre_len = (
pre_len - server_args.chunked_prefill_size
if server_args.chunked_prefill_size > 0
pre_len - get_schedule().chunked_prefill_size
if get_schedule().chunked_prefill_size > 0
else pre_len
)
self._evict_swa(req, pre_len)
@@ -5,6 +5,7 @@ from array import array
from sglang.srt.environ import envs
from sglang.srt.managers.prefill_delayer import PrefillDelayerSinglePassExecutor
from sglang.srt.runtime_context import get_disagg
from sglang.srt.utils import get_bool_env_var
_ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG")
@@ -56,7 +57,6 @@ from sglang.srt.mem_cache.multi_ended_allocator import (
UnifiedMambaTokenToKVPoolAllocator,
)
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
from sglang.srt.runtime_context import get_server_args
from sglang.srt.server_args import ServerArgs
if TYPE_CHECKING:
@@ -195,7 +195,7 @@ class SchedulePolicy:
if (
not isinstance(policy, CacheAwarePolicy)
and self.tree_cache.supports_fast_match_prefix()
and get_server_args().disaggregation_mode != "decode"
and get_disagg().disaggregation_mode != "decode"
):
for r in waiting_queue:
match_prefix_for_req(self.tree_cache, r, include_req=True)
+84 -75
View File
@@ -27,6 +27,20 @@ from functools import partial
from http import HTTPStatus
from typing import Any, Deque, Dict, List, Optional, Tuple, Union
from sglang.srt.runtime_context import (
get_device,
get_disagg,
get_exec,
get_lora,
get_memory,
get_mm,
get_model,
get_observability,
get_schedule,
get_serving,
get_spec,
)
from sglang.srt.utils.common import suppress_noisy_warnings # isort: skip
suppress_noisy_warnings()
@@ -482,9 +496,9 @@ class Scheduler(
attn_tp_cpu_group=self.attn_tp_cpu_group,
tp_cpu_group=self.tp_cpu_group,
attn_cp_cpu_group=self.attn_cp_cpu_group,
enable_metrics=self.server_args.enable_metrics,
enable_metrics=get_observability().enable_metrics,
enable_kv_cache_events=bool(
self.server_args.kv_events_config
get_observability().kv_events_config
and self.ps.pp_rank == 0
and self.ps.attn_tp_rank == 0
and self.ps.attn_cp_rank == 0
@@ -526,8 +540,8 @@ class Scheduler(
self.init_hisparse_coordinator()
if (
self.server_args.disaggregation_mode == "decode"
and self.server_args.disaggregation_decode_enable_offload_kvcache
get_disagg().disaggregation_mode == "decode"
and get_disagg().disaggregation_decode_enable_offload_kvcache
):
self.decode_offload_manager = DecodeKVCacheOffloadManager(
req_to_token_pool=self.req_to_token_pool,
@@ -642,7 +656,7 @@ class Scheduler(
self.dllm_config = ( # For diffusion LLM
DllmConfig.from_server_args(self.server_args)
if self.server_args.dllm_algorithm is not None
if get_exec().dllm.dllm_algorithm is not None
else None
)
@@ -671,10 +685,10 @@ class Scheduler(
port_args=port_args,
is_rank_zero=is_rank_zero,
skip_tokenizer_init=self.server_args.skip_tokenizer_init,
metrics_enabled=self.server_args.enable_metrics
metrics_enabled=get_observability().enable_metrics
and (
self.ps.attn_tp_rank == 0
or self.server_args.enable_metrics_for_all_schedulers
or get_observability().enable_metrics_for_all_schedulers
),
enable_scripted_runtime=envs.SGLANG_TEST_SCRIPTED_RUNTIME.get(),
)
@@ -693,7 +707,7 @@ class Scheduler(
port_args,
self.ps.dp_size,
dp_rank,
publish_interval=self.server_args.load_snapshot_publish_interval,
publish_interval=get_observability().load_snapshot_publish_interval,
)
except Exception as e:
logger.warning("load snapshot writer init failed: %s", e)
@@ -703,7 +717,7 @@ class Scheduler(
self.ps.pp_rank == 0
and self.ps.attn_tp_rank == 0
and self.ps.attn_cp_rank == 0
and self.server_args.sleep_on_idle
and get_device().sleep_on_idle
):
self.idle_sleeper = IdleSleeper(
sockets=[
@@ -737,22 +751,22 @@ class Scheduler(
else:
if self.model_config.is_multimodal:
self.processor = get_processor(
server_args.tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
use_fast=not server_args.disable_fast_image_processor,
tokenizer_backend=server_args.tokenizer_backend,
model_name=server_args.model_path,
get_serving().tokenizer_path,
tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=get_model().trust_remote_code,
revision=get_model().revision,
use_fast=not get_mm().disable_fast_image_processor,
tokenizer_backend=get_serving().tokenizer_backend,
model_name=get_model().model_path,
)
self.tokenizer = get_tokenizer_from_processor(self.processor)
else:
self.tokenizer = get_tokenizer(
server_args.tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
tokenizer_backend=server_args.tokenizer_backend,
get_serving().tokenizer_path,
tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=get_model().trust_remote_code,
revision=get_model().revision,
tokenizer_backend=get_serving().tokenizer_backend,
)
# Load multimodal processor for M-RoPE fallback computation.
@@ -774,9 +788,9 @@ class Scheduler(
)
# Set reasoning_parser and think_end_id if --reasoning_parser is enabled
if self.server_args.reasoning_parser and self.tokenizer:
if get_serving().reasoning_parser and self.tokenizer:
reasoning_parser = ReasoningParser(
model_type=self.server_args.reasoning_parser,
model_type=get_serving().reasoning_parser,
stream_reasoning=False,
tokenizer=self.tokenizer,
)
@@ -847,7 +861,7 @@ class Scheduler(
target_worker=self.tp_worker,
)
if self.server_args.speculative_draft_load_format is not None:
if get_spec().speculative_draft_load_format is not None:
# Write the draft load_format onto server_args (not just the bag):
# the draft worker is built from a copy of self.server_args and
# build_load_config reads server_args.load_format, so a bag-only
@@ -855,10 +869,10 @@ class Scheduler(
# format.
self.server_args.override(
"scheduler.draft_load_format",
load_format=self.server_args.speculative_draft_load_format,
load_format=get_spec().speculative_draft_load_format,
)
logger.info(
f"Using draft model load_format: '{self.server_args.speculative_draft_load_format}'"
f"Using draft model load_format: '{get_spec().speculative_draft_load_format}'"
)
DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args)
@@ -925,8 +939,8 @@ class Scheduler(
model_runner.post_capture_resize_kv_pool()
if (
self.server_args.elastic_ep_backend is not None
and self.server_args.ep_join_mode == "recover"
get_exec().moe.elastic_ep_backend is not None
and get_exec().moe.ep_join_mode == "recover"
):
model_runner.post_capture_elastic_ep_recover()
@@ -955,7 +969,7 @@ class Scheduler(
# --min-free-slots-delay. Built independently of the prefill delayer.
self.min_free_slots_delayer: Optional[MinFreeSlotsDelayer] = None
min_free_slots = resolve_min_free_slots(
self.server_args.min_free_slots_delay,
get_schedule().min_free_slots_delay,
self.max_running_requests,
is_dflash_family=self.spec_algorithm.is_dflash_family(),
)
@@ -1001,14 +1015,14 @@ class Scheduler(
if self.ps.tp_rank == 0:
logger.info(
f"max_total_num_tokens={self.max_total_num_tokens}, "
f"chunked_prefill_size={self.server_args.chunked_prefill_size}, "
f"chunked_prefill_size={get_schedule().chunked_prefill_size}, "
f"max_prefill_tokens={self.max_prefill_tokens}, "
f"max_running_requests={self.max_running_requests}, "
f"context_len={self.model_config.context_len}, "
f"{'available_cpu_mem' if self.device == 'cpu' else 'available_gpu_mem'}={avail_mem:.2f} GB"
)
if self.server_args.enable_metrics:
if get_observability().enable_metrics:
self.metrics_collector.emit_constants(
max_total_num_tokens=self.max_total_num_tokens,
# TODO: max_running_requests_under_SLO has no setter — dead chain.
@@ -1055,7 +1069,7 @@ class Scheduler(
self._engine_paused = False
def init_chunked_prefill(self):
self.chunked_prefill_size = self.server_args.chunked_prefill_size
self.chunked_prefill_size = get_schedule().chunked_prefill_size
uses_transformers_backend = (
get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS
)
@@ -1075,13 +1089,12 @@ class Scheduler(
self.chunked_req = None
self._pending_chunked_abort_req = None
self.is_mixed_chunk = (
self.chunked_prefill_size is not None
and self.server_args.enable_mixed_chunk
self.chunked_prefill_size is not None and get_schedule().enable_mixed_chunk
)
# Init the dynamic chunking predictor for PP
self.enable_dynamic_chunking = (
self.server_args.enable_dynamic_chunking and self.ps.pp_size > 1
get_schedule().enable_dynamic_chunking and self.ps.pp_size > 1
)
if self.enable_dynamic_chunking:
try:
@@ -1117,8 +1130,8 @@ class Scheduler(
)
self.prefill_delayer: Optional[PrefillDelayer] = None
self.max_prefill_bs: int = 0
if self.server_args.enable_prefill_delayer:
if self.server_args.disaggregation_mode == "decode":
if get_schedule().enable_prefill_delayer:
if get_disagg().disaggregation_mode == "decode":
logger.info(
"Ignoring --enable-prefill-delayer on decode engine "
"(no prefill scheduling path; delayer would be a no-op)."
@@ -1135,15 +1148,15 @@ class Scheduler(
if self.metrics_reporter.enable_metrics
else None
),
max_delay_passes=self.server_args.prefill_delayer_max_delay_passes,
token_usage_low_watermark=self.server_args.prefill_delayer_token_usage_low_watermark,
max_delay_passes=get_schedule().prefill_delayer_max_delay_passes,
token_usage_low_watermark=get_schedule().prefill_delayer_token_usage_low_watermark,
device=self.tp_group.device,
)
# NOTE: preemption is enabled by default for priority scheduling.
self.enable_priority_preemption = (
self.enable_priority_scheduling
and not self.server_args.disable_priority_preemption
and not get_schedule().disable_priority_preemption
)
self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args(
@@ -1159,12 +1172,12 @@ class Scheduler(
def init_watch_dog_memory_saver_input_blocker(self):
# Start watchdog thread
self.watchdog = create_scheduler_watchdog(
self, watchdog_timeout=self.server_args.watchdog_timeout
self, watchdog_timeout=get_device().watchdog_timeout
)
# Init memory saver, profiler and metric stats
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=self.server_args.enable_memory_saver
enable=get_exec().features.enable_memory_saver
)
# Init recv skipper and input blocker
@@ -1186,11 +1199,9 @@ class Scheduler(
self.disagg_decode_prealloc_queue = None
self.disagg_decode_transfer_queue = None
self.disaggregation_mode = DisaggregationMode(
self.server_args.disaggregation_mode
)
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
self.transfer_backend = TransferBackend(
self.server_args.disaggregation_transfer_backend
get_disagg().disaggregation_transfer_backend
)
# todo: should we fix this when enabling mtp or it doesn't matter since we only enable mtp in decode node thus we don't transfer draft kvs between P and D?
@@ -1260,10 +1271,10 @@ class Scheduler(
tp_size=self.ps.tp_size,
dp_size=self.server_args.dp_size,
gpu_id=self.ps.gpu_id,
bootstrap_port=self.server_args.disaggregation_bootstrap_port,
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
max_total_num_tokens=self.max_total_num_tokens,
pp_rank=self.ps.pp_rank,
num_reserved_decode_tokens=self.server_args.num_reserved_decode_tokens,
num_reserved_decode_tokens=get_disagg().num_reserved_decode_tokens,
transfer_backend=self.transfer_backend,
)
@@ -1289,7 +1300,7 @@ class Scheduler(
tp_rank=self.ps.tp_rank,
tp_size=self.ps.tp_size,
gpu_id=self.ps.gpu_id,
bootstrap_port=self.server_args.disaggregation_bootstrap_port,
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
gloo_group=self.attn_tp_cpu_group,
max_total_num_tokens=self.max_total_num_tokens,
scheduler=self,
@@ -1303,11 +1314,10 @@ class Scheduler(
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
# Init mm receiver for EPD disaggregation mode
if (
self.server_args.language_only
and self.server_args.encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
):
if get_disagg().language_only and get_disagg().encoder_transfer_backend in [
"zmq_to_scheduler",
"mooncake",
]:
self.mm_receiver = create_mm_receiver(
self.server_args,
dtype=self.model_config.dtype,
@@ -1388,7 +1398,7 @@ class Scheduler(
def init_deterministic_inference_config(self):
"""Initialize deterministic inference configuration for different attention backends."""
if not self.server_args.enable_deterministic_inference:
if not get_exec().deterministic.enable_deterministic_inference:
self.truncation_align_size = None
return
@@ -1794,10 +1804,10 @@ class Scheduler(
)
def init_lora_drainer(self) -> None:
if self.server_args.lora_drain_wait_threshold > 0.0:
if get_lora().lora_drain_wait_threshold > 0.0:
self.lora_drainer = LoRADrainer(
self.server_args.max_loras_per_batch,
self.server_args.lora_drain_wait_threshold,
get_lora().max_loras_per_batch,
get_lora().lora_drain_wait_threshold,
)
else:
self.lora_drainer = None
@@ -1923,7 +1933,7 @@ class Scheduler(
def init_kv_events_publisher(self) -> None:
self.kv_events_publisher = SchedulerKvEventsPublisher(
kv_events_config=self.server_args.kv_events_config,
kv_events_config=get_observability().kv_events_config,
ps=self.ps,
attn_tp_rank=self.ps.attn_tp_rank,
attn_cp_rank=self.ps.attn_cp_rank,
@@ -2107,7 +2117,7 @@ class Scheduler(
return image_inputs
def _get_multimodal_inputs(self, mm_inputs_dict):
if self.server_args.enable_broadcast_mm_inputs_process:
if get_mm().enable_broadcast_mm_inputs_process:
return self._process_and_broadcast_mm_inputs(mm_inputs_dict)
else:
return MultimodalInputs.from_processor_output(mm_inputs_dict)
@@ -2154,7 +2164,7 @@ class Scheduler(
def _maybe_namespace_elastic_radix_cache(self, req: Req) -> None:
if (
self.server_args.elastic_ep_backend is None
get_exec().moe.elastic_ep_backend is None
or self.disable_radix_cache
or not self.tree_cache.is_tree_cache()
):
@@ -2200,8 +2210,7 @@ class Scheduler(
)
# Radix-native sessions use only the top-level session_id.
radix_native_session = (
recv_req.session_id is not None
and self.server_args.enable_session_radix_cache
recv_req.session_id is not None and get_memory().enable_session_radix_cache
)
if session_id is None or radix_native_session:
@@ -2213,7 +2222,7 @@ class Scheduler(
if recv_req.bootstrap_port is None:
# Use default bootstrap port
recv_req.bootstrap_port = self.server_args.disaggregation_bootstrap_port
recv_req.bootstrap_port = get_disagg().disaggregation_bootstrap_port
req = Req(
recv_req.rid,
@@ -2366,7 +2375,7 @@ class Scheduler(
self._add_request_to_queue(req)
return
if req.return_sampling_mask and self.server_args.sampling_backend == "ascend":
if req.return_sampling_mask and get_exec().kernel.sampling_backend == "ascend":
# The ascend backend samples from logits directly and never builds the
# top-k/top-p support, so it cannot produce a sampling mask.
error_msg = (
@@ -2415,7 +2424,7 @@ class Scheduler(
error_msg = validate_input_length(
req,
self.max_req_input_len,
self.server_args.allow_auto_truncate,
get_serving().allow_auto_truncate,
)
if error_msg:
req.set_finish_with_abort(error_msg)
@@ -2693,7 +2702,7 @@ class Scheduler(
error_msg = validate_input_length(
req,
self.max_req_input_len,
self.server_args.allow_auto_truncate,
get_serving().allow_auto_truncate,
)
if error_msg:
self._add_request_to_queue(req)
@@ -2905,7 +2914,7 @@ class Scheduler(
if (
need_mlp_sync
and not self.spec_algorithm.is_none()
and not self.server_args.speculative_skip_dp_mlp_sync
and not get_spec().speculative_skip_dp_mlp_sync
):
# NOTE: This branch makes sure prefill and decode batches will not be mixed when spec and dp-attn is enabled.
# Before merging the new batch into running batch:
@@ -2979,7 +2988,7 @@ class Scheduler(
for req in ready_grammar_requests:
self._add_request_to_queue(req)
if self.enable_hierarchical_cache or self.server_args.enable_flexkv:
if self.enable_hierarchical_cache or get_memory().enable_flexkv:
self.tree_cache.check_hicache_events()
if self.enable_priority_preemption or self.is_hybrid_swa:
@@ -3046,7 +3055,7 @@ class Scheduler(
self.priority_scheduling_preemption_threshold,
max_prefill_bs=self.max_prefill_bs,
max_running_requests=self.max_running_requests,
prefill_max_requests=self.server_args.prefill_max_requests,
prefill_max_requests=get_schedule().prefill_max_requests,
prefill_delayer_single_pass=prefill_delayer_single_pass,
dllm_config=self.dllm_config,
waiting_queue_len=len(self.waiting_queue),
@@ -3619,7 +3628,7 @@ class Scheduler(
def _maybe_report_active_ranks(self) -> None:
if not (
self.enable_dp_attention and self.server_args.elastic_ep_backend is not None
self.enable_dp_attention and get_exec().moe.elastic_ep_backend is not None
):
return
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
@@ -3924,7 +3933,7 @@ class Scheduler(
ok, msg = self.tree_cache.attach_storage_backend(
storage_backend=recv_req.hicache_storage_backend,
storage_backend_extra_config_json=recv_req.hicache_storage_backend_extra_config_json,
served_model_name=self.server_args.served_model_name,
served_model_name=get_serving().served_model_name,
hicache_storage_prefetch_policy=recv_req.hicache_storage_prefetch_policy,
hicache_write_policy=recv_req.hicache_write_policy,
)
@@ -4044,7 +4053,7 @@ class Scheduler(
}
ret["effective_max_running_requests_per_dp"] = self.max_running_requests
if self.server_args.elastic_ep_backend is not None:
if get_exec().moe.elastic_ep_backend is not None:
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
ret["is_scaling_elastic_ep"] = ElasticEPStateManager.is_scaling()
@@ -4583,10 +4592,10 @@ class Scheduler(
return None
def close_session(self, recv_req: CloseSessionReqInput):
if self.server_args.enable_session_radix_cache:
if get_memory().enable_session_radix_cache:
self.tree_cache.release_radix_session(recv_req.session_id)
if recv_req.session_id in self.session_controller or not (
self.server_args.enable_session_radix_cache
get_memory().enable_session_radix_cache
):
self.session_controller.close(recv_req)
@@ -27,7 +27,13 @@ from sglang.srt.mem_cache.common import (
maybe_cache_unfinished_req,
release_kv_cache,
)
from sglang.srt.runtime_context import get_server_args
from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_memory,
get_observability,
get_server_args,
)
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
@@ -84,7 +90,7 @@ class SchedulerBatchResultProcessor:
def process_batch_result_prebuilt(self, batch: ScheduleBatch):
assert self.disaggregation_mode == DisaggregationMode.DECODE
use_free_group = self.server_args.disaggregation_decode_enable_radix_cache
use_free_group = get_disagg().disaggregation_decode_enable_radix_cache
if use_free_group:
self.token_to_kv_pool_allocator.free_group_begin()
for req in batch.reqs:
@@ -92,7 +98,7 @@ class SchedulerBatchResultProcessor:
req.update_finish_state()
if req.finished():
req.time_stats.set_quick_finish_time()
if self.server_args.enable_hisparse:
if get_memory().enable_hisparse:
self.hisparse_coordinator.request_finished(req)
release_kv_cache(req, self.tree_cache)
@@ -243,7 +249,7 @@ class SchedulerBatchResultProcessor:
req.time_stats.set_completion_time()
elif not batch.decoding_reqs or req not in batch.decoding_reqs:
maybe_cache_unfinished_req(req, self.tree_cache)
if self.server_args.enable_hisparse:
if get_memory().enable_hisparse:
self.hisparse_coordinator.admit_request_into_staging(req)
self._maybe_collect_customized_info(i, req, logits_output)
@@ -756,7 +762,7 @@ class SchedulerBatchResultProcessor:
num_block_accept_tokens=result.num_block_accept_tokens,
num_cap_tokens=result.num_cap_tokens,
)
if self.server_args.enable_metrics:
if get_observability().enable_metrics:
self.metrics_collector.increment_decode_cuda_graph_pass(
value=can_run_cuda_graph
)
@@ -939,7 +945,7 @@ class SchedulerBatchResultProcessor:
self._mamba_prefix_cache_update(req, batch, result, i)
if (
self.server_args.disaggregation_decode_enable_offload_kvcache
get_disagg().disaggregation_decode_enable_offload_kvcache
and not req.finished()
):
self.decode_offload_manager.offload_kv_cache(req)
@@ -959,12 +965,12 @@ class SchedulerBatchResultProcessor:
self._maybe_collect_routed_experts(req)
self._maybe_collect_indexer_topk(req)
if self.server_args.disaggregation_decode_enable_offload_kvcache:
if get_disagg().disaggregation_decode_enable_offload_kvcache:
# Asynchronously offload KV cache; release_kv_cache will be called after Device->Host transfer completes
if not self.decode_offload_manager.offload_kv_cache(req):
self.decode_offload_manager.finalize_release_on_finish(req)
else:
if self.server_args.enable_hisparse:
if get_memory().enable_hisparse:
self.hisparse_coordinator.request_finished(req)
prepare_release = getattr(
self.model_worker, "prepare_for_kv_cache_release", None
@@ -1063,7 +1069,7 @@ class SchedulerBatchResultProcessor:
other_idx
].item() == -1 and mamba_lazy_spec_in_window(
req,
server_args.mamba_track_interval,
get_exec().mamba.mamba_track_interval,
server_args.max_speculative_num_draft_tokens,
)
if (
@@ -1102,7 +1108,7 @@ class SchedulerBatchResultProcessor:
For spec decode, the boundary is detected by comparing the
accepted seq_len range against interval boundaries.
"""
interval = get_server_args().mamba_track_interval
interval = get_exec().mamba.mamba_track_interval
if batch.spec_algorithm.is_none():
if req.kv_committed_len % interval == 0:
@@ -26,6 +26,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
)
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.metrics_collector import DPCooperationInfo
from sglang.srt.runtime_context import get_schedule
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils.common import require_mlp_tp_gather
@@ -385,7 +386,7 @@ class SchedulerDPAttnAdapter:
get_idle_batch=self.get_idle_batch,
disable_cuda_graph=cuda_graph_fully_disabled(),
require_mlp_tp_gather=require_mlp_tp_gather(self.server_args),
disable_overlap_schedule=self.server_args.disable_overlap_schedule,
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
offload_tags=self.offload_tags,
dwdp=self.server_args.dwdp_size > 1,
)
@@ -14,6 +14,7 @@ from sglang.srt.managers.load_snapshot import (
QueueMetrics,
SpeculativeMetrics,
)
from sglang.srt.runtime_context import get_lora
if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
@@ -155,7 +156,7 @@ class SchedulerLoadInquirer:
)
lora = None
if self.server_args.enable_lora:
if get_lora().enable_lora:
lora = LoRAMetrics(
slots_used=stats.lora_pool_slots_used,
slots_total=stats.lora_pool_slots_total,
@@ -11,6 +11,7 @@ import torch
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.runtime_context import get_exec
from sglang.srt.server_args import (
MIS_DELIMITER_TOKEN_ID,
ServerArgs,
@@ -164,7 +165,7 @@ class SchedulerLogprobResultProcessor:
delimiter token receive logprobs.
"""
return (
self.server_args.enable_mis
get_exec().features.enable_mis
and req.is_prefill_only
and req.multi_item_delimiter_indices is not None
)
@@ -26,6 +26,7 @@ from sglang.srt.observability.metrics_collector import (
SchedulerStats,
compute_routing_key_stats,
)
from sglang.srt.runtime_context import get_spec
from sglang.srt.utils.device_timer import DeviceTimer
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
@@ -764,12 +765,10 @@ class SchedulerMetricsReporter:
else:
spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct
num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct
if self.scheduler.server_args.speculative_num_draft_tokens:
draft_per_round = (
self.scheduler.server_args.speculative_num_draft_tokens - 1
)
if get_spec().speculative_num_draft_tokens:
draft_per_round = get_spec().speculative_num_draft_tokens - 1
else:
draft_per_round = self.scheduler.server_args.speculative_num_steps or 0
draft_per_round = get_spec().speculative_num_steps or 0
total_draft_tokens = self.spec_num_forward_ct * draft_per_round
spec_accept_rate = (
num_correct_drafts / total_draft_tokens if total_draft_tokens > 0 else 0
@@ -27,6 +27,7 @@ from sglang.srt.managers.schedule_batch import (
Req,
)
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.runtime_context import get_observability, get_serving
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -153,7 +154,7 @@ class SchedulerOutputStreamer:
return_sampling_mask=return_sampling_mask,
spec_algorithm=self.spec_algorithm,
disaggregation_mode=self.disaggregation_mode,
default_stream_interval=self.server_args.stream_interval,
default_stream_interval=get_serving().stream_interval,
default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL,
get_cached_tokens_details=self.get_cached_tokens_details,
rust_server_mode=self.rust_server is not None,
@@ -184,7 +185,7 @@ class SchedulerOutputStreamer:
if (
req.finished()
and self.ps.attn_tp_rank == 0
and self.server_args.enable_request_time_stats_logging
and get_observability().enable_request_time_stats_logging
):
req.log_time_stats()
@@ -19,7 +19,7 @@ from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_server_args
from sglang.srt.runtime_context import get_device
from sglang.srt.utils import is_mps, is_npu
from sglang.srt.utils.profile_merger import ProfileMerger
from sglang.srt.utils.profile_utils import ProfileManager
@@ -257,7 +257,7 @@ class SchedulerProfilerManager:
self.profile_in_progress = True
if "CUDA_PROFILER" in activities:
if self.ps.gpu_id == get_server_args().base_gpu_id:
if self.ps.gpu_id == get_device().base_gpu_id:
torch.cuda.cudart().cudaProfilerStart()
self.profile_in_progress = True
@@ -368,7 +368,7 @@ class SchedulerProfilerManager:
torch.cuda.memory._record_memory_history(enabled=None)
if "CUDA_PROFILER" in self.profiler_activities:
if self.ps.gpu_id == get_server_args().base_gpu_id:
if self.ps.gpu_id == get_device().base_gpu_id:
torch.cuda.cudart().cudaProfilerStop()
merge_message = self._merge_profile_traces()
@@ -27,6 +27,7 @@ from sglang.srt.managers.mm_utils import (
has_shm_features,
unwrap_shm_features,
)
from sglang.srt.runtime_context import get_disagg
from sglang.srt.utils import (
broadcast_pyobj,
point_to_point_pyobj,
@@ -231,8 +232,8 @@ class SchedulerRequestReceiver:
# Process MM requests under EPD-disaggregation mode
if (
self.ps.pp_rank == 0
and self.server_args.language_only
and self.server_args.encoder_transfer_backend
and get_disagg().language_only
and get_disagg().encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
):
recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
@@ -36,6 +36,7 @@ from sglang.srt.model_executor.forward_batch_info import (
PPProxyTensors,
)
from sglang.srt.observability.req_time_stats import set_time_batch
from sglang.srt.runtime_context import get_disagg
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj
from sglang.srt.utils.common import get_device_module, is_xpu
@@ -479,7 +480,7 @@ class SchedulerPPMixin:
)
)
if self.server_args.disaggregation_decode_enable_offload_kvcache:
if get_disagg().disaggregation_decode_enable_offload_kvcache:
self.decode_offload_manager.check_offload_progress()
if rmbs[next_mb_id] is not None:
@@ -549,7 +550,7 @@ class SchedulerPPMixin:
+ len(self.disagg_decode_transfer_queue.queue)
+ len(self.disagg_decode_prealloc_queue.queue)
)
if self.server_args.disaggregation_decode_enable_offload_kvcache:
if get_disagg().disaggregation_decode_enable_offload_kvcache:
queue_size += len(self.decode_offload_manager.ongoing_offload)
if server_is_idle and queue_size == 0:
+11 -10
View File
@@ -47,6 +47,7 @@ from sglang.srt.model_executor.forward_batch_info import (
PPProxyTensors,
)
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.runtime_context import get_exec, get_model, get_schedule, get_spec
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed
from sglang.srt.utils.hf_transformers_utils import (
@@ -408,14 +409,14 @@ class TpModelWorker(BaseTpWorker):
self.model_config = ModelConfig.from_server_args(
self.server_args,
model_path=(
self.server_args.model_path
get_model().model_path
if not self.is_draft_worker
else self.server_args.speculative_draft_model_path
else get_spec().speculative_draft_model_path
),
model_revision=(
self.server_args.revision
get_model().revision
if not self.is_draft_worker
else self.server_args.speculative_draft_model_revision
else get_spec().speculative_draft_model_revision
),
is_draft_model=self.is_draft_worker,
context_length=self.context_length,
@@ -426,7 +427,7 @@ class TpModelWorker(BaseTpWorker):
self._model_runner = ModelRunner(
model_config=self.model_config,
mem_fraction_static=self.server_args.mem_fraction_static,
mem_fraction_static=get_schedule().mem_fraction_static,
gpu_id=self.gpu_id,
ps=self.ps,
nccl_port=self.nccl_port,
@@ -442,11 +443,11 @@ class TpModelWorker(BaseTpWorker):
from sglang.srt.model_executor.model_runner import ModelRunner
self.model_runner_list.append(self.model_runner)
for i in range(1, self.server_args.speculative_num_steps):
for i in range(1, get_spec().speculative_num_steps):
self.model_runner_list.append(
ModelRunner(
model_config=self.model_config,
mem_fraction_static=self.server_args.mem_fraction_static,
mem_fraction_static=get_schedule().mem_fraction_static,
gpu_id=self.gpu_id,
ps=self.ps,
nccl_port=self.nccl_port,
@@ -462,7 +463,7 @@ class TpModelWorker(BaseTpWorker):
def _init_dllm_algorithm(self):
from sglang.srt.dllm.algorithm.base import DllmAlgorithm
if self.server_args.dllm_algorithm is not None:
if get_exec().dllm.dllm_algorithm is not None:
self.dllm_algorithm = DllmAlgorithm.from_server_args(self.server_args)
else:
self.dllm_algorithm = None
@@ -488,9 +489,9 @@ class TpModelWorker(BaseTpWorker):
)
return (
self.model_runner.max_total_num_tokens,
self.server_args.max_prefill_tokens,
get_schedule().max_prefill_tokens,
self.model_runner.max_running_requests,
self.server_args.max_queued_requests,
get_schedule().max_queued_requests,
max_req_len,
max_req_len - 5,
self.random_seed,