config: read resolved config via namespace accessors (#31814)
This commit is contained in:
@@ -48,6 +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 get_exec
|
||||
from sglang.srt.server_args import (
|
||||
DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
||||
PortArgs,
|
||||
@@ -231,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; "
|
||||
@@ -484,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(
|
||||
|
||||
@@ -33,7 +33,13 @@ 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,
|
||||
get_serving,
|
||||
)
|
||||
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
|
||||
@@ -878,7 +884,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"
|
||||
@@ -1287,7 +1293,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
|
||||
)
|
||||
@@ -1967,7 +1973,7 @@ def wrap_shm_features(obj):
|
||||
"""
|
||||
Scan the object for multimodal tensors and wrap them in SHM pointers.
|
||||
"""
|
||||
if _get_is_default_transport() or get_server_args().skip_tokenizer_init:
|
||||
if _get_is_default_transport() or get_serving().skip_tokenizer_init:
|
||||
return obj
|
||||
|
||||
if obj.mm_inputs:
|
||||
@@ -2028,7 +2034,7 @@ def unwrap_shm_features(obj):
|
||||
Restore ShmPointerMMData wrappers back into standard torch.Tensors.
|
||||
Handles both single requests and batch requests.
|
||||
"""
|
||||
if _get_is_default_transport() or get_server_args().skip_tokenizer_init:
|
||||
if _get_is_default_transport() or get_serving().skip_tokenizer_init:
|
||||
return obj
|
||||
# Handle batch requests
|
||||
if isinstance(obj, BaseBatchReq):
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sglang.srt.runtime_context import get_disagg
|
||||
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
@@ -645,15 +647,15 @@ class TokenizerWorker(TokenizerManager):
|
||||
self.tokenizer_ipc_name = port_args.tokenizer_ipc_name
|
||||
|
||||
# For PD disaggregtion
|
||||
self.server_args.override(
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
"tokenizer_worker.restore_disaggregation_mode",
|
||||
disaggregation_mode=disaggregation_mode,
|
||||
)
|
||||
self.disaggregation_mode = DisaggregationMode(
|
||||
self.server_args.disaggregation_mode
|
||||
)
|
||||
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||
self.disaggregation_transfer_backend = TransferBackend(
|
||||
self.server_args.disaggregation_transfer_backend
|
||||
get_disagg().disaggregation_transfer_backend
|
||||
)
|
||||
|
||||
# Register this worker with the router for pause/continue broadcasting
|
||||
|
||||
@@ -77,10 +77,7 @@ from sglang.srt.managers.embed_types import PositionalEmbeds
|
||||
from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import (
|
||||
NewTokenRatioTracker,
|
||||
)
|
||||
from sglang.srt.mem_cache.allocation import (
|
||||
alloc_for_decode,
|
||||
alloc_for_extend,
|
||||
)
|
||||
from sglang.srt.mem_cache.allocation import alloc_for_decode, alloc_for_extend
|
||||
from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
@@ -105,7 +102,12 @@ from sglang.srt.observability.req_time_stats import (
|
||||
DPControllerReqTimeStats,
|
||||
SchedulerReqTimeStats,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_serving,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -1094,7 +1096,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
|
||||
@@ -1115,7 +1117,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
|
||||
|
||||
|
||||
@@ -56,7 +56,7 @@ 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.runtime_context import get_disagg
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -193,7 +193,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)
|
||||
|
||||
@@ -210,9 +210,7 @@ from sglang.srt.managers.scheduler_components.pool_stats_observer import (
|
||||
from sglang.srt.managers.scheduler_components.profiler_manager import (
|
||||
SchedulerProfilerManager,
|
||||
)
|
||||
from sglang.srt.managers.scheduler_components.recv_skipper import (
|
||||
SchedulerRecvSkipper,
|
||||
)
|
||||
from sglang.srt.managers.scheduler_components.recv_skipper import SchedulerRecvSkipper
|
||||
from sglang.srt.managers.scheduler_components.request_receiver import (
|
||||
SchedulerRequestReceiver,
|
||||
)
|
||||
@@ -241,7 +239,20 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
|
||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.plugins import load_plugins
|
||||
from sglang.srt.runtime_context import get_context, get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_context,
|
||||
get_device,
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_lora,
|
||||
get_memory,
|
||||
get_mm,
|
||||
get_observability,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_serving,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
@@ -443,9 +454,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
|
||||
@@ -471,8 +482,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,
|
||||
@@ -583,7 +594,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
|
||||
)
|
||||
|
||||
@@ -611,11 +622,11 @@ class Scheduler(
|
||||
self.ipc_channels = SchedulerIpcChannels.create(
|
||||
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
|
||||
skip_tokenizer_init=get_serving().skip_tokenizer_init,
|
||||
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(),
|
||||
)
|
||||
@@ -631,7 +642,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)
|
||||
@@ -641,7 +652,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=[
|
||||
@@ -712,9 +723,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,
|
||||
)
|
||||
@@ -785,7 +796,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
|
||||
@@ -793,10 +804,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)
|
||||
@@ -887,7 +898,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(),
|
||||
)
|
||||
@@ -933,14 +944,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.
|
||||
@@ -987,7 +998,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
|
||||
)
|
||||
@@ -1007,13 +1018,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:
|
||||
@@ -1049,8 +1059,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)."
|
||||
@@ -1067,15 +1077,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(
|
||||
@@ -1091,12 +1101,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
|
||||
@@ -1118,11 +1128,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?
|
||||
@@ -1192,10 +1200,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,
|
||||
)
|
||||
|
||||
@@ -1221,7 +1229,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,
|
||||
@@ -1235,11 +1243,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,
|
||||
@@ -1320,7 +1327,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
|
||||
|
||||
@@ -1329,7 +1336,7 @@ class Scheduler(
|
||||
"triton": ("SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE", 4096),
|
||||
}
|
||||
env_var, default_size = backend_sizes.get(
|
||||
self.server_args.attention_backend, (None, None)
|
||||
get_exec().kernel.attention_backend, (None, None)
|
||||
)
|
||||
self.truncation_align_size = (
|
||||
get_int_env_var(env_var, default_size) if env_var else None
|
||||
@@ -1725,10 +1732,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
|
||||
@@ -1830,7 +1837,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,
|
||||
@@ -2006,7 +2013,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)
|
||||
@@ -2053,7 +2060,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()
|
||||
):
|
||||
@@ -2099,8 +2106,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:
|
||||
@@ -2112,7 +2118,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,
|
||||
@@ -2265,7 +2271,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 = (
|
||||
@@ -2314,7 +2320,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)
|
||||
@@ -2592,7 +2598,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)
|
||||
@@ -2804,7 +2810,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:
|
||||
@@ -2878,7 +2884,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:
|
||||
@@ -2945,7 +2951,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),
|
||||
@@ -3516,7 +3522,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
|
||||
@@ -3792,7 +3798,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,
|
||||
)
|
||||
@@ -3912,7 +3918,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()
|
||||
@@ -4445,10 +4451,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)
|
||||
|
||||
|
||||
@@ -2,14 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Callable,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
@@ -23,11 +16,14 @@ from sglang.srt.managers.schedule_batch import (
|
||||
ScheduleBatch,
|
||||
mamba_lazy_spec_in_window,
|
||||
)
|
||||
from sglang.srt.mem_cache.common import (
|
||||
maybe_cache_unfinished_req,
|
||||
release_kv_cache,
|
||||
from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_memory,
|
||||
get_observability,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.runtime_context import 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
|
||||
@@ -48,10 +44,7 @@ if TYPE_CHECKING:
|
||||
SchedulerOutputStreamer,
|
||||
)
|
||||
from sglang.srt.managers.tp_worker import BaseTpWorker
|
||||
from sglang.srt.managers.utils import (
|
||||
EmbeddingBatchResult,
|
||||
GenerationBatchResult,
|
||||
)
|
||||
from sglang.srt.managers.utils import EmbeddingBatchResult, GenerationBatchResult
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
@@ -84,7 +77,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 +85,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 +236,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 +749,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 +932,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 +952,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
|
||||
@@ -1102,7 +1095,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:
|
||||
|
||||
@@ -12,9 +12,7 @@ from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import world_dp_gather_enabled
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.managers.scheduler_components.recv_skipper import (
|
||||
SchedulerRecvSkipper,
|
||||
)
|
||||
from sglang.srt.managers.scheduler_components.recv_skipper import SchedulerRecvSkipper
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
@@ -26,6 +24,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 +384,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
|
||||
@@ -144,7 +145,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,
|
||||
|
||||
@@ -1,20 +1,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import (
|
||||
List,
|
||||
Tuple,
|
||||
)
|
||||
from typing import List, Tuple
|
||||
|
||||
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.server_args import (
|
||||
MIS_DELIMITER_TOKEN_ID,
|
||||
ServerArgs,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID, ServerArgs
|
||||
|
||||
|
||||
@dataclass(kw_only=True, slots=True, frozen=True)
|
||||
@@ -164,7 +159,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
|
||||
)
|
||||
|
||||
@@ -2,12 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
List,
|
||||
Optional,
|
||||
)
|
||||
from typing import Any, Callable, List, Optional
|
||||
|
||||
import torch
|
||||
import zmq
|
||||
@@ -21,11 +16,9 @@ from sglang.srt.managers.io_struct import (
|
||||
CachedTokensDetails,
|
||||
wrap_as_pickle,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import (
|
||||
BaseFinishReason,
|
||||
Req,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import BaseFinishReason, 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
|
||||
|
||||
@@ -144,7 +137,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,
|
||||
)
|
||||
@@ -171,7 +164,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()
|
||||
|
||||
|
||||
@@ -5,13 +5,7 @@ import os
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
List,
|
||||
Optional,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Callable, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
@@ -19,7 +13,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
|
||||
@@ -255,7 +249,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
|
||||
|
||||
@@ -365,7 +359,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()
|
||||
|
||||
@@ -2,14 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from http import HTTPStatus
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
List,
|
||||
Optional,
|
||||
Union,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Callable, List, Optional, Union
|
||||
|
||||
import zmq
|
||||
from torch.distributed import barrier
|
||||
@@ -22,14 +15,9 @@ from sglang.srt.managers.io_struct import (
|
||||
TokenizedGenerateReqInput,
|
||||
sock_recv,
|
||||
)
|
||||
from sglang.srt.managers.mm_utils import (
|
||||
has_shm_features,
|
||||
unwrap_shm_features,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
broadcast_pyobj,
|
||||
point_to_point_pyobj,
|
||||
)
|
||||
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
|
||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -220,8 +208,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:
|
||||
|
||||
@@ -74,6 +74,7 @@ from sglang.srt.managers.io_struct import (
|
||||
UpdateWeightsFromTensorReqOutput,
|
||||
)
|
||||
from sglang.srt.managers.load_snapshot import LoadSnapshot
|
||||
from sglang.srt.runtime_context import get_lora
|
||||
from sglang.srt.server_args import LoRARef, ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
@@ -569,7 +570,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
|
||||
try:
|
||||
if not self.server_args.enable_lora:
|
||||
if not get_lora().enable_lora:
|
||||
raise ValueError(
|
||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||
)
|
||||
@@ -602,10 +603,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
|
||||
@@ -619,7 +620,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(
|
||||
@@ -647,7 +648,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
|
||||
try:
|
||||
if not self.server_args.enable_lora:
|
||||
if not get_lora().enable_lora:
|
||||
raise ValueError(
|
||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||
)
|
||||
@@ -672,10 +673,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
|
||||
@@ -689,7 +690,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(
|
||||
@@ -717,7 +718,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
|
||||
try:
|
||||
if not self.server_args.enable_lora:
|
||||
if not get_lora().enable_lora:
|
||||
raise ValueError(
|
||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||
)
|
||||
@@ -893,6 +894,8 @@ class TokenizerControlMixin:
|
||||
) -> None:
|
||||
"""Update weight version if provided."""
|
||||
if weight_version is not None:
|
||||
self.server_args.override(
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
"tokenizer.weight_version", weight_version=weight_version
|
||||
)
|
||||
|
||||
@@ -110,6 +110,14 @@ from sglang.srt.observability.request_metrics_exporter import (
|
||||
RequestMetricsExporterManager,
|
||||
)
|
||||
from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_disagg,
|
||||
get_lora,
|
||||
get_model,
|
||||
get_observability,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.server_args import (
|
||||
PortArgs,
|
||||
@@ -463,10 +471,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# TODO: Refactor and organize the log export code.
|
||||
# Request logging
|
||||
self.request_logger = RequestLogger(
|
||||
log_requests=self.server_args.log_requests,
|
||||
log_requests_level=self.server_args.log_requests_level,
|
||||
log_requests_format=self.server_args.log_requests_format,
|
||||
log_requests_target=self.server_args.log_requests_target,
|
||||
log_requests=get_observability().log_requests,
|
||||
log_requests_level=get_observability().log_requests_level,
|
||||
log_requests_format=get_observability().log_requests_format,
|
||||
log_requests_target=get_observability().log_requests_target,
|
||||
)
|
||||
|
||||
# Dumping
|
||||
@@ -489,7 +497,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
def init_weight_update(self):
|
||||
# Initial weights status
|
||||
self.initial_weights_loaded = True
|
||||
if self.server_args.checkpoint_engine_wait_weights_before_ready:
|
||||
if get_model().checkpoint_engine_wait_weights_before_ready:
|
||||
self.initial_weights_loaded = False
|
||||
|
||||
# Weight updates
|
||||
@@ -509,7 +517,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# The registry dynamically updates as adapters are loaded / unloaded during runtime. It
|
||||
# serves as the source of truth for available adapters and maps user-friendly LoRA names
|
||||
# to internally used unique LoRA IDs.
|
||||
self.lora_registry = LoRARegistry(self.server_args.lora_paths)
|
||||
self.lora_registry = LoRARegistry(get_lora().lora_paths)
|
||||
# Lock to serialize LoRA update operations.
|
||||
# Please note that, unlike `model_update_lock`, this does not block inference, allowing
|
||||
# LoRA updates and inference to overlap.
|
||||
@@ -518,15 +526,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# point to their latest LoRARef objects, so that they can be
|
||||
# dynamically loaded if needed for inference
|
||||
self.lora_ref_cache: Dict[str, LoRARef] = {}
|
||||
if self.server_args.lora_paths is not None:
|
||||
for lora_ref in self.server_args.lora_paths:
|
||||
if get_lora().lora_paths is not None:
|
||||
for lora_ref in get_lora().lora_paths:
|
||||
self.lora_ref_cache[lora_ref.lora_name] = lora_ref
|
||||
|
||||
def init_disaggregation(self):
|
||||
# PD Disaggregation
|
||||
self.disaggregation_mode = DisaggregationMode(
|
||||
self.server_args.disaggregation_mode
|
||||
)
|
||||
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||
# Keep a reference so the bootstrap server is not garbage-collected.
|
||||
self.bootstrap_server = start_disagg_service(self.server_args)
|
||||
# Single-source counter for auto-assigning fake bootstrap_room.
|
||||
@@ -535,18 +541,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# Encoder Disaggregation
|
||||
self.encoder_bootstrap_server = None
|
||||
if self.server_args.language_only:
|
||||
from sglang.srt.disaggregation.encode_receiver import (
|
||||
EncoderBootstrapServer,
|
||||
)
|
||||
from sglang.srt.disaggregation.encode_receiver import EncoderBootstrapServer
|
||||
|
||||
# Shared mutable URL list: the bootstrap server appends / removes
|
||||
# entries as encoders register, the receiver reads from the same
|
||||
# list. Pre-populated with static --encoder-urls so the legacy
|
||||
# CLI flag still works (alongside dynamic registrations).
|
||||
self.encoder_urls: List[str] = list(self.server_args.encoder_urls)
|
||||
self.encoder_urls: List[str] = list(get_disagg().encoder_urls)
|
||||
self.encoder_bootstrap_server = EncoderBootstrapServer(
|
||||
host=self.server_args.host,
|
||||
port=self.server_args.encoder_bootstrap_port,
|
||||
host=get_serving().host,
|
||||
port=get_disagg().encoder_bootstrap_port,
|
||||
urls=self.encoder_urls,
|
||||
)
|
||||
self.mm_receiver = create_mm_receiver(
|
||||
@@ -560,20 +564,22 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# Metrics
|
||||
if self.enable_metrics:
|
||||
engine_type = DisaggregationMode.to_engine_type(
|
||||
self.server_args.disaggregation_mode
|
||||
get_disagg().disaggregation_mode
|
||||
)
|
||||
|
||||
labels = {
|
||||
"model_name": self.server_args.served_model_name,
|
||||
"model_name": get_serving().served_model_name,
|
||||
"engine_type": engine_type,
|
||||
}
|
||||
if self.enable_priority_scheduling:
|
||||
labels["priority"] = ""
|
||||
if self.server_args.tokenizer_metrics_allowed_custom_labels:
|
||||
for label in self.server_args.tokenizer_metrics_allowed_custom_labels:
|
||||
if get_observability().tokenizer_metrics_allowed_custom_labels:
|
||||
for (
|
||||
label
|
||||
) in get_observability().tokenizer_metrics_allowed_custom_labels:
|
||||
labels[label] = ""
|
||||
if self.server_args.extra_metric_labels:
|
||||
labels.update(self.server_args.extra_metric_labels)
|
||||
if get_observability().extra_metric_labels:
|
||||
labels.update(get_observability().extra_metric_labels)
|
||||
tokenizer_collector_cls = resolve_collector_class(
|
||||
self.server_args,
|
||||
STAT_LOGGER_ROLE_TOKENIZER,
|
||||
@@ -582,18 +588,18 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
self.metrics_collector = tokenizer_collector_cls(
|
||||
server_args=self.server_args,
|
||||
labels=labels,
|
||||
bucket_time_to_first_token=self.server_args.bucket_time_to_first_token,
|
||||
bucket_e2e_request_latency=self.server_args.bucket_e2e_request_latency,
|
||||
bucket_inter_token_latency=self.server_args.bucket_inter_token_latency,
|
||||
bucket_time_to_first_token=get_observability().bucket_time_to_first_token,
|
||||
bucket_e2e_request_latency=get_observability().bucket_e2e_request_latency,
|
||||
bucket_inter_token_latency=get_observability().bucket_inter_token_latency,
|
||||
)
|
||||
|
||||
start_cpu_monitor_thread("tokenizer")
|
||||
|
||||
if self.server_args.gc_warning_threshold_secs > 0.0:
|
||||
configure_gc_warning(self.server_args.gc_warning_threshold_secs)
|
||||
if get_observability().gc_warning_threshold_secs > 0.0:
|
||||
configure_gc_warning(get_observability().gc_warning_threshold_secs)
|
||||
self.soft_watchdog = Watchdog.create(
|
||||
debug_name="TokenizerManager",
|
||||
watchdog_timeout=self.server_args.soft_watchdog_timeout,
|
||||
watchdog_timeout=get_device().soft_watchdog_timeout,
|
||||
soft=True,
|
||||
test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(),
|
||||
)
|
||||
@@ -1757,7 +1763,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
|
||||
# default the load format to the server_args
|
||||
if obj.load_format is None:
|
||||
obj.load_format = self.server_args.load_format
|
||||
obj.load_format = get_model().load_format
|
||||
logger.info("Start update_weights. Load format=%s", obj.load_format)
|
||||
|
||||
if obj.abort_all_requests:
|
||||
@@ -1783,7 +1789,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
|
||||
def _update_model_path_info(self, model_path: str, load_format: str):
|
||||
self.served_model_name = model_path
|
||||
self.server_args.override(
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
"tokenizer.update_weights", model_path=model_path, load_format=load_format
|
||||
)
|
||||
self.model_path = model_path
|
||||
@@ -1927,7 +1935,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
"id": rid,
|
||||
"finish_reason": recv_obj.finished_reasons[i],
|
||||
"prompt_tokens": recv_obj.prompt_tokens[i],
|
||||
"weight_version": self.server_args.weight_version,
|
||||
"weight_version": get_serving().weight_version,
|
||||
"num_retractions": recv_obj.retraction_counts[i],
|
||||
}
|
||||
|
||||
@@ -2801,7 +2809,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
meta_info = {
|
||||
"id": recv_obj.rid,
|
||||
"finish_reason": finish_reason,
|
||||
"weight_version": self.server_args.weight_version,
|
||||
"weight_version": get_serving().weight_version,
|
||||
"e2e_latency": state.time_stats.get_e2e_latency(),
|
||||
}
|
||||
is_stream = getattr(state.obj, "stream", False)
|
||||
|
||||
@@ -597,7 +597,10 @@ class TokenizerManagerScoreMixin:
|
||||
f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})"
|
||||
)
|
||||
|
||||
# Check if multi-item scoring is enabled
|
||||
# Check if multi-item scoring is enabled. enable_mis is a static startup
|
||||
# feature flag (never overridden post-publish), and score_request is also
|
||||
# exercised on a bare mixin without a published context, so read it off
|
||||
# server_args rather than the resolved-config bag.
|
||||
use_multi_item_scoring = self.server_args.enable_mis
|
||||
|
||||
input_ids = None
|
||||
|
||||
@@ -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 (
|
||||
@@ -405,14 +406,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,
|
||||
@@ -423,7 +424,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,
|
||||
@@ -439,11 +440,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,
|
||||
@@ -459,7 +460,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
|
||||
@@ -485,9 +486,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,
|
||||
|
||||
Reference in New Issue
Block a user