Revert RuntimeContext config-namespace reads/roles (#31813–#31817) (#32100)
This commit is contained in:
@@ -48,7 +48,6 @@ 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,
|
||||
@@ -232,7 +231,7 @@ class DataParallelController:
|
||||
sock_send(worker, obj)
|
||||
|
||||
def update_active_ranks(self, ranks: ActiveRanksOutput):
|
||||
if get_exec().moe.elastic_ep_backend is not None:
|
||||
if self.server_args.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 +484,7 @@ class DataParallelController:
|
||||
logger.debug("Worker port broadcast completed")
|
||||
return worker_ports
|
||||
finally:
|
||||
if get_exec().moe.elastic_ep_backend is None:
|
||||
if self.server_args.elastic_ep_backend is None:
|
||||
rep_socket.close()
|
||||
else:
|
||||
threading.Thread(
|
||||
@@ -816,12 +815,6 @@ def run_data_parallel_controller_process(
|
||||
kill_itself_when_parent_died()
|
||||
parent_process = psutil.Process().parent()
|
||||
|
||||
# Publish the resolved config at DP-controller process entry: this process
|
||||
# reads config namespaces (e.g. get_exec().moe.*) in its own address space
|
||||
# before spawning schedulers.
|
||||
from sglang.srt.runtime_context import publish
|
||||
|
||||
publish(server_args, role="scheduler")
|
||||
configure_logger(server_args)
|
||||
if server_args.enable_trace:
|
||||
process_tracing_init(
|
||||
|
||||
@@ -33,13 +33,7 @@ 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_disagg,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_server_args,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, 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
|
||||
@@ -884,7 +878,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_schedule().chunked_prefill_size
|
||||
chunked_prefill_size = get_server_args().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"
|
||||
@@ -1293,7 +1287,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_disagg().language_only:
|
||||
if get_server_args().language_only:
|
||||
precomputed_embeddings = getattr(
|
||||
mm_item, "precomputed_embeddings", None
|
||||
)
|
||||
@@ -1973,7 +1967,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_serving().skip_tokenizer_init:
|
||||
if _get_is_default_transport() or get_server_args().skip_tokenizer_init:
|
||||
return obj
|
||||
|
||||
if obj.mm_inputs:
|
||||
@@ -2034,7 +2028,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_serving().skip_tokenizer_init:
|
||||
if _get_is_default_transport() or get_server_args().skip_tokenizer_init:
|
||||
return obj
|
||||
# Handle batch requests
|
||||
if isinstance(obj, BaseBatchReq):
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
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.
|
||||
@@ -647,15 +645,15 @@ class TokenizerWorker(TokenizerManager):
|
||||
self.tokenizer_ipc_name = port_args.tokenizer_ipc_name
|
||||
|
||||
# For PD disaggregtion
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
self.server_args.override(
|
||||
"tokenizer_worker.restore_disaggregation_mode",
|
||||
disaggregation_mode=disaggregation_mode,
|
||||
)
|
||||
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||
self.disaggregation_mode = DisaggregationMode(
|
||||
self.server_args.disaggregation_mode
|
||||
)
|
||||
self.disaggregation_transfer_backend = TransferBackend(
|
||||
get_disagg().disaggregation_transfer_backend
|
||||
self.server_args.disaggregation_transfer_backend
|
||||
)
|
||||
|
||||
# Register this worker with the router for pause/continue broadcasting
|
||||
|
||||
@@ -77,7 +77,10 @@ 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 (
|
||||
@@ -102,12 +105,7 @@ from sglang.srt.observability.req_time_stats import (
|
||||
DPControllerReqTimeStats,
|
||||
SchedulerReqTimeStats,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_serving,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -1096,7 +1094,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_spec().speculative_algorithm
|
||||
spec_alg = get_server_args().speculative_algorithm
|
||||
return self.sampling_params.max_new_tokens == 0 and spec_alg is None
|
||||
|
||||
@property
|
||||
@@ -1117,7 +1115,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_serving().strip_thinking_cache and self.reasoning_tokens > 0:
|
||||
if get_server_args().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_disagg
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
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_disagg().disaggregation_mode != "decode"
|
||||
and get_server_args().disaggregation_mode != "decode"
|
||||
):
|
||||
for r in waiting_queue:
|
||||
match_prefix_for_req(self.tree_cache, r, include_req=True)
|
||||
|
||||
@@ -210,7 +210,9 @@ 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,
|
||||
)
|
||||
@@ -239,20 +241,7 @@ 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_device,
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_lora,
|
||||
get_memory,
|
||||
get_mm,
|
||||
get_observability,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_serving,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_context, get_parallel
|
||||
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
|
||||
@@ -454,9 +443,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=get_observability().enable_metrics,
|
||||
enable_metrics=self.server_args.enable_metrics,
|
||||
enable_kv_cache_events=bool(
|
||||
get_observability().kv_events_config
|
||||
self.server_args.kv_events_config
|
||||
and self.ps.pp_rank == 0
|
||||
and self.ps.attn_tp_rank == 0
|
||||
and self.ps.attn_cp_rank == 0
|
||||
@@ -482,8 +471,8 @@ class Scheduler(
|
||||
self.init_hisparse_coordinator()
|
||||
|
||||
if (
|
||||
get_disagg().disaggregation_mode == "decode"
|
||||
and get_disagg().disaggregation_decode_enable_offload_kvcache
|
||||
self.server_args.disaggregation_mode == "decode"
|
||||
and self.server_args.disaggregation_decode_enable_offload_kvcache
|
||||
):
|
||||
self.decode_offload_manager = DecodeKVCacheOffloadManager(
|
||||
req_to_token_pool=self.req_to_token_pool,
|
||||
@@ -594,7 +583,7 @@ class Scheduler(
|
||||
|
||||
self.dllm_config = ( # For diffusion LLM
|
||||
DllmConfig.from_server_args(self.server_args)
|
||||
if get_exec().dllm.dllm_algorithm is not None
|
||||
if self.server_args.dllm_algorithm is not None
|
||||
else None
|
||||
)
|
||||
|
||||
@@ -622,11 +611,11 @@ class Scheduler(
|
||||
self.ipc_channels = SchedulerIpcChannels.create(
|
||||
port_args=port_args,
|
||||
is_rank_zero=is_rank_zero,
|
||||
skip_tokenizer_init=get_serving().skip_tokenizer_init,
|
||||
metrics_enabled=get_observability().enable_metrics
|
||||
skip_tokenizer_init=self.server_args.skip_tokenizer_init,
|
||||
metrics_enabled=self.server_args.enable_metrics
|
||||
and (
|
||||
self.ps.attn_tp_rank == 0
|
||||
or get_observability().enable_metrics_for_all_schedulers
|
||||
or self.server_args.enable_metrics_for_all_schedulers
|
||||
),
|
||||
enable_scripted_runtime=envs.SGLANG_TEST_SCRIPTED_RUNTIME.get(),
|
||||
)
|
||||
@@ -642,7 +631,7 @@ class Scheduler(
|
||||
port_args,
|
||||
self.ps.dp_size,
|
||||
dp_rank,
|
||||
publish_interval=get_observability().load_snapshot_publish_interval,
|
||||
publish_interval=self.server_args.load_snapshot_publish_interval,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("load snapshot writer init failed: %s", e)
|
||||
@@ -652,7 +641,7 @@ class Scheduler(
|
||||
self.ps.pp_rank == 0
|
||||
and self.ps.attn_tp_rank == 0
|
||||
and self.ps.attn_cp_rank == 0
|
||||
and get_device().sleep_on_idle
|
||||
and self.server_args.sleep_on_idle
|
||||
):
|
||||
self.idle_sleeper = IdleSleeper(
|
||||
sockets=[
|
||||
@@ -723,9 +712,9 @@ class Scheduler(
|
||||
)
|
||||
|
||||
# Set reasoning_parser and think_end_id if --reasoning_parser is enabled
|
||||
if get_serving().reasoning_parser and self.tokenizer:
|
||||
if self.server_args.reasoning_parser and self.tokenizer:
|
||||
reasoning_parser = ReasoningParser(
|
||||
model_type=get_serving().reasoning_parser,
|
||||
model_type=self.server_args.reasoning_parser,
|
||||
stream_reasoning=False,
|
||||
tokenizer=self.tokenizer,
|
||||
)
|
||||
@@ -796,7 +785,7 @@ class Scheduler(
|
||||
target_worker=self.tp_worker,
|
||||
)
|
||||
|
||||
if get_spec().speculative_draft_load_format is not None:
|
||||
if self.server_args.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
|
||||
@@ -804,10 +793,10 @@ class Scheduler(
|
||||
# format.
|
||||
self.server_args.override(
|
||||
"scheduler.draft_load_format",
|
||||
load_format=get_spec().speculative_draft_load_format,
|
||||
load_format=self.server_args.speculative_draft_load_format,
|
||||
)
|
||||
logger.info(
|
||||
f"Using draft model load_format: '{get_spec().speculative_draft_load_format}'"
|
||||
f"Using draft model load_format: '{self.server_args.speculative_draft_load_format}'"
|
||||
)
|
||||
|
||||
DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args)
|
||||
@@ -898,7 +887,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(
|
||||
get_schedule().min_free_slots_delay,
|
||||
self.server_args.min_free_slots_delay,
|
||||
self.max_running_requests,
|
||||
is_dflash_family=self.spec_algorithm.is_dflash_family(),
|
||||
)
|
||||
@@ -944,14 +933,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={get_schedule().chunked_prefill_size}, "
|
||||
f"chunked_prefill_size={self.server_args.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 get_observability().enable_metrics:
|
||||
if self.server_args.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.
|
||||
@@ -998,7 +987,7 @@ class Scheduler(
|
||||
self._engine_paused = False
|
||||
|
||||
def init_chunked_prefill(self):
|
||||
self.chunked_prefill_size = get_schedule().chunked_prefill_size
|
||||
self.chunked_prefill_size = self.server_args.chunked_prefill_size
|
||||
uses_transformers_backend = (
|
||||
get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS
|
||||
)
|
||||
@@ -1018,12 +1007,13 @@ class Scheduler(
|
||||
self.chunked_req = None
|
||||
self._pending_chunked_abort_req = None
|
||||
self.is_mixed_chunk = (
|
||||
self.chunked_prefill_size is not None and get_schedule().enable_mixed_chunk
|
||||
self.chunked_prefill_size is not None
|
||||
and self.server_args.enable_mixed_chunk
|
||||
)
|
||||
|
||||
# Init the dynamic chunking predictor for PP
|
||||
self.enable_dynamic_chunking = (
|
||||
get_schedule().enable_dynamic_chunking and self.ps.pp_size > 1
|
||||
self.server_args.enable_dynamic_chunking and self.ps.pp_size > 1
|
||||
)
|
||||
if self.enable_dynamic_chunking:
|
||||
try:
|
||||
@@ -1059,8 +1049,8 @@ class Scheduler(
|
||||
)
|
||||
self.prefill_delayer: Optional[PrefillDelayer] = None
|
||||
self.max_prefill_bs: int = 0
|
||||
if get_schedule().enable_prefill_delayer:
|
||||
if get_disagg().disaggregation_mode == "decode":
|
||||
if self.server_args.enable_prefill_delayer:
|
||||
if self.server_args.disaggregation_mode == "decode":
|
||||
logger.info(
|
||||
"Ignoring --enable-prefill-delayer on decode engine "
|
||||
"(no prefill scheduling path; delayer would be a no-op)."
|
||||
@@ -1077,15 +1067,15 @@ class Scheduler(
|
||||
if self.metrics_reporter.enable_metrics
|
||||
else None
|
||||
),
|
||||
max_delay_passes=get_schedule().prefill_delayer_max_delay_passes,
|
||||
token_usage_low_watermark=get_schedule().prefill_delayer_token_usage_low_watermark,
|
||||
max_delay_passes=self.server_args.prefill_delayer_max_delay_passes,
|
||||
token_usage_low_watermark=self.server_args.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 get_schedule().disable_priority_preemption
|
||||
and not self.server_args.disable_priority_preemption
|
||||
)
|
||||
|
||||
self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args(
|
||||
@@ -1101,12 +1091,12 @@ class Scheduler(
|
||||
def init_watch_dog_memory_saver_input_blocker(self):
|
||||
# Start watchdog thread
|
||||
self.watchdog = create_scheduler_watchdog(
|
||||
self, watchdog_timeout=get_device().watchdog_timeout
|
||||
self, watchdog_timeout=self.server_args.watchdog_timeout
|
||||
)
|
||||
|
||||
# Init memory saver, profiler and metric stats
|
||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
enable=get_exec().features.enable_memory_saver
|
||||
enable=self.server_args.enable_memory_saver
|
||||
)
|
||||
|
||||
# Init recv skipper and input blocker
|
||||
@@ -1128,9 +1118,11 @@ class Scheduler(
|
||||
self.disagg_decode_prealloc_queue = None
|
||||
self.disagg_decode_transfer_queue = None
|
||||
|
||||
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||
self.disaggregation_mode = DisaggregationMode(
|
||||
self.server_args.disaggregation_mode
|
||||
)
|
||||
self.transfer_backend = TransferBackend(
|
||||
get_disagg().disaggregation_transfer_backend
|
||||
self.server_args.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?
|
||||
@@ -1198,12 +1190,12 @@ class Scheduler(
|
||||
gloo_group=self.attn_tp_cpu_group,
|
||||
tp_rank=self.ps.tp_rank,
|
||||
tp_size=self.ps.tp_size,
|
||||
dp_size=get_parallel().dp_size,
|
||||
dp_size=self.server_args.dp_size,
|
||||
gpu_id=self.ps.gpu_id,
|
||||
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
||||
bootstrap_port=self.server_args.disaggregation_bootstrap_port,
|
||||
max_total_num_tokens=self.max_total_num_tokens,
|
||||
pp_rank=self.ps.pp_rank,
|
||||
num_reserved_decode_tokens=get_disagg().num_reserved_decode_tokens,
|
||||
num_reserved_decode_tokens=self.server_args.num_reserved_decode_tokens,
|
||||
transfer_backend=self.transfer_backend,
|
||||
)
|
||||
|
||||
@@ -1229,7 +1221,7 @@ class Scheduler(
|
||||
tp_rank=self.ps.tp_rank,
|
||||
tp_size=self.ps.tp_size,
|
||||
gpu_id=self.ps.gpu_id,
|
||||
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
||||
bootstrap_port=self.server_args.disaggregation_bootstrap_port,
|
||||
gloo_group=self.attn_tp_cpu_group,
|
||||
max_total_num_tokens=self.max_total_num_tokens,
|
||||
scheduler=self,
|
||||
@@ -1243,10 +1235,11 @@ class Scheduler(
|
||||
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||
|
||||
# Init mm receiver for EPD disaggregation mode
|
||||
if get_disagg().language_only and get_disagg().encoder_transfer_backend in [
|
||||
"zmq_to_scheduler",
|
||||
"mooncake",
|
||||
]:
|
||||
if (
|
||||
self.server_args.language_only
|
||||
and self.server_args.encoder_transfer_backend
|
||||
in ["zmq_to_scheduler", "mooncake"]
|
||||
):
|
||||
self.mm_receiver = create_mm_receiver(
|
||||
self.server_args,
|
||||
dtype=self.model_config.dtype,
|
||||
@@ -1327,7 +1320,7 @@ class Scheduler(
|
||||
|
||||
def init_deterministic_inference_config(self):
|
||||
"""Initialize deterministic inference configuration for different attention backends."""
|
||||
if not get_exec().deterministic.enable_deterministic_inference:
|
||||
if not self.server_args.enable_deterministic_inference:
|
||||
self.truncation_align_size = None
|
||||
return
|
||||
|
||||
@@ -1336,7 +1329,7 @@ class Scheduler(
|
||||
"triton": ("SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE", 4096),
|
||||
}
|
||||
env_var, default_size = backend_sizes.get(
|
||||
get_exec().kernel.attention_backend, (None, None)
|
||||
self.server_args.attention_backend, (None, None)
|
||||
)
|
||||
self.truncation_align_size = (
|
||||
get_int_env_var(env_var, default_size) if env_var else None
|
||||
@@ -1732,10 +1725,10 @@ class Scheduler(
|
||||
)
|
||||
|
||||
def init_lora_drainer(self) -> None:
|
||||
if get_lora().lora_drain_wait_threshold > 0.0:
|
||||
if self.server_args.lora_drain_wait_threshold > 0.0:
|
||||
self.lora_drainer = LoRADrainer(
|
||||
get_lora().max_loras_per_batch,
|
||||
get_lora().lora_drain_wait_threshold,
|
||||
self.server_args.max_loras_per_batch,
|
||||
self.server_args.lora_drain_wait_threshold,
|
||||
)
|
||||
else:
|
||||
self.lora_drainer = None
|
||||
@@ -1837,7 +1830,7 @@ class Scheduler(
|
||||
|
||||
def init_kv_events_publisher(self) -> None:
|
||||
self.kv_events_publisher = SchedulerKvEventsPublisher(
|
||||
kv_events_config=get_observability().kv_events_config,
|
||||
kv_events_config=self.server_args.kv_events_config,
|
||||
ps=self.ps,
|
||||
attn_tp_rank=self.ps.attn_tp_rank,
|
||||
attn_cp_rank=self.ps.attn_cp_rank,
|
||||
@@ -2013,7 +2006,7 @@ class Scheduler(
|
||||
return image_inputs
|
||||
|
||||
def _get_multimodal_inputs(self, mm_inputs_dict):
|
||||
if get_mm().enable_broadcast_mm_inputs_process:
|
||||
if self.server_args.enable_broadcast_mm_inputs_process:
|
||||
return self._process_and_broadcast_mm_inputs(mm_inputs_dict)
|
||||
else:
|
||||
return MultimodalInputs.from_processor_output(mm_inputs_dict)
|
||||
@@ -2060,7 +2053,7 @@ class Scheduler(
|
||||
|
||||
def _maybe_namespace_elastic_radix_cache(self, req: Req) -> None:
|
||||
if (
|
||||
get_exec().moe.elastic_ep_backend is None
|
||||
self.server_args.elastic_ep_backend is None
|
||||
or self.disable_radix_cache
|
||||
or not self.tree_cache.is_tree_cache()
|
||||
):
|
||||
@@ -2106,7 +2099,8 @@ class Scheduler(
|
||||
)
|
||||
# Radix-native sessions use only the top-level session_id.
|
||||
radix_native_session = (
|
||||
recv_req.session_id is not None and get_memory().enable_session_radix_cache
|
||||
recv_req.session_id is not None
|
||||
and self.server_args.enable_session_radix_cache
|
||||
)
|
||||
|
||||
if session_id is None or radix_native_session:
|
||||
@@ -2118,7 +2112,7 @@ class Scheduler(
|
||||
|
||||
if recv_req.bootstrap_port is None:
|
||||
# Use default bootstrap port
|
||||
recv_req.bootstrap_port = get_disagg().disaggregation_bootstrap_port
|
||||
recv_req.bootstrap_port = self.server_args.disaggregation_bootstrap_port
|
||||
|
||||
req = Req(
|
||||
recv_req.rid,
|
||||
@@ -2271,7 +2265,7 @@ class Scheduler(
|
||||
self._add_request_to_queue(req)
|
||||
return
|
||||
|
||||
if req.return_sampling_mask and get_exec().kernel.sampling_backend == "ascend":
|
||||
if req.return_sampling_mask and self.server_args.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 = (
|
||||
@@ -2320,7 +2314,7 @@ class Scheduler(
|
||||
error_msg = validate_input_length(
|
||||
req,
|
||||
self.max_req_input_len,
|
||||
get_serving().allow_auto_truncate,
|
||||
self.server_args.allow_auto_truncate,
|
||||
)
|
||||
if error_msg:
|
||||
req.set_finish_with_abort(error_msg)
|
||||
@@ -2598,7 +2592,7 @@ class Scheduler(
|
||||
error_msg = validate_input_length(
|
||||
req,
|
||||
self.max_req_input_len,
|
||||
get_serving().allow_auto_truncate,
|
||||
self.server_args.allow_auto_truncate,
|
||||
)
|
||||
if error_msg:
|
||||
self._add_request_to_queue(req)
|
||||
@@ -2810,7 +2804,7 @@ class Scheduler(
|
||||
if (
|
||||
need_mlp_sync
|
||||
and not self.spec_algorithm.is_none()
|
||||
and not get_spec().speculative_skip_dp_mlp_sync
|
||||
and not self.server_args.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:
|
||||
@@ -2884,7 +2878,7 @@ class Scheduler(
|
||||
for req in ready_grammar_requests:
|
||||
self._add_request_to_queue(req)
|
||||
|
||||
if self.enable_hierarchical_cache or get_memory().enable_flexkv:
|
||||
if self.enable_hierarchical_cache or self.server_args.enable_flexkv:
|
||||
self.tree_cache.check_hicache_events()
|
||||
|
||||
if self.enable_priority_preemption or self.is_hybrid_swa:
|
||||
@@ -2951,7 +2945,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=get_schedule().prefill_max_requests,
|
||||
prefill_max_requests=self.server_args.prefill_max_requests,
|
||||
prefill_delayer_single_pass=prefill_delayer_single_pass,
|
||||
dllm_config=self.dllm_config,
|
||||
waiting_queue_len=len(self.waiting_queue),
|
||||
@@ -3522,7 +3516,7 @@ class Scheduler(
|
||||
|
||||
def _maybe_report_active_ranks(self) -> None:
|
||||
if not (
|
||||
self.enable_dp_attention and get_exec().moe.elastic_ep_backend is not None
|
||||
self.enable_dp_attention and self.server_args.elastic_ep_backend is not None
|
||||
):
|
||||
return
|
||||
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
||||
@@ -3798,7 +3792,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=get_serving().served_model_name,
|
||||
served_model_name=self.server_args.served_model_name,
|
||||
hicache_storage_prefetch_policy=recv_req.hicache_storage_prefetch_policy,
|
||||
hicache_write_policy=recv_req.hicache_write_policy,
|
||||
)
|
||||
@@ -3918,7 +3912,7 @@ class Scheduler(
|
||||
}
|
||||
ret["effective_max_running_requests_per_dp"] = self.max_running_requests
|
||||
|
||||
if get_exec().moe.elastic_ep_backend is not None:
|
||||
if self.server_args.elastic_ep_backend is not None:
|
||||
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
||||
|
||||
ret["is_scaling_elastic_ep"] = ElasticEPStateManager.is_scaling()
|
||||
@@ -4313,7 +4307,7 @@ class Scheduler(
|
||||
|
||||
old_ep_size = ElasticEPStateManager.get_effective_ep_size()
|
||||
new_ep_size = recv_req.new_ep_size
|
||||
max_ep_size = get_parallel().max_ep_size or old_ep_size
|
||||
max_ep_size = self.server_args.max_ep_size or old_ep_size
|
||||
|
||||
logger.debug(
|
||||
"[Elastic EP][scale] request received: new_ep_size=%d "
|
||||
@@ -4451,10 +4445,10 @@ class Scheduler(
|
||||
return None
|
||||
|
||||
def close_session(self, recv_req: CloseSessionReqInput):
|
||||
if get_memory().enable_session_radix_cache:
|
||||
if self.server_args.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 (
|
||||
get_memory().enable_session_radix_cache
|
||||
self.server_args.enable_session_radix_cache
|
||||
):
|
||||
self.session_controller.close(recv_req)
|
||||
|
||||
@@ -4633,13 +4627,6 @@ def run_scheduler_process(
|
||||
display_dp_rank=display_dp_rank,
|
||||
display_moe_ep_rank=display_moe_ep_rank,
|
||||
)
|
||||
# Publish the resolved config at scheduler process entry so the config
|
||||
# namespaces (get_serving()/get_device()/get_exec()/...) are available to
|
||||
# Scheduler.__init__ and its init_* helpers, which read them before the
|
||||
# model worker's own publish. ModelRunner re-publishes idempotently.
|
||||
from sglang.srt.runtime_context import publish
|
||||
|
||||
publish(server_args, role="scheduler")
|
||||
parent_process = psutil.Process().parent()
|
||||
|
||||
# Set up tracing
|
||||
|
||||
@@ -2,7 +2,14 @@ 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
|
||||
|
||||
@@ -16,14 +23,11 @@ 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.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_memory,
|
||||
get_observability,
|
||||
get_server_args,
|
||||
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.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
|
||||
@@ -44,7 +48,10 @@ 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
|
||||
@@ -77,7 +84,7 @@ class SchedulerBatchResultProcessor:
|
||||
|
||||
def process_batch_result_prebuilt(self, batch: ScheduleBatch):
|
||||
assert self.disaggregation_mode == DisaggregationMode.DECODE
|
||||
use_free_group = get_disagg().disaggregation_decode_enable_radix_cache
|
||||
use_free_group = self.server_args.disaggregation_decode_enable_radix_cache
|
||||
if use_free_group:
|
||||
self.token_to_kv_pool_allocator.free_group_begin()
|
||||
for req in batch.reqs:
|
||||
@@ -85,7 +92,7 @@ class SchedulerBatchResultProcessor:
|
||||
req.update_finish_state()
|
||||
if req.finished():
|
||||
req.time_stats.set_quick_finish_time()
|
||||
if get_memory().enable_hisparse:
|
||||
if self.server_args.enable_hisparse:
|
||||
self.hisparse_coordinator.request_finished(req)
|
||||
release_kv_cache(req, self.tree_cache)
|
||||
|
||||
@@ -236,7 +243,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 get_memory().enable_hisparse:
|
||||
if self.server_args.enable_hisparse:
|
||||
self.hisparse_coordinator.admit_request_into_staging(req)
|
||||
|
||||
self._maybe_collect_customized_info(i, req, logits_output)
|
||||
@@ -749,7 +756,7 @@ class SchedulerBatchResultProcessor:
|
||||
num_block_accept_tokens=result.num_block_accept_tokens,
|
||||
num_cap_tokens=result.num_cap_tokens,
|
||||
)
|
||||
if get_observability().enable_metrics:
|
||||
if self.server_args.enable_metrics:
|
||||
self.metrics_collector.increment_decode_cuda_graph_pass(
|
||||
value=can_run_cuda_graph
|
||||
)
|
||||
@@ -932,7 +939,7 @@ class SchedulerBatchResultProcessor:
|
||||
self._mamba_prefix_cache_update(req, batch, result, i)
|
||||
|
||||
if (
|
||||
get_disagg().disaggregation_decode_enable_offload_kvcache
|
||||
self.server_args.disaggregation_decode_enable_offload_kvcache
|
||||
and not req.finished()
|
||||
):
|
||||
self.decode_offload_manager.offload_kv_cache(req)
|
||||
@@ -952,12 +959,12 @@ class SchedulerBatchResultProcessor:
|
||||
self._maybe_collect_routed_experts(req)
|
||||
self._maybe_collect_indexer_topk(req)
|
||||
|
||||
if get_disagg().disaggregation_decode_enable_offload_kvcache:
|
||||
if self.server_args.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 get_memory().enable_hisparse:
|
||||
if self.server_args.enable_hisparse:
|
||||
self.hisparse_coordinator.request_finished(req)
|
||||
prepare_release = getattr(
|
||||
self.model_worker, "prepare_for_kv_cache_release", None
|
||||
@@ -1095,7 +1102,7 @@ class SchedulerBatchResultProcessor:
|
||||
For spec decode, the boundary is detected by comparing the
|
||||
accepted seq_len range against interval boundaries.
|
||||
"""
|
||||
interval = get_exec().mamba.mamba_track_interval
|
||||
interval = get_server_args().mamba_track_interval
|
||||
|
||||
if batch.spec_algorithm.is_none():
|
||||
if req.kv_committed_len % interval == 0:
|
||||
|
||||
@@ -12,7 +12,9 @@ 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
|
||||
@@ -24,7 +26,6 @@ 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_parallel, 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
|
||||
@@ -377,14 +378,14 @@ class SchedulerDPAttnAdapter:
|
||||
def prepare_mlp_sync_batch(self, local_batch: ScheduleBatch):
|
||||
return prepare_mlp_sync_batch_raw(
|
||||
local_batch,
|
||||
dp_size=get_parallel().dp_size,
|
||||
dp_size=self.server_args.dp_size,
|
||||
attn_tp_size=self.ps.attn_tp_size,
|
||||
attn_cp_size=self.ps.attn_cp_size,
|
||||
tp_group=self.tp_group,
|
||||
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=get_schedule().disable_overlap_schedule,
|
||||
disable_overlap_schedule=self.server_args.disable_overlap_schedule,
|
||||
offload_tags=self.offload_tags,
|
||||
dwdp=self.server_args.dwdp_size > 1,
|
||||
)
|
||||
|
||||
@@ -14,7 +14,6 @@ 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
|
||||
@@ -145,7 +144,7 @@ class SchedulerLoadInquirer:
|
||||
)
|
||||
|
||||
lora = None
|
||||
if get_lora().enable_lora:
|
||||
if self.server_args.enable_lora:
|
||||
lora = LoRAMetrics(
|
||||
slots_used=stats.lora_pool_slots_used,
|
||||
slots_total=stats.lora_pool_slots_total,
|
||||
|
||||
@@ -1,15 +1,20 @@
|
||||
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.runtime_context import get_exec
|
||||
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID, ServerArgs
|
||||
from sglang.srt.server_args import (
|
||||
MIS_DELIMITER_TOKEN_ID,
|
||||
ServerArgs,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(kw_only=True, slots=True, frozen=True)
|
||||
@@ -159,7 +164,7 @@ class SchedulerLogprobResultProcessor:
|
||||
delimiter token receive logprobs.
|
||||
"""
|
||||
return (
|
||||
get_exec().features.enable_mis
|
||||
self.server_args.enable_mis
|
||||
and req.is_prefill_only
|
||||
and req.multi_item_delimiter_indices is not None
|
||||
)
|
||||
|
||||
@@ -2,7 +2,12 @@ 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
|
||||
@@ -16,9 +21,11 @@ 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
|
||||
|
||||
@@ -137,7 +144,7 @@ class SchedulerOutputStreamer:
|
||||
return_sampling_mask=return_sampling_mask,
|
||||
spec_algorithm=self.spec_algorithm,
|
||||
disaggregation_mode=self.disaggregation_mode,
|
||||
default_stream_interval=get_serving().stream_interval,
|
||||
default_stream_interval=self.server_args.stream_interval,
|
||||
default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL,
|
||||
get_cached_tokens_details=self.get_cached_tokens_details,
|
||||
)
|
||||
@@ -164,7 +171,7 @@ class SchedulerOutputStreamer:
|
||||
if (
|
||||
req.finished()
|
||||
and self.ps.attn_tp_rank == 0
|
||||
and get_observability().enable_request_time_stats_logging
|
||||
and self.server_args.enable_request_time_stats_logging
|
||||
):
|
||||
req.log_time_stats()
|
||||
|
||||
|
||||
@@ -5,7 +5,13 @@ 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
|
||||
|
||||
@@ -13,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_device
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
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
|
||||
@@ -249,7 +255,7 @@ class SchedulerProfilerManager:
|
||||
self.profile_in_progress = True
|
||||
|
||||
if "CUDA_PROFILER" in activities:
|
||||
if self.ps.gpu_id == get_device().base_gpu_id:
|
||||
if self.ps.gpu_id == get_server_args().base_gpu_id:
|
||||
torch.cuda.cudart().cudaProfilerStart()
|
||||
self.profile_in_progress = True
|
||||
|
||||
@@ -359,7 +365,7 @@ class SchedulerProfilerManager:
|
||||
torch.cuda.memory._record_memory_history(enabled=None)
|
||||
|
||||
if "CUDA_PROFILER" in self.profiler_activities:
|
||||
if self.ps.gpu_id == get_device().base_gpu_id:
|
||||
if self.ps.gpu_id == get_server_args().base_gpu_id:
|
||||
torch.cuda.cudart().cudaProfilerStop()
|
||||
|
||||
merge_message = self._merge_profile_traces()
|
||||
|
||||
@@ -2,7 +2,14 @@ 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
|
||||
@@ -15,9 +22,14 @@ 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.runtime_context import get_disagg, get_parallel
|
||||
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.utils import (
|
||||
broadcast_pyobj,
|
||||
point_to_point_pyobj,
|
||||
)
|
||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -127,7 +139,7 @@ class SchedulerRequestReceiver:
|
||||
return recv_reqs
|
||||
|
||||
def _broadcast_reqs_across_ranks(self, recv_reqs: Optional[List]) -> List:
|
||||
if get_parallel().enable_dp_attention:
|
||||
if self.server_args.enable_dp_attention:
|
||||
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0:
|
||||
work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs)
|
||||
else:
|
||||
@@ -156,7 +168,7 @@ class SchedulerRequestReceiver:
|
||||
# instead of the full tp_group. This avoids an expensive
|
||||
# all-ranks gloo sync.
|
||||
_local_ctrl = (
|
||||
get_parallel().enable_dp_attention_local_control_broadcast
|
||||
self.server_args.enable_dp_attention_local_control_broadcast
|
||||
or self.server_args.is_ep_scale_joiner
|
||||
)
|
||||
if _local_ctrl:
|
||||
@@ -208,8 +220,8 @@ class SchedulerRequestReceiver:
|
||||
# Process MM requests under EPD-disaggregation mode
|
||||
if (
|
||||
self.ps.pp_rank == 0
|
||||
and get_disagg().language_only
|
||||
and get_disagg().encoder_transfer_backend
|
||||
and self.server_args.language_only
|
||||
and self.server_args.encoder_transfer_backend
|
||||
in ["zmq_to_scheduler", "mooncake"]
|
||||
):
|
||||
recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
|
||||
@@ -233,7 +245,7 @@ class SchedulerRequestReceiver:
|
||||
# peer ranks may still be unpickling ShmPointerMMData
|
||||
# (-> shm_open). Synchronize the same CPU groups that carried
|
||||
# SHM-backed work requests before materialize() unlinks them.
|
||||
if get_parallel().enable_dp_attention:
|
||||
if self.server_args.enable_dp_attention:
|
||||
if self.ps.attn_tp_size > 1:
|
||||
barrier(group=self.attn_tp_cpu_group)
|
||||
if self.ps.attn_cp_size > 1:
|
||||
|
||||
@@ -36,7 +36,6 @@ 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, get_parallel
|
||||
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
|
||||
@@ -123,7 +122,7 @@ class SchedulerPPMixin:
|
||||
next_pp_outputs = None
|
||||
next_batch_result = None
|
||||
d2h_event = None
|
||||
if get_parallel().pp_async_batch_depth > 0:
|
||||
if self.server_args.pp_async_batch_depth > 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -139,7 +138,7 @@ class SchedulerPPMixin:
|
||||
self.mb_metadata,
|
||||
self.last_rank_comm_queue,
|
||||
)
|
||||
if get_parallel().pp_async_batch_depth == 0:
|
||||
if self.server_args.pp_async_batch_depth == 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -269,7 +268,7 @@ class SchedulerPPMixin:
|
||||
server_is_idle = False
|
||||
pp_proxy_tensors = self._pp_recv_proxy_tensors()
|
||||
|
||||
if get_parallel().pp_async_batch_depth > 0:
|
||||
if self.server_args.pp_async_batch_depth > 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -285,7 +284,7 @@ class SchedulerPPMixin:
|
||||
self.mb_metadata,
|
||||
self.last_rank_comm_queue,
|
||||
)
|
||||
if get_parallel().pp_async_batch_depth == 0:
|
||||
if self.server_args.pp_async_batch_depth == 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -428,7 +427,7 @@ class SchedulerPPMixin:
|
||||
pp_proxy_tensors = self._pp_recv_proxy_tensors()
|
||||
|
||||
# early send output if possible
|
||||
if get_parallel().pp_async_batch_depth > 0:
|
||||
if self.server_args.pp_async_batch_depth > 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -446,7 +445,7 @@ class SchedulerPPMixin:
|
||||
self.last_rank_comm_queue,
|
||||
)
|
||||
|
||||
if get_parallel().pp_async_batch_depth == 0:
|
||||
if self.server_args.pp_async_batch_depth == 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -480,7 +479,7 @@ class SchedulerPPMixin:
|
||||
)
|
||||
)
|
||||
|
||||
if get_disagg().disaggregation_decode_enable_offload_kvcache:
|
||||
if self.server_args.disaggregation_decode_enable_offload_kvcache:
|
||||
self.decode_offload_manager.check_offload_progress()
|
||||
|
||||
if rmbs[next_mb_id] is not None:
|
||||
@@ -550,17 +549,17 @@ class SchedulerPPMixin:
|
||||
+ len(self.disagg_decode_transfer_queue.queue)
|
||||
+ len(self.disagg_decode_prealloc_queue.queue)
|
||||
)
|
||||
if get_disagg().disaggregation_decode_enable_offload_kvcache:
|
||||
if self.server_args.disaggregation_decode_enable_offload_kvcache:
|
||||
queue_size += len(self.decode_offload_manager.ongoing_offload)
|
||||
|
||||
if server_is_idle and queue_size == 0:
|
||||
self.on_idle()
|
||||
|
||||
def init_pp_loop_state(self: Scheduler):
|
||||
self.pp_loop_size: int = self.ps.pp_size + get_parallel().pp_async_batch_depth
|
||||
self.pp_loop_size: int = self.ps.pp_size + self.server_args.pp_async_batch_depth
|
||||
# In CP mode, attention weights are duplicated, eliminating the need for the attention TP all-gather operation.
|
||||
self.require_attn_tp_allgather = (
|
||||
not get_parallel().enable_dsa_prefill_context_parallel
|
||||
not self.server_args.enable_dsa_prefill_context_parallel
|
||||
)
|
||||
self.mbs = [None] * self.pp_loop_size
|
||||
self.last_mbs = [None] * self.pp_loop_size
|
||||
|
||||
@@ -74,7 +74,6 @@ from sglang.srt.managers.io_struct import (
|
||||
UpdateWeightsFromTensorReqOutput,
|
||||
)
|
||||
from sglang.srt.managers.load_snapshot import LoadSnapshot
|
||||
from sglang.srt.runtime_context import get_lora, get_parallel
|
||||
from sglang.srt.server_args import LoRARef, ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
@@ -146,8 +145,8 @@ class TokenizerControlMixin:
|
||||
|
||||
def update_control_communicator_fan_out(self: TokenizerManager, worker_count: int):
|
||||
primary_group_control = (
|
||||
get_parallel().enable_dp_attention
|
||||
and not get_parallel().enable_dp_attention_local_control_broadcast
|
||||
self.server_args.enable_dp_attention
|
||||
and not self.server_args.enable_dp_attention_local_control_broadcast
|
||||
)
|
||||
if primary_group_control:
|
||||
control_fan_out = (
|
||||
@@ -397,7 +396,7 @@ class TokenizerControlMixin:
|
||||
) -> Tuple[bool, str]:
|
||||
self.auto_create_handle_loop()
|
||||
assert (
|
||||
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
|
||||
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||
), "dp_size must be 1 or dp attention must be enabled for update weights from distributed"
|
||||
|
||||
results = await self.init_weights_update_group_communicator(obj)
|
||||
@@ -410,7 +409,7 @@ class TokenizerControlMixin:
|
||||
) -> Tuple[bool, str]:
|
||||
self.auto_create_handle_loop()
|
||||
assert (
|
||||
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
|
||||
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||
), "dp_size must be 1 or dp attention must be enabled for destroy parameter update group"
|
||||
|
||||
results = await self.destroy_weights_update_group_communicator(obj)
|
||||
@@ -423,7 +422,7 @@ class TokenizerControlMixin:
|
||||
) -> Tuple[bool, str]:
|
||||
self.auto_create_handle_loop()
|
||||
assert (
|
||||
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
|
||||
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||
), "dp_size must be 1 or dp attention must be enabled for update weights from distributed"
|
||||
|
||||
if obj.abort_all_requests:
|
||||
@@ -454,7 +453,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
# TODO: support DP
|
||||
assert (
|
||||
get_parallel().dp_size == 1
|
||||
self.server_args.dp_size == 1
|
||||
), "dp_size must be 1 for init_weights_send_group_for_remote_instance"
|
||||
result = (
|
||||
await self.init_weights_send_group_for_remote_instance_communicator(obj)
|
||||
@@ -469,7 +468,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
# TODO: support DP
|
||||
assert (
|
||||
get_parallel().dp_size == 1
|
||||
self.server_args.dp_size == 1
|
||||
), "dp_size must be 1 for send_weights_to_remote_instance"
|
||||
result = (await self.send_weights_to_remote_instance_communicator(obj))[0]
|
||||
return result.success, result.message
|
||||
@@ -481,7 +480,7 @@ class TokenizerControlMixin:
|
||||
) -> Tuple[bool, str]:
|
||||
self.auto_create_handle_loop()
|
||||
assert (
|
||||
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
|
||||
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||
), "dp_size must be 1 or dp attention must be enabled for update weights from tensor"
|
||||
|
||||
if obj.abort_all_requests:
|
||||
@@ -517,7 +516,7 @@ class TokenizerControlMixin:
|
||||
try:
|
||||
# For now, we only support single data parallel instance
|
||||
assert (
|
||||
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
|
||||
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||
), "dp_size must be 1 or dp attention must be enabled for update weights from IPC"
|
||||
logger.info("Starting IPC weight update")
|
||||
|
||||
@@ -570,7 +569,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
|
||||
try:
|
||||
if not get_lora().enable_lora:
|
||||
if not self.server_args.enable_lora:
|
||||
raise ValueError(
|
||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||
)
|
||||
@@ -578,7 +577,7 @@ class TokenizerControlMixin:
|
||||
# TODO (lifuhuang): Remove this after we verify that dynamic lora loading works
|
||||
# with dp_size > 1.
|
||||
assert (
|
||||
get_parallel().dp_size == 1
|
||||
self.server_args.dp_size == 1
|
||||
), "dp_size must be 1 for dynamic lora loading"
|
||||
logger.info(
|
||||
"Start load Lora adapter. Lora name=%s, path=%s",
|
||||
@@ -603,10 +602,10 @@ class TokenizerControlMixin:
|
||||
await self.lora_registry.register(new_adapter)
|
||||
self.lora_ref_cache[obj.lora_name] = new_adapter
|
||||
|
||||
if get_lora().max_loaded_loras is not None:
|
||||
if self.server_args.max_loaded_loras is not None:
|
||||
while (
|
||||
self.lora_registry.num_registered_loras
|
||||
> get_lora().max_loaded_loras
|
||||
> self.server_args.max_loaded_loras
|
||||
):
|
||||
lru_lora_name = await self.lora_registry.lru_lora_name(
|
||||
exclude_pinned=True
|
||||
@@ -620,7 +619,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: {get_lora().max_loaded_loras})"
|
||||
f"max allowed: {self.server_args.max_loaded_loras})"
|
||||
)
|
||||
|
||||
unload_result = await self._unload_lora_adapter_locked(
|
||||
@@ -648,13 +647,13 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
|
||||
try:
|
||||
if not get_lora().enable_lora:
|
||||
if not self.server_args.enable_lora:
|
||||
raise ValueError(
|
||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||
)
|
||||
|
||||
assert (
|
||||
get_parallel().dp_size == 1
|
||||
self.server_args.dp_size == 1
|
||||
), "dp_size must be 1 for dynamic lora loading"
|
||||
logger.info(
|
||||
"Start load Lora adapter from tensors. Lora name=%s",
|
||||
@@ -673,10 +672,10 @@ class TokenizerControlMixin:
|
||||
if result.success:
|
||||
await self.lora_registry.register(new_adapter)
|
||||
self.lora_ref_cache[obj.lora_name] = new_adapter
|
||||
if get_lora().max_loaded_loras is not None:
|
||||
if self.server_args.max_loaded_loras is not None:
|
||||
while (
|
||||
self.lora_registry.num_registered_loras
|
||||
> get_lora().max_loaded_loras
|
||||
> self.server_args.max_loaded_loras
|
||||
):
|
||||
lru_lora_name = await self.lora_registry.lru_lora_name(
|
||||
exclude_pinned=True
|
||||
@@ -690,7 +689,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: {get_lora().max_loaded_loras})"
|
||||
f"max allowed: {self.server_args.max_loaded_loras})"
|
||||
)
|
||||
|
||||
unload_result = await self._unload_lora_adapter_locked(
|
||||
@@ -718,7 +717,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
|
||||
try:
|
||||
if not get_lora().enable_lora:
|
||||
if not self.server_args.enable_lora:
|
||||
raise ValueError(
|
||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||
)
|
||||
@@ -730,7 +729,7 @@ class TokenizerControlMixin:
|
||||
# TODO (lifuhuang): Remove this after we verify that dynamic lora loading works
|
||||
# with dp_size > 1.
|
||||
assert (
|
||||
get_parallel().dp_size == 1
|
||||
self.server_args.dp_size == 1
|
||||
), "dp_size must be 1 for dynamic lora loading"
|
||||
logger.info(
|
||||
"Start unload Lora adapter. Lora name=%s",
|
||||
@@ -750,7 +749,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
results = await self.get_weights_by_name_communicator(obj)
|
||||
all_parameters = [r.parameter for r in results]
|
||||
if get_parallel().dp_size == 1:
|
||||
if self.server_args.dp_size == 1:
|
||||
return all_parameters[0]
|
||||
else:
|
||||
return all_parameters
|
||||
@@ -894,8 +893,6 @@ class TokenizerControlMixin:
|
||||
) -> None:
|
||||
"""Update weight version if provided."""
|
||||
if weight_version is not None:
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
self.server_args.override(
|
||||
"tokenizer.weight_version", weight_version=weight_version
|
||||
)
|
||||
|
||||
@@ -110,15 +110,6 @@ 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_parallel,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.server_args import (
|
||||
PortArgs,
|
||||
@@ -472,10 +463,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# TODO: Refactor and organize the log export code.
|
||||
# Request logging
|
||||
self.request_logger = RequestLogger(
|
||||
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,
|
||||
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,
|
||||
)
|
||||
|
||||
# Dumping
|
||||
@@ -498,7 +489,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
def init_weight_update(self):
|
||||
# Initial weights status
|
||||
self.initial_weights_loaded = True
|
||||
if get_model().checkpoint_engine_wait_weights_before_ready:
|
||||
if self.server_args.checkpoint_engine_wait_weights_before_ready:
|
||||
self.initial_weights_loaded = False
|
||||
|
||||
# Weight updates
|
||||
@@ -518,7 +509,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(get_lora().lora_paths)
|
||||
self.lora_registry = LoRARegistry(self.server_args.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.
|
||||
@@ -527,13 +518,15 @@ 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 get_lora().lora_paths is not None:
|
||||
for lora_ref in get_lora().lora_paths:
|
||||
if self.server_args.lora_paths is not None:
|
||||
for lora_ref in self.server_args.lora_paths:
|
||||
self.lora_ref_cache[lora_ref.lora_name] = lora_ref
|
||||
|
||||
def init_disaggregation(self):
|
||||
# PD Disaggregation
|
||||
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||
self.disaggregation_mode = DisaggregationMode(
|
||||
self.server_args.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.
|
||||
@@ -542,16 +535,18 @@ 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(get_disagg().encoder_urls)
|
||||
self.encoder_urls: List[str] = list(self.server_args.encoder_urls)
|
||||
self.encoder_bootstrap_server = EncoderBootstrapServer(
|
||||
host=get_serving().host,
|
||||
port=get_disagg().encoder_bootstrap_port,
|
||||
host=self.server_args.host,
|
||||
port=self.server_args.encoder_bootstrap_port,
|
||||
urls=self.encoder_urls,
|
||||
)
|
||||
self.mm_receiver = create_mm_receiver(
|
||||
@@ -565,22 +560,20 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# Metrics
|
||||
if self.enable_metrics:
|
||||
engine_type = DisaggregationMode.to_engine_type(
|
||||
get_disagg().disaggregation_mode
|
||||
self.server_args.disaggregation_mode
|
||||
)
|
||||
|
||||
labels = {
|
||||
"model_name": get_serving().served_model_name,
|
||||
"model_name": self.server_args.served_model_name,
|
||||
"engine_type": engine_type,
|
||||
}
|
||||
if self.enable_priority_scheduling:
|
||||
labels["priority"] = ""
|
||||
if get_observability().tokenizer_metrics_allowed_custom_labels:
|
||||
for (
|
||||
label
|
||||
) in get_observability().tokenizer_metrics_allowed_custom_labels:
|
||||
if self.server_args.tokenizer_metrics_allowed_custom_labels:
|
||||
for label in self.server_args.tokenizer_metrics_allowed_custom_labels:
|
||||
labels[label] = ""
|
||||
if get_observability().extra_metric_labels:
|
||||
labels.update(get_observability().extra_metric_labels)
|
||||
if self.server_args.extra_metric_labels:
|
||||
labels.update(self.server_args.extra_metric_labels)
|
||||
tokenizer_collector_cls = resolve_collector_class(
|
||||
self.server_args,
|
||||
STAT_LOGGER_ROLE_TOKENIZER,
|
||||
@@ -589,18 +582,18 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
self.metrics_collector = tokenizer_collector_cls(
|
||||
server_args=self.server_args,
|
||||
labels=labels,
|
||||
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,
|
||||
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,
|
||||
)
|
||||
|
||||
start_cpu_monitor_thread("tokenizer")
|
||||
|
||||
if get_observability().gc_warning_threshold_secs > 0.0:
|
||||
configure_gc_warning(get_observability().gc_warning_threshold_secs)
|
||||
if self.server_args.gc_warning_threshold_secs > 0.0:
|
||||
configure_gc_warning(self.server_args.gc_warning_threshold_secs)
|
||||
self.soft_watchdog = Watchdog.create(
|
||||
debug_name="TokenizerManager",
|
||||
watchdog_timeout=get_device().soft_watchdog_timeout,
|
||||
watchdog_timeout=self.server_args.soft_watchdog_timeout,
|
||||
soft=True,
|
||||
test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(),
|
||||
)
|
||||
@@ -1366,7 +1359,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
return batch_size > 0 and (
|
||||
self.server_args.enable_tokenizer_batch_encode
|
||||
or (
|
||||
(not get_parallel().enable_dp_attention)
|
||||
(not self.server_args.enable_dp_attention)
|
||||
and (not self._batch_has_text(batch_size, requests))
|
||||
)
|
||||
)
|
||||
@@ -1764,7 +1757,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
|
||||
# default the load format to the server_args
|
||||
if obj.load_format is None:
|
||||
obj.load_format = get_model().load_format
|
||||
obj.load_format = self.server_args.load_format
|
||||
logger.info("Start update_weights. Load format=%s", obj.load_format)
|
||||
|
||||
if obj.abort_all_requests:
|
||||
@@ -1790,9 +1783,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
|
||||
def _update_model_path_info(self, model_path: str, load_format: str):
|
||||
self.served_model_name = model_path
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
self.server_args.override(
|
||||
"tokenizer.update_weights", model_path=model_path, load_format=load_format
|
||||
)
|
||||
self.model_path = model_path
|
||||
@@ -1936,7 +1927,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
"id": rid,
|
||||
"finish_reason": recv_obj.finished_reasons[i],
|
||||
"prompt_tokens": recv_obj.prompt_tokens[i],
|
||||
"weight_version": get_serving().weight_version,
|
||||
"weight_version": self.server_args.weight_version,
|
||||
"num_retractions": recv_obj.retraction_counts[i],
|
||||
}
|
||||
|
||||
@@ -2810,7 +2801,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
meta_info = {
|
||||
"id": recv_obj.rid,
|
||||
"finish_reason": finish_reason,
|
||||
"weight_version": get_serving().weight_version,
|
||||
"weight_version": self.server_args.weight_version,
|
||||
"e2e_latency": state.time_stats.get_e2e_latency(),
|
||||
}
|
||||
is_stream = getattr(state.obj, "stream", False)
|
||||
|
||||
@@ -597,10 +597,7 @@ class TokenizerManagerScoreMixin:
|
||||
f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})"
|
||||
)
|
||||
|
||||
# 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.
|
||||
# Check if multi-item scoring is enabled
|
||||
use_multi_item_scoring = self.server_args.enable_mis
|
||||
|
||||
input_ids = None
|
||||
|
||||
@@ -47,7 +47,6 @@ 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 (
|
||||
@@ -406,14 +405,14 @@ class TpModelWorker(BaseTpWorker):
|
||||
self.model_config = ModelConfig.from_server_args(
|
||||
self.server_args,
|
||||
model_path=(
|
||||
get_model().model_path
|
||||
self.server_args.model_path
|
||||
if not self.is_draft_worker
|
||||
else get_spec().speculative_draft_model_path
|
||||
else self.server_args.speculative_draft_model_path
|
||||
),
|
||||
model_revision=(
|
||||
get_model().revision
|
||||
self.server_args.revision
|
||||
if not self.is_draft_worker
|
||||
else get_spec().speculative_draft_model_revision
|
||||
else self.server_args.speculative_draft_model_revision
|
||||
),
|
||||
is_draft_model=self.is_draft_worker,
|
||||
context_length=self.context_length,
|
||||
@@ -424,7 +423,7 @@ class TpModelWorker(BaseTpWorker):
|
||||
|
||||
self._model_runner = ModelRunner(
|
||||
model_config=self.model_config,
|
||||
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||
mem_fraction_static=self.server_args.mem_fraction_static,
|
||||
gpu_id=self.gpu_id,
|
||||
ps=self.ps,
|
||||
nccl_port=self.nccl_port,
|
||||
@@ -440,11 +439,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, get_spec().speculative_num_steps):
|
||||
for i in range(1, self.server_args.speculative_num_steps):
|
||||
self.model_runner_list.append(
|
||||
ModelRunner(
|
||||
model_config=self.model_config,
|
||||
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||
mem_fraction_static=self.server_args.mem_fraction_static,
|
||||
gpu_id=self.gpu_id,
|
||||
ps=self.ps,
|
||||
nccl_port=self.nccl_port,
|
||||
@@ -460,7 +459,7 @@ class TpModelWorker(BaseTpWorker):
|
||||
def _init_dllm_algorithm(self):
|
||||
from sglang.srt.dllm.algorithm.base import DllmAlgorithm
|
||||
|
||||
if get_exec().dllm.dllm_algorithm is not None:
|
||||
if self.server_args.dllm_algorithm is not None:
|
||||
self.dllm_algorithm = DllmAlgorithm.from_server_args(self.server_args)
|
||||
else:
|
||||
self.dllm_algorithm = None
|
||||
@@ -486,9 +485,9 @@ class TpModelWorker(BaseTpWorker):
|
||||
)
|
||||
return (
|
||||
self.model_runner.max_total_num_tokens,
|
||||
get_schedule().max_prefill_tokens,
|
||||
self.server_args.max_prefill_tokens,
|
||||
self.model_runner.max_running_requests,
|
||||
get_schedule().max_queued_requests,
|
||||
self.server_args.max_queued_requests,
|
||||
max_req_len,
|
||||
max_req_len - 5,
|
||||
self.random_seed,
|
||||
|
||||
Reference in New Issue
Block a user