Pack scattered scheduler IPC channel state into a dedicated container (#25714)

This commit is contained in:
fzyzcjy
2026-05-19 09:19:02 +08:00
committed by GitHub
parent 07b4f262b7
commit 954b5c5846
4 changed files with 117 additions and 73 deletions
+38 -67
View File
@@ -33,7 +33,6 @@ import psutil
import setproctitle import setproctitle
import torch import torch
import torch.distributed import torch.distributed
import zmq
from torch.cuda import Stream as CudaStream from torch.cuda import Stream as CudaStream
from torch.distributed import barrier from torch.distributed import barrier
@@ -171,6 +170,9 @@ from sglang.srt.managers.scheduler_components.invariant_checker import (
SchedulerInvariantChecker, SchedulerInvariantChecker,
create_scheduler_watchdog, create_scheduler_watchdog,
) )
from sglang.srt.managers.scheduler_components.ipc_channels import (
SchedulerIpcChannels,
)
from sglang.srt.managers.scheduler_components.kv_events_publisher import ( from sglang.srt.managers.scheduler_components.kv_events_publisher import (
SchedulerKvEventsPublisher, SchedulerKvEventsPublisher,
) )
@@ -185,7 +187,6 @@ 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,
) )
@@ -250,7 +251,6 @@ from sglang.srt.utils.hf_transformers_utils import (
get_tokenizer, get_tokenizer,
get_tokenizer_from_processor, 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.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
@@ -552,8 +552,8 @@ class Scheduler(
self.grammar_manager = GrammarManager(self) self.grammar_manager = GrammarManager(self)
self.request_receiver = SchedulerRequestReceiver( self.request_receiver = SchedulerRequestReceiver(
recv_from_tokenizer=self.recv_from_tokenizer, recv_from_tokenizer=self.ipc_channels.recv_from_tokenizer,
recv_from_rpc=self.recv_from_rpc, recv_from_rpc=self.ipc_channels.recv_from_rpc,
recv_skipper=self.recv_skipper, recv_skipper=self.recv_skipper,
input_blocker=self.input_blocker, input_blocker=self.input_blocker,
mm_receiver=self.mm_receiver, mm_receiver=self.mm_receiver,
@@ -629,7 +629,7 @@ class Scheduler(
attn_dp_rank=self.ps.attn_dp_rank, attn_dp_rank=self.ps.attn_dp_rank,
dp_rank=self.ps.dp_rank, dp_rank=self.ps.dp_rank,
tree_cache=self.tree_cache, 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_running_requests=self.max_running_requests,
max_total_num_tokens=self.max_total_num_tokens, max_total_num_tokens=self.max_total_num_tokens,
get_stats=lambda: self.metrics_reporter.stats, get_stats=lambda: self.metrics_reporter.stats,
@@ -658,7 +658,7 @@ class Scheduler(
) )
self.output_streamer = SchedulerOutputStreamer( 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, tree_cache=self.tree_cache,
ps=self.ps, ps=self.ps,
server_args=self.server_args, server_args=self.server_args,
@@ -726,50 +726,21 @@ class Scheduler(
self.page_size = self.dllm_config.block_size self.page_size = self.dllm_config.block_size
def init_ipc_channels(self, port_args: PortArgs): def init_ipc_channels(self, port_args: PortArgs):
context = zmq.Context(2) is_rank_zero = (
self.send_metrics_from_scheduler = None
if (
self.ps.pp_rank == 0 self.ps.pp_rank == 0
and self.ps.attn_tp_rank == 0 and self.ps.attn_tp_rank == 0
and self.ps.attn_cp_rank == 0 and self.ps.attn_cp_rank == 0
): )
self.recv_from_tokenizer = get_zmq_socket( self.ipc_channels = SchedulerIpcChannels.create(
context, zmq.PULL, port_args.scheduler_input_ipc_name, False port_args=port_args,
) is_rank_zero=is_rank_zero,
self.recv_from_rpc = get_zmq_socket( skip_tokenizer_init=self.server_args.skip_tokenizer_init,
context, zmq.DEALER, port_args.rpc_ipc_name, False metrics_enabled=self.server_args.enable_metrics
) and (
self.ps.attn_tp_rank == 0
send_to_tokenizer = get_zmq_socket( or self.server_args.enable_metrics_for_all_schedulers
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
)
def init_idle_sleeper(self) -> None: def init_idle_sleeper(self) -> None:
if ( if (
@@ -780,8 +751,8 @@ class Scheduler(
): ):
self.idle_sleeper = IdleSleeper( self.idle_sleeper = IdleSleeper(
sockets=[ sockets=[
self.recv_from_tokenizer, self.ipc_channels.recv_from_tokenizer,
self.recv_from_rpc, self.ipc_channels.recv_from_rpc,
], ],
) )
else: else:
@@ -919,7 +890,7 @@ class Scheduler(
self.external_corpus_manager = ExternalCorpusManager( self.external_corpus_manager = ExternalCorpusManager(
self.draft_worker, self.draft_worker,
self.send_to_tokenizer.send_output, self.ipc_channels.send_to_tokenizer.send_output,
) )
else: else:
self.external_corpus_manager = None self.external_corpus_manager = None
@@ -1660,10 +1631,10 @@ class Scheduler(
output = self._request_dispatcher(recv_req) output = self._request_dispatcher(recv_req)
if output is not None: if output is not None:
if not isinstance(output, RpcReqOutput): 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: else:
if self.recv_from_rpc is not None: if self.ipc_channels.recv_from_rpc is not None:
self.recv_from_rpc.send_pyobj(output) self.ipc_channels.recv_from_rpc.send_pyobj(output)
self._check_pending_flush() self._check_pending_flush()
if self.external_corpus_manager is not None: if self.external_corpus_manager is not None:
@@ -2093,7 +2064,7 @@ class Scheduler(
rid=req.rid, rid=req.rid,
) )
req.time_stats.trace_ctx.abort(abort_info=abort_req.finished_reason) 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 False
return True return True
@@ -2132,7 +2103,7 @@ class Scheduler(
req_to_abort = candidate_req req_to_abort = candidate_req
message = "The request is aborted by a higher priority request." 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( AbortReq(
finished_reason={ finished_reason={
"type": "abort", "type": "abort",
@@ -2158,7 +2129,7 @@ class Scheduler(
if self.enable_hicache_storage: if self.enable_hicache_storage:
# Release prefetch events associated with the request # Release prefetch events associated with the request
self.tree_cache.release_aborted_request(req.rid) self.tree_cache.release_aborted_request(req.rid)
self.send_to_tokenizer.send_output( self.ipc_channels.send_to_tokenizer.send_output(
AbortReq( AbortReq(
finished_reason={ finished_reason={
"type": "abort", "type": "abort",
@@ -2750,7 +2721,7 @@ class Scheduler(
self.new_token_ratio = new_token_ratio self.new_token_ratio = new_token_ratio
for req in reqs_to_abort: for req in reqs_to_abort:
abort_reason: FINISH_ABORT = req.to_finish abort_reason: FINISH_ABORT = req.to_finish
self.send_to_tokenizer.send_output( self.ipc_channels.send_to_tokenizer.send_output(
AbortReq( AbortReq(
finished_reason=abort_reason.to_json(), finished_reason=abort_reason.to_json(),
rid=req.rid, rid=req.rid,
@@ -2971,7 +2942,7 @@ class Scheduler(
tp_active_ranks_cpu = self.tp_group.active_ranks_cpu.detach().numpy() tp_active_ranks_cpu = self.tp_group.active_ranks_cpu.detach().numpy()
tp_active_ranks &= tp_active_ranks_cpu tp_active_ranks &= tp_active_ranks_cpu
dp_active_ranks = tp_active_ranks.reshape(self.ps.dp_size, -1).prod(axis=1) 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()) ActiveRanksOutput(status=dp_active_ranks.tolist())
) )
@@ -3039,7 +3010,7 @@ class Scheduler(
# Return some signal for the health check. # Return some signal for the health check.
# This is used to prevent the health check signal being blocked by long context prefill. # 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. # 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( HealthCheckOutput(
http_worker_ipc=self.return_health_check_ipcs.popleft() http_worker_ipc=self.return_health_check_ipcs.popleft()
) )
@@ -3054,7 +3025,7 @@ class Scheduler(
if self.is_fully_idle(): if self.is_fully_idle():
success = self.flush_cache() success = self.flush_cache()
self._pending_flush = None self._pending_flush = None
self.send_to_tokenizer.send_output( self.ipc_channels.send_to_tokenizer.send_output(
FlushCacheReqOutput(success=success), pending_req FlushCacheReqOutput(success=success), pending_req
) )
return return
@@ -3064,7 +3035,7 @@ class Scheduler(
"Deferred flush_cache timed out while waiting for idle state." "Deferred flush_cache timed out while waiting for idle state."
) )
self._pending_flush = None self._pending_flush = None
self.send_to_tokenizer.send_output( self.ipc_channels.send_to_tokenizer.send_output(
FlushCacheReqOutput( FlushCacheReqOutput(
success=False, message="Timed out waiting for idle state." success=False, message="Timed out waiting for idle state."
), ),
@@ -3461,7 +3432,7 @@ class Scheduler(
if self.enable_hicache_storage: if self.enable_hicache_storage:
# to release prefetch events associated with the request # to release prefetch events associated with the request
self.tree_cache.release_aborted_request(req.rid) 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. # For disaggregation decode mode, the request in the waiting queue has KV cache allocated.
if self.disaggregation_mode == DisaggregationMode.DECODE: if self.disaggregation_mode == DisaggregationMode.DECODE:
release_kv_cache(req, self.tree_cache) 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): if recv_req.abort_all or decode_req.rid.startswith(recv_req.rid):
assert hasattr(decode_req, "kv_cache_cpu") assert hasattr(decode_req, "kv_cache_cpu")
del 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 AbortReq(rid=decode_req.rid), decode_req
) )
else: else:
@@ -3687,7 +3658,7 @@ class Scheduler(
def handle_freeze_gc(self, recv_req: FreezeGCReq): def handle_freeze_gc(self, recv_req: FreezeGCReq):
"""Handle freeze_gc request: freeze scheduler's GC and forward to detokenizer.""" """Handle freeze_gc request: freeze scheduler's GC and forward to detokenizer."""
freeze_gc("Scheduler") 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 return None
def handle_dumper_control(self, recv_req: DumperControlReqInput): def handle_dumper_control(self, recv_req: DumperControlReqInput):
@@ -3702,12 +3673,12 @@ class Scheduler(
response = dumper._http_manager.handle_request( response = dumper._http_manager.handle_request(
method=recv_req.method, body=recv_req.body 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 DumperControlReqOutput(success=True, response=response), recv_req
) )
except Exception as e: except Exception as e:
print(f"[Scheduler] handle_dumper_control error: {e}", flush=True) 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)), DumperControlReqOutput(success=False, response=[], error=str(e)),
recv_req, recv_req,
) )
@@ -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,
)
@@ -30,7 +30,7 @@ class TestDisaggregationPriorityQueueing(unittest.TestCase):
scheduler.model_config = SimpleNamespace(num_key_value_heads=8) scheduler.model_config = SimpleNamespace(num_key_value_heads=8)
scheduler.disagg_prefill_bootstrap_queue = MagicMock() scheduler.disagg_prefill_bootstrap_queue = MagicMock()
scheduler.disagg_decode_prealloc_queue = MagicMock() scheduler.disagg_decode_prealloc_queue = MagicMock()
scheduler.send_to_tokenizer = MagicMock() scheduler.ipc_channels = MagicMock()
return scheduler return scheduler
def _new_req(self, priority=None): def _new_req(self, priority=None):
@@ -72,7 +72,7 @@ class TestDisaggregationPriorityQueueing(unittest.TestCase):
scheduler._add_request_to_queue(req) scheduler._add_request_to_queue(req)
scheduler.disagg_decode_prealloc_queue.add.assert_not_called() 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() req.time_stats.trace_ctx.abort.assert_called_once()
@@ -16,7 +16,7 @@ class TestSchedulerFlushCache(unittest.TestCase):
def _new_scheduler(self) -> Scheduler: def _new_scheduler(self) -> Scheduler:
scheduler = Scheduler.__new__(Scheduler) scheduler = Scheduler.__new__(Scheduler)
scheduler._pending_flush = None scheduler._pending_flush = None
scheduler.send_to_tokenizer = MagicMock() scheduler.ipc_channels = MagicMock()
scheduler.flush_cache = MagicMock(return_value=True) scheduler.flush_cache = MagicMock(return_value=True)
scheduler.is_fully_idle = MagicMock(return_value=False) scheduler.is_fully_idle = MagicMock(return_value=False)
return scheduler return scheduler
@@ -82,7 +82,7 @@ class TestSchedulerFlushCache(unittest.TestCase):
self.assertIsNone(scheduler._pending_flush) self.assertIsNone(scheduler._pending_flush)
scheduler.flush_cache.assert_called_once() 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) self.assertTrue(out.success)
def test_pending_flush_expires_on_timeout(self): def test_pending_flush_expires_on_timeout(self):
@@ -95,7 +95,7 @@ class TestSchedulerFlushCache(unittest.TestCase):
self.assertIsNone(scheduler._pending_flush) self.assertIsNone(scheduler._pending_flush)
scheduler.flush_cache.assert_not_called() 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) self.assertFalse(out.success)
def test_pending_flush_survives_before_deadline(self): def test_pending_flush_survives_before_deadline(self):
@@ -107,7 +107,7 @@ class TestSchedulerFlushCache(unittest.TestCase):
Scheduler._check_pending_flush(scheduler) Scheduler._check_pending_flush(scheduler)
self.assertIsNotNone(scheduler._pending_flush) 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__": if __name__ == "__main__":