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