Move module-level helpers out of scheduler.py (#25638)
This commit is contained in:
@@ -22,7 +22,6 @@ import sys
|
|||||||
import time
|
import time
|
||||||
from collections import deque
|
from collections import deque
|
||||||
from contextlib import contextmanager, nullcontext
|
from contextlib import contextmanager, nullcontext
|
||||||
from dataclasses import dataclass
|
|
||||||
from http import HTTPStatus
|
from http import HTTPStatus
|
||||||
from typing import Any, Deque, Dict, List, Optional, Tuple, Union
|
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.jit_kernel.ngram_embedding import update_token_table
|
||||||
from sglang.srt.configs.model_config import ModelConfig, ModelImpl
|
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.constrained.grammar_manager import GrammarManager
|
||||||
from sglang.srt.disaggregation.decode import (
|
from sglang.srt.disaggregation.decode import (
|
||||||
DecodePreallocQueue,
|
DecodePreallocQueue,
|
||||||
@@ -90,8 +88,6 @@ from sglang.srt.managers.io_struct import (
|
|||||||
AddExternalCorpusReqOutput,
|
AddExternalCorpusReqOutput,
|
||||||
AttachHiCacheStorageReqInput,
|
AttachHiCacheStorageReqInput,
|
||||||
AttachHiCacheStorageReqOutput,
|
AttachHiCacheStorageReqOutput,
|
||||||
BaseBatchReq,
|
|
||||||
BaseReq,
|
|
||||||
BatchTokenizedEmbeddingReqInput,
|
BatchTokenizedEmbeddingReqInput,
|
||||||
BatchTokenizedGenerateReqInput,
|
BatchTokenizedGenerateReqInput,
|
||||||
CheckWeightsReqInput,
|
CheckWeightsReqInput,
|
||||||
@@ -170,8 +166,10 @@ from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
|||||||
from sglang.srt.managers.scheduler_components.dp_attn import (
|
from sglang.srt.managers.scheduler_components.dp_attn import (
|
||||||
SchedulerDPAttnAdapter,
|
SchedulerDPAttnAdapter,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.scheduler_components.idle_sleeper import IdleSleeper
|
||||||
from sglang.srt.managers.scheduler_components.invariant_checker import (
|
from sglang.srt.managers.scheduler_components.invariant_checker import (
|
||||||
SchedulerInvariantChecker,
|
SchedulerInvariantChecker,
|
||||||
|
create_scheduler_watchdog,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.scheduler_components.kv_events_publisher import (
|
from sglang.srt.managers.scheduler_components.kv_events_publisher import (
|
||||||
SchedulerKvEventsPublisher,
|
SchedulerKvEventsPublisher,
|
||||||
@@ -187,6 +185,7 @@ from sglang.srt.managers.scheduler_components.metrics_reporter import (
|
|||||||
PrefillStats,
|
PrefillStats,
|
||||||
SchedulerMetricsReporter,
|
SchedulerMetricsReporter,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.scheduler_components.output_sender import SenderWrapper
|
||||||
from sglang.srt.managers.scheduler_components.output_streamer import (
|
from sglang.srt.managers.scheduler_components.output_streamer import (
|
||||||
SchedulerOutputStreamer,
|
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_input_blocker import SchedulerInputBlocker
|
||||||
from sglang.srt.managers.scheduler_pp_mixin import SchedulerPPMixin
|
from sglang.srt.managers.scheduler_pp_mixin import SchedulerPPMixin
|
||||||
from sglang.srt.managers.scheduler_recv_skipper import SchedulerRecvSkipper
|
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 import kv_cache_builder
|
||||||
from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache
|
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
|
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.multiplex.multiplexing_mixin import SchedulerMultiplexMixin
|
||||||
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
|
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
|
||||||
from sglang.srt.observability.req_time_stats import (
|
from sglang.srt.observability.req_time_stats import (
|
||||||
real_time,
|
|
||||||
set_schedule_time_batch,
|
set_schedule_time_batch,
|
||||||
set_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.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs, get_global_server_args
|
from sglang.srt.server_args import PortArgs, ServerArgs, get_global_server_args
|
||||||
from sglang.srt.session.session_controller import SessionController
|
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.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
DynamicGradMode,
|
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.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.tensor_bridge import use_mlx
|
||||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
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
|
from sglang.utils import TypeBasedDispatcher, get_exception_traceback
|
||||||
|
|
||||||
if is_mps():
|
if is_mps():
|
||||||
@@ -273,92 +276,6 @@ TEST_RETRACT_NO_PREFILL_BS = envs.SGLANG_TEST_RETRACT_NO_PREFILL_BS.get()
|
|||||||
_is_npu = is_npu()
|
_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(
|
class Scheduler(
|
||||||
SchedulerDisaggregationDecodeMixin,
|
SchedulerDisaggregationDecodeMixin,
|
||||||
SchedulerDisaggregationPrefillMixin,
|
SchedulerDisaggregationPrefillMixin,
|
||||||
@@ -3790,41 +3707,6 @@ class Scheduler(
|
|||||||
pass
|
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):
|
def is_work_request(recv_req):
|
||||||
return isinstance(
|
return isinstance(
|
||||||
recv_req,
|
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):
|
def dispatch_event_loop(scheduler: Scheduler):
|
||||||
# Dispatch to the appropriate event loop based on the disaggregation mode
|
# Dispatch to the appropriate event loop based on the disaggregation mode
|
||||||
server_args = scheduler.server_args
|
server_args = scheduler.server_args
|
||||||
|
|||||||
@@ -44,6 +44,10 @@ if TYPE_CHECKING:
|
|||||||
SchedulerOutputStreamer,
|
SchedulerOutputStreamer,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.tp_worker import BaseTpWorker
|
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.allocator import BaseTokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
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
|
import warnings
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import (
|
from typing import (
|
||||||
|
TYPE_CHECKING,
|
||||||
Callable,
|
Callable,
|
||||||
List,
|
List,
|
||||||
Optional,
|
Optional,
|
||||||
@@ -24,6 +25,11 @@ from sglang.srt.utils.common import (
|
|||||||
ceil_align,
|
ceil_align,
|
||||||
raise_error_or_warn,
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -267,3 +273,27 @@ class SchedulerInvariantChecker:
|
|||||||
or (self.is_hybrid_ssm and self.tree_cache.supports_mamba())
|
or (self.is_hybrid_ssm and self.tree_cache.supports_mamba())
|
||||||
):
|
):
|
||||||
self.tree_cache.sanity_check()
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
from sglang.srt.managers.schedule_policy import PrefillAdder
|
from sglang.srt.managers.schedule_policy import PrefillAdder
|
||||||
from sglang.srt.managers.scheduler import (
|
from sglang.srt.managers.scheduler import Scheduler
|
||||||
EmbeddingBatchResult,
|
from sglang.srt.managers.utils import EmbeddingBatchResult
|
||||||
Scheduler,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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.mm_utils import TensorTransportMode, wrap_shm_features
|
||||||
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
|
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
|
||||||
from sglang.srt.managers.schedule_batch import MultimodalDataItem
|
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.scheduler_input_blocker import input_blocker_guard_region
|
||||||
from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin
|
from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin
|
||||||
from sglang.srt.managers.tokenizer_manager_score_mixin import (
|
from sglang.srt.managers.tokenizer_manager_score_mixin import (
|
||||||
TokenizerManagerScoreMixin,
|
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.cpu_monitor import start_cpu_monitor_thread
|
||||||
from sglang.srt.observability.metrics_collector import TokenizerMetricsCollector
|
from sglang.srt.observability.metrics_collector import TokenizerMetricsCollector
|
||||||
from sglang.srt.observability.req_time_stats import (
|
from sglang.srt.observability.req_time_stats import (
|
||||||
|
|||||||
@@ -2,10 +2,12 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import logging
|
import logging
|
||||||
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, Any, List, Optional, Union
|
from typing import TYPE_CHECKING, Any, List, Optional, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX
|
||||||
from sglang.srt.eplb.expert_distribution import ExpertDistributionMetrics
|
from sglang.srt.eplb.expert_distribution import ExpertDistributionMetrics
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
from sglang.srt.managers.overlap_utils import FutureIndices
|
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(
|
raise NotImplementedError(
|
||||||
"get_alloc_len_per_decode not implemented for page_size > 1 and spec_topk > 1"
|
"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
|
import torch.nn.functional as F
|
||||||
|
|
||||||
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
|
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
|
from sglang.srt.utils import is_cuda, is_musa
|
||||||
|
|
||||||
DEFAULT_DFLASH_MASK_TOKEN = "<|MASK|>"
|
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)
|
accept_pos = accept_index[row_ids, correct_len.to(torch.long)].to(torch.long)
|
||||||
bonus = predicts[accept_pos].to(torch.int64)
|
bonus = predicts[accept_pos].to(torch.int64)
|
||||||
return correct_len, bonus
|
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