diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 58670f279..905d7b3a1 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -33,7 +33,6 @@ import psutil import setproctitle import torch import torch.distributed -import zmq from torch.cuda import Stream as CudaStream from torch.distributed import barrier @@ -171,6 +170,9 @@ from sglang.srt.managers.scheduler_components.invariant_checker import ( SchedulerInvariantChecker, create_scheduler_watchdog, ) +from sglang.srt.managers.scheduler_components.ipc_channels import ( + SchedulerIpcChannels, +) from sglang.srt.managers.scheduler_components.kv_events_publisher import ( SchedulerKvEventsPublisher, ) @@ -185,7 +187,6 @@ 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, ) @@ -250,7 +251,6 @@ from sglang.srt.utils.hf_transformers_utils import ( get_tokenizer, get_tokenizer_from_processor, ) -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 @@ -552,8 +552,8 @@ class Scheduler( self.grammar_manager = GrammarManager(self) self.request_receiver = SchedulerRequestReceiver( - recv_from_tokenizer=self.recv_from_tokenizer, - recv_from_rpc=self.recv_from_rpc, + recv_from_tokenizer=self.ipc_channels.recv_from_tokenizer, + recv_from_rpc=self.ipc_channels.recv_from_rpc, recv_skipper=self.recv_skipper, input_blocker=self.input_blocker, mm_receiver=self.mm_receiver, @@ -629,7 +629,7 @@ class Scheduler( attn_dp_rank=self.ps.attn_dp_rank, dp_rank=self.ps.dp_rank, tree_cache=self.tree_cache, - send_metrics_from_scheduler=self.send_metrics_from_scheduler, + send_metrics_from_scheduler=self.ipc_channels.send_metrics_from_scheduler, max_running_requests=self.max_running_requests, max_total_num_tokens=self.max_total_num_tokens, get_stats=lambda: self.metrics_reporter.stats, @@ -658,7 +658,7 @@ class Scheduler( ) self.output_streamer = SchedulerOutputStreamer( - send_to_detokenizer=self.send_to_detokenizer, + send_to_detokenizer=self.ipc_channels.send_to_detokenizer, tree_cache=self.tree_cache, ps=self.ps, server_args=self.server_args, @@ -726,50 +726,21 @@ class Scheduler( self.page_size = self.dllm_config.block_size def init_ipc_channels(self, port_args: PortArgs): - context = zmq.Context(2) - self.send_metrics_from_scheduler = None - - if ( + is_rank_zero = ( self.ps.pp_rank == 0 and self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0 - ): - self.recv_from_tokenizer = get_zmq_socket( - context, zmq.PULL, port_args.scheduler_input_ipc_name, False - ) - self.recv_from_rpc = get_zmq_socket( - context, zmq.DEALER, port_args.rpc_ipc_name, False - ) - - send_to_tokenizer = get_zmq_socket( - context, zmq.PUSH, port_args.tokenizer_ipc_name, False - ) - if self.server_args.skip_tokenizer_init: - # Directly send to the TokenizerManager - send_to_detokenizer = get_zmq_socket( - context, zmq.PUSH, port_args.tokenizer_ipc_name, False - ) - else: - # Send to the DetokenizerManager - send_to_detokenizer = get_zmq_socket( - context, zmq.PUSH, port_args.detokenizer_ipc_name, False - ) - - self.send_to_tokenizer = SenderWrapper(send_to_tokenizer) - self.send_to_detokenizer = SenderWrapper(send_to_detokenizer) - else: - self.recv_from_tokenizer = None - self.recv_from_rpc = None - self.send_to_tokenizer = SenderWrapper(None) - self.send_to_detokenizer = SenderWrapper(None) - - if self.server_args.enable_metrics and ( - self.ps.attn_tp_rank == 0 - or self.server_args.enable_metrics_for_all_schedulers - ): - self.send_metrics_from_scheduler = get_zmq_socket( - context, zmq.PUSH, port_args.metrics_ipc_name, False - ) + ) + self.ipc_channels = SchedulerIpcChannels.create( + 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 + and ( + self.ps.attn_tp_rank == 0 + or self.server_args.enable_metrics_for_all_schedulers + ), + ) def init_idle_sleeper(self) -> None: if ( @@ -780,8 +751,8 @@ class Scheduler( ): self.idle_sleeper = IdleSleeper( sockets=[ - self.recv_from_tokenizer, - self.recv_from_rpc, + self.ipc_channels.recv_from_tokenizer, + self.ipc_channels.recv_from_rpc, ], ) else: @@ -919,7 +890,7 @@ class Scheduler( self.external_corpus_manager = ExternalCorpusManager( self.draft_worker, - self.send_to_tokenizer.send_output, + self.ipc_channels.send_to_tokenizer.send_output, ) else: self.external_corpus_manager = None @@ -1660,10 +1631,10 @@ class Scheduler( output = self._request_dispatcher(recv_req) if output is not None: if not isinstance(output, RpcReqOutput): - self.send_to_tokenizer.send_output(output, recv_req) + self.ipc_channels.send_to_tokenizer.send_output(output, recv_req) else: - if self.recv_from_rpc is not None: - self.recv_from_rpc.send_pyobj(output) + if self.ipc_channels.recv_from_rpc is not None: + self.ipc_channels.recv_from_rpc.send_pyobj(output) self._check_pending_flush() if self.external_corpus_manager is not None: @@ -2093,7 +2064,7 @@ class Scheduler( rid=req.rid, ) req.time_stats.trace_ctx.abort(abort_info=abort_req.finished_reason) - self.send_to_tokenizer.send_output(abort_req, req) + self.ipc_channels.send_to_tokenizer.send_output(abort_req, req) return False return True @@ -2132,7 +2103,7 @@ class Scheduler( req_to_abort = candidate_req message = "The request is aborted by a higher priority request." - self.send_to_tokenizer.send_output( + self.ipc_channels.send_to_tokenizer.send_output( AbortReq( finished_reason={ "type": "abort", @@ -2158,7 +2129,7 @@ class Scheduler( if self.enable_hicache_storage: # Release prefetch events associated with the request self.tree_cache.release_aborted_request(req.rid) - self.send_to_tokenizer.send_output( + self.ipc_channels.send_to_tokenizer.send_output( AbortReq( finished_reason={ "type": "abort", @@ -2750,7 +2721,7 @@ class Scheduler( self.new_token_ratio = new_token_ratio for req in reqs_to_abort: abort_reason: FINISH_ABORT = req.to_finish - self.send_to_tokenizer.send_output( + self.ipc_channels.send_to_tokenizer.send_output( AbortReq( finished_reason=abort_reason.to_json(), rid=req.rid, @@ -2971,7 +2942,7 @@ class Scheduler( tp_active_ranks_cpu = self.tp_group.active_ranks_cpu.detach().numpy() tp_active_ranks &= tp_active_ranks_cpu dp_active_ranks = tp_active_ranks.reshape(self.ps.dp_size, -1).prod(axis=1) - self.send_to_tokenizer.send_output( + self.ipc_channels.send_to_tokenizer.send_output( ActiveRanksOutput(status=dp_active_ranks.tolist()) ) @@ -3039,7 +3010,7 @@ class Scheduler( # Return some signal for the health check. # This is used to prevent the health check signal being blocked by long context prefill. # However, one minor issue is that this code path does not check the status of detokenizer manager. - self.send_to_tokenizer.send_output( + self.ipc_channels.send_to_tokenizer.send_output( HealthCheckOutput( http_worker_ipc=self.return_health_check_ipcs.popleft() ) @@ -3054,7 +3025,7 @@ class Scheduler( if self.is_fully_idle(): success = self.flush_cache() self._pending_flush = None - self.send_to_tokenizer.send_output( + self.ipc_channels.send_to_tokenizer.send_output( FlushCacheReqOutput(success=success), pending_req ) return @@ -3064,7 +3035,7 @@ class Scheduler( "Deferred flush_cache timed out while waiting for idle state." ) self._pending_flush = None - self.send_to_tokenizer.send_output( + self.ipc_channels.send_to_tokenizer.send_output( FlushCacheReqOutput( success=False, message="Timed out waiting for idle state." ), @@ -3461,7 +3432,7 @@ class Scheduler( if self.enable_hicache_storage: # to release prefetch events associated with the request self.tree_cache.release_aborted_request(req.rid) - self.send_to_tokenizer.send_output(AbortReq(rid=req.rid), req) + self.ipc_channels.send_to_tokenizer.send_output(AbortReq(rid=req.rid), req) # For disaggregation decode mode, the request in the waiting queue has KV cache allocated. if self.disaggregation_mode == DisaggregationMode.DECODE: release_kv_cache(req, self.tree_cache) @@ -3525,7 +3496,7 @@ class Scheduler( if recv_req.abort_all or decode_req.rid.startswith(recv_req.rid): assert hasattr(decode_req, "kv_cache_cpu") del decode_req.kv_cache_cpu - self.send_to_tokenizer.send_output( + self.ipc_channels.send_to_tokenizer.send_output( AbortReq(rid=decode_req.rid), decode_req ) else: @@ -3687,7 +3658,7 @@ class Scheduler( def handle_freeze_gc(self, recv_req: FreezeGCReq): """Handle freeze_gc request: freeze scheduler's GC and forward to detokenizer.""" freeze_gc("Scheduler") - self.send_to_detokenizer.send_output(recv_req, recv_req) + self.ipc_channels.send_to_detokenizer.send_output(recv_req, recv_req) return None def handle_dumper_control(self, recv_req: DumperControlReqInput): @@ -3702,12 +3673,12 @@ class Scheduler( response = dumper._http_manager.handle_request( method=recv_req.method, body=recv_req.body ) - self.send_to_tokenizer.send_output( + self.ipc_channels.send_to_tokenizer.send_output( DumperControlReqOutput(success=True, response=response), recv_req ) except Exception as e: print(f"[Scheduler] handle_dumper_control error: {e}", flush=True) - self.send_to_tokenizer.send_output( + self.ipc_channels.send_to_tokenizer.send_output( DumperControlReqOutput(success=False, response=[], error=str(e)), recv_req, ) diff --git a/python/sglang/srt/managers/scheduler_components/ipc_channels.py b/python/sglang/srt/managers/scheduler_components/ipc_channels.py new file mode 100644 index 000000000..246655bf0 --- /dev/null +++ b/python/sglang/srt/managers/scheduler_components/ipc_channels.py @@ -0,0 +1,73 @@ +from dataclasses import dataclass +from typing import Optional + +import zmq + +from sglang.srt.managers.scheduler_components.output_sender import SenderWrapper +from sglang.srt.server_args import PortArgs +from sglang.srt.utils.network import get_zmq_socket + + +@dataclass(frozen=True, slots=True, kw_only=True) +class SchedulerIpcChannels: + recv_from_tokenizer: Optional[zmq.Socket] + recv_from_rpc: Optional[zmq.Socket] + send_to_tokenizer: SenderWrapper + send_to_detokenizer: SenderWrapper + send_metrics_from_scheduler: Optional[zmq.Socket] + + @classmethod + def create( + cls, + *, + port_args: PortArgs, + is_rank_zero: bool, + skip_tokenizer_init: bool, + metrics_enabled: bool, + ) -> "SchedulerIpcChannels": + context = zmq.Context(2) + + if is_rank_zero: + recv_from_tokenizer = get_zmq_socket( + context, zmq.PULL, port_args.scheduler_input_ipc_name, False + ) + recv_from_rpc = get_zmq_socket( + context, zmq.DEALER, port_args.rpc_ipc_name, False + ) + + send_to_tokenizer_raw = get_zmq_socket( + context, zmq.PUSH, port_args.tokenizer_ipc_name, False + ) + if skip_tokenizer_init: + # Directly send to the TokenizerManager + send_to_detokenizer_raw = get_zmq_socket( + context, zmq.PUSH, port_args.tokenizer_ipc_name, False + ) + else: + # Send to the DetokenizerManager + send_to_detokenizer_raw = get_zmq_socket( + context, zmq.PUSH, port_args.detokenizer_ipc_name, False + ) + + send_to_tokenizer = SenderWrapper(send_to_tokenizer_raw) + send_to_detokenizer = SenderWrapper(send_to_detokenizer_raw) + else: + recv_from_tokenizer = None + recv_from_rpc = None + send_to_tokenizer = SenderWrapper(None) + send_to_detokenizer = SenderWrapper(None) + + if metrics_enabled: + send_metrics_from_scheduler = get_zmq_socket( + context, zmq.PUSH, port_args.metrics_ipc_name, False + ) + else: + send_metrics_from_scheduler = None + + return cls( + recv_from_tokenizer=recv_from_tokenizer, + recv_from_rpc=recv_from_rpc, + send_to_tokenizer=send_to_tokenizer, + send_to_detokenizer=send_to_detokenizer, + send_metrics_from_scheduler=send_metrics_from_scheduler, + ) diff --git a/test/registered/unit/managers/test_priority_scheduling_disaggregation.py b/test/registered/unit/managers/test_priority_scheduling_disaggregation.py index 3f287001c..8468ad4e9 100644 --- a/test/registered/unit/managers/test_priority_scheduling_disaggregation.py +++ b/test/registered/unit/managers/test_priority_scheduling_disaggregation.py @@ -30,7 +30,7 @@ class TestDisaggregationPriorityQueueing(unittest.TestCase): scheduler.model_config = SimpleNamespace(num_key_value_heads=8) scheduler.disagg_prefill_bootstrap_queue = MagicMock() scheduler.disagg_decode_prealloc_queue = MagicMock() - scheduler.send_to_tokenizer = MagicMock() + scheduler.ipc_channels = MagicMock() return scheduler def _new_req(self, priority=None): @@ -72,7 +72,7 @@ class TestDisaggregationPriorityQueueing(unittest.TestCase): scheduler._add_request_to_queue(req) scheduler.disagg_decode_prealloc_queue.add.assert_not_called() - scheduler.send_to_tokenizer.send_output.assert_called_once() + scheduler.ipc_channels.send_to_tokenizer.send_output.assert_called_once() req.time_stats.trace_ctx.abort.assert_called_once() diff --git a/test/registered/unit/managers/test_scheduler_flush_cache.py b/test/registered/unit/managers/test_scheduler_flush_cache.py index 73a522d26..3693d725c 100644 --- a/test/registered/unit/managers/test_scheduler_flush_cache.py +++ b/test/registered/unit/managers/test_scheduler_flush_cache.py @@ -16,7 +16,7 @@ class TestSchedulerFlushCache(unittest.TestCase): def _new_scheduler(self) -> Scheduler: scheduler = Scheduler.__new__(Scheduler) scheduler._pending_flush = None - scheduler.send_to_tokenizer = MagicMock() + scheduler.ipc_channels = MagicMock() scheduler.flush_cache = MagicMock(return_value=True) scheduler.is_fully_idle = MagicMock(return_value=False) return scheduler @@ -82,7 +82,7 @@ class TestSchedulerFlushCache(unittest.TestCase): self.assertIsNone(scheduler._pending_flush) scheduler.flush_cache.assert_called_once() - out = scheduler.send_to_tokenizer.send_output.call_args.args[0] + out = scheduler.ipc_channels.send_to_tokenizer.send_output.call_args.args[0] self.assertTrue(out.success) def test_pending_flush_expires_on_timeout(self): @@ -95,7 +95,7 @@ class TestSchedulerFlushCache(unittest.TestCase): self.assertIsNone(scheduler._pending_flush) scheduler.flush_cache.assert_not_called() - out = scheduler.send_to_tokenizer.send_output.call_args.args[0] + out = scheduler.ipc_channels.send_to_tokenizer.send_output.call_args.args[0] self.assertFalse(out.success) def test_pending_flush_survives_before_deadline(self): @@ -107,7 +107,7 @@ class TestSchedulerFlushCache(unittest.TestCase): Scheduler._check_pending_flush(scheduler) self.assertIsNotNone(scheduler._pending_flush) - scheduler.send_to_tokenizer.send_output.assert_not_called() + scheduler.ipc_channels.send_to_tokenizer.send_output.assert_not_called() if __name__ == "__main__":