Move module-level helpers out of scheduler.py (#25638)
This commit is contained in:
@@ -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 (
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user