[fix] /pause_generation and /continue_generation wrong for --tokenizer-worker-num > 1 (#24462)

Co-authored-by: lawrence-harmonic <185285563+lawrence-harmonic@users.noreply.github.com>
This commit is contained in:
maocheng23
2026-05-07 21:32:21 -07:00
committed by GitHub
co-authored by lawrence-harmonic
parent 2afb450501
commit 7deed98e1b
2 changed files with 116 additions and 6 deletions
+14
View File
@@ -1381,6 +1381,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
@@ -28,7 +28,7 @@ import sys
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
@@ -43,6 +43,10 @@ from sglang.srt.managers.io_struct import (
BatchEmbeddingOutput,
BatchStrOutput,
BatchTokenIDOutput,
ContinueGenerationReqInput,
PauseContinueBroadcast,
PauseGenerationReqInput,
TokenizerWorkerRegistration,
)
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.server_args import PortArgs, ServerArgs
@@ -318,7 +322,12 @@ class MultiHttpWorkerDetokenizerMixin:
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,
@@ -342,29 +351,59 @@ 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, 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):
@@ -407,6 +446,63 @@ class TokenizerWorker(TokenizerManager):
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):