Move module-level helpers out of scheduler.py (#25638)

This commit is contained in:
fzyzcjy
2026-05-18 18:45:38 +08:00
committed by GitHub
parent 99ad2b0894
commit c54b34c007
9 changed files with 180 additions and 156 deletions
+10 -151
View File
@@ -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
@@ -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
@@ -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()
@@ -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,
)
@@ -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__)
@@ -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)
@@ -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 (
+51
View File
@@ -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)
@@ -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