diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 56d2228a1..a573e672a 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1381,20 +1381,6 @@ class ContinueGenerationReqInput(BaseReq): pass -@dataclass -class TokenizerWorkerRegisterReq: - """Sent by each TokenizerWorker on startup to register its IPC name with the router.""" - - worker_ipc_name: str - - -@dataclass -class PauseContinueBroadcast: - """Broadcast from router to all workers to set is_pause state.""" - - is_pause: bool - - @dataclass class UpdateWeightFromDiskReqInput(BaseReq): # The model path with the new weights diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index ad8be95ff..0620cb3d6 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -28,7 +28,7 @@ import sys import threading from functools import partialmethod from multiprocessing import shared_memory -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, Union import setproctitle import zmq @@ -43,10 +43,6 @@ from sglang.srt.managers.io_struct import ( BatchEmbeddingOutput, BatchStrOutput, BatchTokenIDOutput, - ContinueGenerationReqInput, - PauseContinueBroadcast, - PauseGenerationReqInput, - TokenizerWorkerRegisterReq, ) from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.server_args import PortArgs, ServerArgs @@ -322,12 +318,7 @@ class MultiHttpWorkerDetokenizerMixin: class MultiTokenizerRouter: - """A router between tokenizer managers and the scheduler/detokenizer manager. - - Forward: tokenizer managers → router → scheduler. - Backward: detokenizer manager → router → tokenizer managers. - Also broadcasts pause/continue to all tokenizer managers for consistent is_pause state. - """ + """A router to receive requests from TokenizerWorker""" def __init__( self, @@ -351,59 +342,29 @@ class MultiTokenizerRouter: self._task = asyncio.run_coroutine_threadsafe( self.router_worker_obj(), self._loop ) + # Start handle_loop simultaneously self._handle_task = asyncio.run_coroutine_threadsafe( print_exception_wrapper(self.handle_loop), self._loop ) self.disaggregation_bootstrap_server = start_disagg_service(self.server_args) - # Worker IPC names for pause/continue broadcasting - self.all_worker_ipcs: set[str] = set() - # Shared socket mapping (both coroutines run on self._loop, so safe) - self.socket_mapping = SocketMapping() - def _run_loop(self): self._loop.run_forever() async def router_worker_obj(self): - """Forward path: workers → scheduler, with pause/continue broadcast.""" while True: recv_obj = await self.receive_from_worker.recv_pyobj() - - if isinstance(recv_obj, TokenizerWorkerRegisterReq): - if recv_obj.worker_ipc_name not in self.all_worker_ipcs: - self.all_worker_ipcs.add(recv_obj.worker_ipc_name) - logger.info( - f"Router registered worker IPC: {recv_obj.worker_ipc_name} " - f"(total: {len(self.all_worker_ipcs)})" - ) - continue - - if isinstance( - recv_obj, (PauseGenerationReqInput, ContinueGenerationReqInput) - ): - # Broadcast to ALL workers so every worker's is_pause is set - is_pause = isinstance(recv_obj, PauseGenerationReqInput) - broadcast = PauseContinueBroadcast(is_pause=is_pause) - for ipc_name in self.all_worker_ipcs: - self.socket_mapping.send_output(ipc_name, broadcast) - # Forward to scheduler rank 0 (it broadcasts to all TP/PP/DP - # ranks internally). Skip for abort mode which drains via polling. - if not ( - isinstance(recv_obj, PauseGenerationReqInput) - and recv_obj.mode == "abort" - ): - await self.send_to_scheduler.send_pyobj(recv_obj) - continue - await self.send_to_scheduler.send_pyobj(recv_obj) async def handle_loop(self): - """Backward path: detokenizer → route results to correct worker.""" + # special reqs will recv from scheduler, need to route to right worker + self.socket_mapping = SocketMapping() while True: recv_obj = await self.recv_from_detokenizer.recv_pyobj() await self._distribute_result_to_workers(recv_obj) async def _distribute_result_to_workers(self, recv_obj): + # Distribute result to each worker if isinstance(recv_obj, BaseReq): ipc_names = [recv_obj.http_worker_ipc] elif isinstance(recv_obj, BaseBatchReq): @@ -446,63 +407,6 @@ class TokenizerWorker(TokenizerManager): self.send_to_scheduler, 2 ) - # Register this worker with the router for pause/continue broadcasting - reg = TokenizerWorkerRegisterReq(worker_ipc_name=self.tokenizer_ipc_name) - self.send_to_scheduler.send_pyobj(reg) - - # Future for awaiting pause/continue broadcast confirmation - self._pause_continue_future: Optional[asyncio.Future] = None - - # Register PauseContinueBroadcast in the result dispatcher so - # handle_loop routes it to _handle_pause_continue_broadcast - from sglang.utils import TypeBasedDispatcher - - self._result_dispatcher += TypeBasedDispatcher( - [(PauseContinueBroadcast, self._handle_pause_continue_broadcast)] - ) - - async def pause_generation(self, obj: PauseGenerationReqInput): - loop = asyncio.get_event_loop() - self._pause_continue_future = loop.create_future() - # Send to router which will broadcast to all workers - # (router also handles forwarding to scheduler for non-abort modes) - self.send_to_scheduler.send_pyobj(obj) - await self._pause_continue_future - - if obj.mode == "abort": - # Abort polling: only the originator checks its own lock state - while True: - self.abort_request(abort_all=True) - is_locked = await self.model_update_lock.is_locked() - if not is_locked: - break - await asyncio.sleep(1.0) - - async def continue_generation(self, obj: ContinueGenerationReqInput): - loop = asyncio.get_event_loop() - self._pause_continue_future = loop.create_future() - self.send_to_scheduler.send_pyobj(obj) - await self._pause_continue_future - - def _handle_pause_continue_broadcast(self, obj: PauseContinueBroadcast): - """Called from handle_loop when a broadcast arrives from the router.""" - loop = asyncio.get_event_loop() - loop.create_task(self._apply_pause_continue_broadcast(obj)) - - async def _apply_pause_continue_broadcast(self, obj: PauseContinueBroadcast): - """Apply pause/continue state under the condition lock.""" - async with self.is_pause_cond: - if obj.is_pause: - self.is_pause = True - else: - self.is_pause = False - self.is_pause_cond.notify_all() - - # Resolve the pending future if this worker initiated the pause/continue - if self._pause_continue_future and not self._pause_continue_future.done(): - self._pause_continue_future.set_result(True) - self._pause_continue_future = None - def _attach_multi_http_worker_info(self, req: Union[BaseReq, BaseBatchReq]): if isinstance(req, BaseReq):