Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions python/sglang/srt/managers/io_struct.py
Original file line number Diff line number Diff line change
Expand Up @@ -1223,6 +1223,20 @@ class ContinueGenerationReqInput(BaseReq):
pass


@dataclass
class TokenizerWorkerRegistration:
"""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
Expand Down
108 changes: 102 additions & 6 deletions python/sglang/srt/managers/multi_tokenizer_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@
import threading
from functools import partialmethod
from multiprocessing import shared_memory
from typing import TYPE_CHECKING, Any, Dict, Union
from typing import TYPE_CHECKING, Any, Dict, Optional, Union

import setproctitle
import zmq
Expand All @@ -42,6 +42,10 @@
BatchEmbeddingOutput,
BatchStrOutput,
BatchTokenIDOutput,
ContinueGenerationReqInput,
PauseContinueBroadcast,
PauseGenerationReqInput,
TokenizerWorkerRegistration,
)
from sglang.srt.managers.tokenizer_communicator_mixin import _Communicator
from sglang.srt.managers.tokenizer_manager import TokenizerManager
Expand Down Expand Up @@ -315,7 +319,12 @@ def multi_http_worker_event_loop(self: DetokenizerManager):


class MultiTokenizerRouter:
"""A router to receive requests from TokenizerWorker"""
"""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.
"""

def __init__(
self,
Expand All @@ -339,29 +348,59 @@ def __init__(
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, TokenizerWorkerRegistration):
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):
# special reqs will recv from scheduler, need to route to right worker
self.socket_mapping = SocketMapping()
"""Backward path: detokenizer → route results to correct worker."""
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):
Expand Down Expand Up @@ -404,6 +443,63 @@ def __init__(
self.send_to_scheduler, 2
)

# Register this worker with the router for pause/continue broadcasting
reg = TokenizerWorkerRegistration(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):
Expand Down
Loading