diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 7308add1d..7667672a7 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -22,7 +22,6 @@ import sys import time from collections import deque from contextlib import contextmanager, nullcontext -from dataclasses import dataclass from http import HTTPStatus from typing import Any, Deque, Dict, List, Optional, Tuple, Union @@ -40,7 +39,6 @@ from torch.distributed import barrier from sglang.jit_kernel.ngram_embedding import update_token_table from sglang.srt.configs.model_config import ModelConfig, ModelImpl -from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX from sglang.srt.constrained.grammar_manager import GrammarManager from sglang.srt.disaggregation.decode import ( DecodePreallocQueue, @@ -90,8 +88,6 @@ from sglang.srt.managers.io_struct import ( AddExternalCorpusReqOutput, AttachHiCacheStorageReqInput, AttachHiCacheStorageReqOutput, - BaseBatchReq, - BaseReq, BatchTokenizedEmbeddingReqInput, BatchTokenizedGenerateReqInput, CheckWeightsReqInput, @@ -170,8 +166,10 @@ from sglang.srt.managers.scheduler_components.batch_result_processor import ( from sglang.srt.managers.scheduler_components.dp_attn import ( SchedulerDPAttnAdapter, ) +from sglang.srt.managers.scheduler_components.idle_sleeper import IdleSleeper from sglang.srt.managers.scheduler_components.invariant_checker import ( SchedulerInvariantChecker, + create_scheduler_watchdog, ) from sglang.srt.managers.scheduler_components.kv_events_publisher import ( SchedulerKvEventsPublisher, @@ -187,6 +185,7 @@ from sglang.srt.managers.scheduler_components.metrics_reporter import ( PrefillStats, SchedulerMetricsReporter, ) +from sglang.srt.managers.scheduler_components.output_sender import SenderWrapper from sglang.srt.managers.scheduler_components.output_streamer import ( SchedulerOutputStreamer, ) @@ -205,7 +204,12 @@ from sglang.srt.managers.scheduler_components.weight_updater import ( from sglang.srt.managers.scheduler_input_blocker import SchedulerInputBlocker from sglang.srt.managers.scheduler_pp_mixin import SchedulerPPMixin from sglang.srt.managers.scheduler_recv_skipper import SchedulerRecvSkipper -from sglang.srt.managers.utils import GenerationBatchResult, validate_input_length +from sglang.srt.managers.utils import ( + EmbeddingBatchResult, + GenerationBatchResult, + is_health_check_generate_req, + validate_input_length, +) from sglang.srt.mem_cache import kv_cache_builder from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache from sglang.srt.model_executor.forward_batch_info import ForwardMode, PPProxyTensors @@ -213,7 +217,6 @@ from sglang.srt.model_loader.utils import get_resolved_model_impl from sglang.srt.multiplex.multiplexing_mixin import SchedulerMultiplexMixin from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector from sglang.srt.observability.req_time_stats import ( - real_time, set_schedule_time_batch, set_time_batch, ) @@ -224,6 +227,7 @@ from sglang.srt.plugins import load_plugins from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.server_args import PortArgs, ServerArgs, get_global_server_args from sglang.srt.session.session_controller import SessionController +from sglang.srt.speculative.dflash_utils import validate_dflash_request from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( DynamicGradMode, @@ -250,7 +254,6 @@ from sglang.srt.utils.network import get_zmq_socket from sglang.srt.utils.numa_utils import get_numa_node_if_available, numa_bind_to_node from sglang.srt.utils.tensor_bridge import use_mlx from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter -from sglang.srt.utils.watchdog import WatchdogRaw from sglang.utils import TypeBasedDispatcher, get_exception_traceback if is_mps(): @@ -273,92 +276,6 @@ TEST_RETRACT_NO_PREFILL_BS = envs.SGLANG_TEST_RETRACT_NO_PREFILL_BS.get() _is_npu = is_npu() -@dataclass -class EmbeddingBatchResult: - """Result from an embedding/classification forward pass. - - Attributes: - embeddings: Model output — pooled embeddings or classification logits. - pooled_hidden_states: Raw hidden states before the task head. Present - only when the batch contained ``return_pooled_hidden_states=True`` - requests. Tensor (uniform shapes) or list of tensors (MIS). - copy_done: CUDA event recorded after the async CPU copy completes. - """ - - embeddings: torch.Tensor - pooled_hidden_states: Optional[torch.Tensor] = None - copy_done: Optional[torch.cuda.Event] = None - - def copy_to_cpu(self): - """Copy embeddings and pooled hidden states to CPU for overlap scheduling.""" - if isinstance(self.embeddings, torch.Tensor): - self.copy_done = torch.get_device_module(self.embeddings.device).Event() - self.embeddings = self.embeddings.to("cpu", non_blocking=True) - else: - assert isinstance(self.embeddings, list) - if len(self.embeddings) == 0: - return - - self.copy_done = torch.get_device_module(self.embeddings[0].device).Event() - self.embeddings = [ - emb.to("cpu", non_blocking=True) for emb in self.embeddings - ] - - if self.pooled_hidden_states is not None: - if isinstance(self.pooled_hidden_states, list): - self.pooled_hidden_states = [ - t.to("cpu", non_blocking=True) for t in self.pooled_hidden_states - ] - else: - self.pooled_hidden_states = self.pooled_hidden_states.to( - "cpu", non_blocking=True - ) - - self.copy_done.record() - - -def validate_dflash_request(req: Req) -> Optional[str]: - if req.return_logprob: - return "DFLASH speculative decoding does not support return_logprob yet." - - if ( - req.sampling_params.json_schema is not None - or req.sampling_params.regex is not None - or req.sampling_params.ebnf is not None - or req.sampling_params.structural_tag is not None - ): - return ( - "DFLASH speculative decoding does not support " - "grammar-constrained decoding yet." - ) - - return None - - -def create_scheduler_watchdog( - scheduler: "Scheduler", watchdog_timeout: float, soft: bool = False -) -> WatchdogRaw: - def dump_info() -> str: - if scheduler.is_initializing: - return "" - _, messages = scheduler.invariant_checker._check_all_pools( - scheduler.pool_stats_observer.get_pool_stats(), - ) - return ( - f"{scheduler.cur_batch.batch_size()=}\n" - f"{scheduler.cur_batch.reqs=}\n" + "\n".join(messages) - ) - - return WatchdogRaw( - debug_name="Scheduler", - get_counter=lambda: scheduler.forward_ct, - is_active=lambda: scheduler.is_initializing or scheduler.cur_batch is not None, - watchdog_timeout=watchdog_timeout, - soft=soft, - dump_info=dump_info, - ) - - class Scheduler( SchedulerDisaggregationDecodeMixin, SchedulerDisaggregationPrefillMixin, @@ -3790,41 +3707,6 @@ class Scheduler( pass -class IdleSleeper: - """ - In setups which have long inactivity periods it is desirable to reduce - system power consumption when sglang does nothing. This would lead not only - to power savings, but also to more CPU thermal headroom when a request - eventually comes. This is important in cases when multiple GPUs are connected - as each GPU would otherwise pin one thread at 100% CPU usage. - - The simplest solution is to use zmq.Poller on all sockets that may receive - data that needs handling immediately. - """ - - def __init__(self, sockets): - self.poller = zmq.Poller() - self.last_empty_time = real_time() - for s in sockets: - self.poller.register(s, zmq.POLLIN) - - self.empty_cache_interval = envs.SGLANG_EMPTY_CACHE_INTERVAL.get() - - def maybe_sleep(self): - self.poller.poll(1000) - if ( - self.empty_cache_interval > 0 - and real_time() - self.last_empty_time > self.empty_cache_interval - ): - self.last_empty_time = real_time() - current_platform.empty_cache() - - -def is_health_check_generate_req(recv_req): - rid = getattr(recv_req, "rid", None) - return rid is not None and rid.startswith(HEALTH_CHECK_RID_PREFIX) - - def is_work_request(recv_req): return isinstance( recv_req, @@ -3837,29 +3719,6 @@ def is_work_request(recv_req): ) -class SenderWrapper: - def __init__(self, socket: zmq.Socket): - self.socket = socket - - def send_output( - self, - output: Union[BaseReq, BaseBatchReq], - recv_obj: Optional[Union[BaseReq, BaseBatchReq]] = None, - ): - if self.socket is None: - return - - if ( - isinstance(recv_obj, BaseReq) - and recv_obj.http_worker_ipc is not None - and output.http_worker_ipc is None - ): - # handle communicator reqs for multi-http worker case - output.http_worker_ipc = recv_obj.http_worker_ipc - - self.socket.send_pyobj(output) - - def dispatch_event_loop(scheduler: Scheduler): # Dispatch to the appropriate event loop based on the disaggregation mode server_args = scheduler.server_args diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index 17974b7e7..5b50e1bbd 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -44,6 +44,10 @@ if TYPE_CHECKING: SchedulerOutputStreamer, ) from sglang.srt.managers.tp_worker import BaseTpWorker + 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 diff --git a/python/sglang/srt/managers/scheduler_components/idle_sleeper.py b/python/sglang/srt/managers/scheduler_components/idle_sleeper.py new file mode 100644 index 000000000..38c0c7a2e --- /dev/null +++ b/python/sglang/srt/managers/scheduler_components/idle_sleeper.py @@ -0,0 +1,35 @@ +import zmq + +from sglang.srt.environ import envs +from sglang.srt.observability.req_time_stats import real_time +from sglang.srt.platforms import current_platform + + +class IdleSleeper: + """ + In setups which have long inactivity periods it is desirable to reduce + system power consumption when sglang does nothing. This would lead not only + to power savings, but also to more CPU thermal headroom when a request + eventually comes. This is important in cases when multiple GPUs are connected + as each GPU would otherwise pin one thread at 100% CPU usage. + + The simplest solution is to use zmq.Poller on all sockets that may receive + data that needs handling immediately. + """ + + def __init__(self, sockets): + self.poller = zmq.Poller() + self.last_empty_time = real_time() + for s in sockets: + self.poller.register(s, zmq.POLLIN) + + self.empty_cache_interval = envs.SGLANG_EMPTY_CACHE_INTERVAL.get() + + def maybe_sleep(self): + self.poller.poll(1000) + if ( + self.empty_cache_interval > 0 + and real_time() - self.last_empty_time > self.empty_cache_interval + ): + self.last_empty_time = real_time() + current_platform.empty_cache() diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index e7ed22aa1..237a5a606 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -4,6 +4,7 @@ import logging import warnings from dataclasses import dataclass from typing import ( + TYPE_CHECKING, Callable, List, Optional, @@ -24,6 +25,11 @@ from sglang.srt.utils.common import ( ceil_align, raise_error_or_warn, ) +from sglang.srt.utils.watchdog import WatchdogRaw + +if TYPE_CHECKING: + from sglang.srt.managers.scheduler import Scheduler + logger = logging.getLogger(__name__) @@ -267,3 +273,27 @@ class SchedulerInvariantChecker: or (self.is_hybrid_ssm and self.tree_cache.supports_mamba()) ): self.tree_cache.sanity_check() + + +def create_scheduler_watchdog( + scheduler: "Scheduler", watchdog_timeout: float, soft: bool = False +) -> WatchdogRaw: + def dump_info() -> str: + if scheduler.is_initializing: + return "" + _, messages = scheduler.invariant_checker._check_all_pools( + scheduler.pool_stats_observer.get_pool_stats(), + ) + return ( + f"{scheduler.cur_batch.batch_size()=}\n" + f"{scheduler.cur_batch.reqs=}\n" + "\n".join(messages) + ) + + return WatchdogRaw( + debug_name="Scheduler", + get_counter=lambda: scheduler.forward_ct, + is_active=lambda: scheduler.is_initializing or scheduler.cur_batch is not None, + watchdog_timeout=watchdog_timeout, + soft=soft, + dump_info=dump_info, + ) diff --git a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py index c0a1b0d35..17101633a 100644 --- a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py +++ b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py @@ -32,10 +32,8 @@ from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_policy import PrefillAdder - from sglang.srt.managers.scheduler import ( - EmbeddingBatchResult, - Scheduler, - ) + from sglang.srt.managers.scheduler import Scheduler + from sglang.srt.managers.utils import EmbeddingBatchResult logger = logging.getLogger(__name__) diff --git a/python/sglang/srt/managers/scheduler_components/output_sender.py b/python/sglang/srt/managers/scheduler_components/output_sender.py new file mode 100644 index 000000000..a8eec2fd1 --- /dev/null +++ b/python/sglang/srt/managers/scheduler_components/output_sender.py @@ -0,0 +1,28 @@ +from typing import Optional, Union + +import zmq + +from sglang.srt.managers.io_struct import BaseBatchReq, BaseReq + + +class SenderWrapper: + def __init__(self, socket: zmq.Socket): + self.socket = socket + + def send_output( + self, + output: Union[BaseReq, BaseBatchReq], + recv_obj: Optional[Union[BaseReq, BaseBatchReq]] = None, + ): + if self.socket is None: + return + + if ( + isinstance(recv_obj, BaseReq) + and recv_obj.http_worker_ipc is not None + and output.http_worker_ipc is None + ): + # handle communicator reqs for multi-http worker case + output.http_worker_ipc = recv_obj.http_worker_ipc + + self.socket.send_pyobj(output) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index c097aa4b8..83d8bf6b8 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -75,12 +75,12 @@ from sglang.srt.managers.io_struct import ( from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors from sglang.srt.managers.schedule_batch import MultimodalDataItem -from sglang.srt.managers.scheduler import is_health_check_generate_req from sglang.srt.managers.scheduler_input_blocker import input_blocker_guard_region from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin from sglang.srt.managers.tokenizer_manager_score_mixin import ( TokenizerManagerScoreMixin, ) +from sglang.srt.managers.utils import is_health_check_generate_req from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread from sglang.srt.observability.metrics_collector import TokenizerMetricsCollector from sglang.srt.observability.req_time_stats import ( diff --git a/python/sglang/srt/managers/utils.py b/python/sglang/srt/managers/utils.py index 4b7879b5e..31c5b375c 100644 --- a/python/sglang/srt/managers/utils.py +++ b/python/sglang/srt/managers/utils.py @@ -2,10 +2,12 @@ from __future__ import annotations import dataclasses import logging +from dataclasses import dataclass from typing import TYPE_CHECKING, Any, List, Optional, Union import torch +from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX from sglang.srt.eplb.expert_distribution import ExpertDistributionMetrics from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.managers.overlap_utils import FutureIndices @@ -240,3 +242,52 @@ def get_alloc_len_per_decode(server_args: Optional[ServerArgs] = None) -> int: raise NotImplementedError( "get_alloc_len_per_decode not implemented for page_size > 1 and spec_topk > 1" ) + + +@dataclass +class EmbeddingBatchResult: + """Result from an embedding/classification forward pass. + + Attributes: + embeddings: Model output — pooled embeddings or classification logits. + pooled_hidden_states: Raw hidden states before the task head. Present + only when the batch contained ``return_pooled_hidden_states=True`` + requests. Tensor (uniform shapes) or list of tensors (MIS). + copy_done: CUDA event recorded after the async CPU copy completes. + """ + + embeddings: torch.Tensor + pooled_hidden_states: Optional[torch.Tensor] = None + copy_done: Optional[torch.cuda.Event] = None + + def copy_to_cpu(self): + """Copy embeddings and pooled hidden states to CPU for overlap scheduling.""" + if isinstance(self.embeddings, torch.Tensor): + self.copy_done = torch.get_device_module(self.embeddings.device).Event() + self.embeddings = self.embeddings.to("cpu", non_blocking=True) + else: + assert isinstance(self.embeddings, list) + if len(self.embeddings) == 0: + return + + self.copy_done = torch.get_device_module(self.embeddings[0].device).Event() + self.embeddings = [ + emb.to("cpu", non_blocking=True) for emb in self.embeddings + ] + + if self.pooled_hidden_states is not None: + if isinstance(self.pooled_hidden_states, list): + self.pooled_hidden_states = [ + t.to("cpu", non_blocking=True) for t in self.pooled_hidden_states + ] + else: + self.pooled_hidden_states = self.pooled_hidden_states.to( + "cpu", non_blocking=True + ) + + self.copy_done.record() + + +def is_health_check_generate_req(recv_req): + rid = getattr(recv_req, "rid", None) + return rid is not None and rid.startswith(HEALTH_CHECK_RID_PREFIX) diff --git a/python/sglang/srt/speculative/dflash_utils.py b/python/sglang/srt/speculative/dflash_utils.py index f1ea1d794..982772690 100644 --- a/python/sglang/srt/speculative/dflash_utils.py +++ b/python/sglang/srt/speculative/dflash_utils.py @@ -8,6 +8,7 @@ import torch import torch.nn.functional as F from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod +from sglang.srt.managers.schedule_batch import Req from sglang.srt.utils import is_cuda, is_musa DEFAULT_DFLASH_MASK_TOKEN = "<|MASK|>" @@ -636,3 +637,21 @@ def compute_dflash_sampling_correct_drafts_and_bonus( accept_pos = accept_index[row_ids, correct_len.to(torch.long)].to(torch.long) bonus = predicts[accept_pos].to(torch.int64) return correct_len, bonus + + +def validate_dflash_request(req: Req) -> Optional[str]: + if req.return_logprob: + return "DFLASH speculative decoding does not support return_logprob yet." + + if ( + req.sampling_params.json_schema is not None + or req.sampling_params.regex is not None + or req.sampling_params.ebnf is not None + or req.sampling_params.structural_tag is not None + ): + return ( + "DFLASH speculative decoding does not support " + "grammar-constrained decoding yet." + ) + + return None