Revert "[fix] /pause_generation and /continue_generation wrong for --tokenizer-worker-num > 1" (#24461)
This commit is contained in:
@@ -1381,20 +1381,6 @@ class ContinueGenerationReqInput(BaseReq):
|
|||||||
pass
|
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
|
@dataclass
|
||||||
class UpdateWeightFromDiskReqInput(BaseReq):
|
class UpdateWeightFromDiskReqInput(BaseReq):
|
||||||
# The model path with the new weights
|
# The model path with the new weights
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ import sys
|
|||||||
import threading
|
import threading
|
||||||
from functools import partialmethod
|
from functools import partialmethod
|
||||||
from multiprocessing import shared_memory
|
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 setproctitle
|
||||||
import zmq
|
import zmq
|
||||||
@@ -43,10 +43,6 @@ from sglang.srt.managers.io_struct import (
|
|||||||
BatchEmbeddingOutput,
|
BatchEmbeddingOutput,
|
||||||
BatchStrOutput,
|
BatchStrOutput,
|
||||||
BatchTokenIDOutput,
|
BatchTokenIDOutput,
|
||||||
ContinueGenerationReqInput,
|
|
||||||
PauseContinueBroadcast,
|
|
||||||
PauseGenerationReqInput,
|
|
||||||
TokenizerWorkerRegisterReq,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
@@ -322,12 +318,7 @@ class MultiHttpWorkerDetokenizerMixin:
|
|||||||
|
|
||||||
|
|
||||||
class MultiTokenizerRouter:
|
class MultiTokenizerRouter:
|
||||||
"""A router between tokenizer managers and the scheduler/detokenizer manager.
|
"""A router to receive requests from TokenizerWorker"""
|
||||||
|
|
||||||
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -351,59 +342,29 @@ class MultiTokenizerRouter:
|
|||||||
self._task = asyncio.run_coroutine_threadsafe(
|
self._task = asyncio.run_coroutine_threadsafe(
|
||||||
self.router_worker_obj(), self._loop
|
self.router_worker_obj(), self._loop
|
||||||
)
|
)
|
||||||
|
# Start handle_loop simultaneously
|
||||||
self._handle_task = asyncio.run_coroutine_threadsafe(
|
self._handle_task = asyncio.run_coroutine_threadsafe(
|
||||||
print_exception_wrapper(self.handle_loop), self._loop
|
print_exception_wrapper(self.handle_loop), self._loop
|
||||||
)
|
)
|
||||||
self.disaggregation_bootstrap_server = start_disagg_service(self.server_args)
|
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):
|
def _run_loop(self):
|
||||||
self._loop.run_forever()
|
self._loop.run_forever()
|
||||||
|
|
||||||
async def router_worker_obj(self):
|
async def router_worker_obj(self):
|
||||||
"""Forward path: workers → scheduler, with pause/continue broadcast."""
|
|
||||||
while True:
|
while True:
|
||||||
recv_obj = await self.receive_from_worker.recv_pyobj()
|
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)
|
await self.send_to_scheduler.send_pyobj(recv_obj)
|
||||||
|
|
||||||
async def handle_loop(self):
|
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:
|
while True:
|
||||||
recv_obj = await self.recv_from_detokenizer.recv_pyobj()
|
recv_obj = await self.recv_from_detokenizer.recv_pyobj()
|
||||||
await self._distribute_result_to_workers(recv_obj)
|
await self._distribute_result_to_workers(recv_obj)
|
||||||
|
|
||||||
async def _distribute_result_to_workers(self, recv_obj):
|
async def _distribute_result_to_workers(self, recv_obj):
|
||||||
|
# Distribute result to each worker
|
||||||
if isinstance(recv_obj, BaseReq):
|
if isinstance(recv_obj, BaseReq):
|
||||||
ipc_names = [recv_obj.http_worker_ipc]
|
ipc_names = [recv_obj.http_worker_ipc]
|
||||||
elif isinstance(recv_obj, BaseBatchReq):
|
elif isinstance(recv_obj, BaseBatchReq):
|
||||||
@@ -446,63 +407,6 @@ class TokenizerWorker(TokenizerManager):
|
|||||||
self.send_to_scheduler, 2
|
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]):
|
def _attach_multi_http_worker_info(self, req: Union[BaseReq, BaseBatchReq]):
|
||||||
|
|
||||||
if isinstance(req, BaseReq):
|
if isinstance(req, BaseReq):
|
||||||
|
|||||||
Reference in New Issue
Block a user