fix(load-snapshot): avoid duplicate zmq bind in multi-tokenizer mode (#27145)

Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
This commit is contained in:
ybyang
2026-06-03 18:24:56 -07:00
committed by GitHub
co-authored by Lianmin Zheng
parent 14ed9b448e
commit 687baf9471
5 changed files with 157 additions and 19 deletions
@@ -165,7 +165,7 @@ class DataParallelController:
self.load_snapshot_reader = create_load_snapshot_reader(
server_args,
port_args,
caller="dp_controller",
caller="DataParallelController",
)
self._last_refresh_time = 0.0
+52 -15
View File
@@ -28,7 +28,8 @@ transport backends are supported:
Shared memory does not work across nodes, so multi-node DP attention
requires the ZMQ transport. The ``ZmqShmLoadSnapshotReader`` on node 0
receives snapshots from all schedulers via zmq PUSH/PULL and writes them
into the local SHM file. All readers (tokenizer, dp_controller) on
into the local SHM file. All readers (TokenizerManager,
DataParallelController) on
node 0 then read from SHM.
``zmq_reader_owner()`` decides which process on node 0 binds the zmq
@@ -89,30 +90,49 @@ def should_use_zmq(server_args) -> bool:
_LOAD_AWARE_METHODS = frozenset({"total_requests", "total_tokens"})
def _tokenizer_load_snapshot_owner_caller(server_args) -> str:
"""The caller that plays the tokenizer-side zmq owner role.
In multi-tokenizer mode (``tokenizer_worker_num > 1``) there are N
independent ``TokenizerWorker`` processes that would all try to bind the
same zmq PULL endpoint. Instead, the single ``MultiTokenizerRouter``
process owns the socket (polls zmq -> SHM) and every worker reads SHM.
"""
if server_args.tokenizer_worker_num > 1:
return "MultiTokenizerRouter"
return "TokenizerManager"
def zmq_reader_owner(server_args, caller: str) -> bool:
"""Decide which process owns the zmq PULL socket.
Exactly one of ``"dp_controller"`` or ``"tokenizer"`` must return True
when zmq mode is active. The owner polls zmq -> SHM; the other reads SHM.
Exactly one of ``"DataParallelController"``, ``"TokenizerManager"``, or
``"MultiTokenizerRouter"`` must return True when zmq mode is active. The
owner polls zmq -> SHM; the others read SHM.
Rules:
- Non-zero node_rank: no tokenizer, dp_controller only launches
schedulers and waits -> nobody owns it.
- dp_size == 1: no dp_controller exists -> tokenizer owns it.
- dp_size > 1, load-aware method: dp_controller polls on every
dispatch via refresh_load_budget() -> dp_controller owns it.
- dp_size > 1, round-robin / other: dp_controller never reads
load data -> tokenizer owns it (polls on /v1/loads calls).
- Non-zero node_rank: no TokenizerManager, DataParallelController only
launches schedulers and waits -> nobody owns it.
- dp_size == 1: no DataParallelController exists -> tokenizer-side owner
owns it.
- dp_size > 1, load-aware method: DataParallelController polls on every
dispatch via refresh_load_budget() -> DataParallelController owns it.
- dp_size > 1, round-robin / other: DataParallelController never reads
load data -> tokenizer-side owner owns it (polls on /v1/loads calls).
The tokenizer-side owner is the ``"MultiTokenizerRouter"`` caller in
multi-tokenizer mode, otherwise the ``"TokenizerManager"`` caller.
"""
if not should_use_zmq(server_args):
return False
if server_args.node_rank != 0:
return False
tokenizer_owner = _tokenizer_load_snapshot_owner_caller(server_args)
if server_args.dp_size == 1:
return caller == "tokenizer"
return caller == tokenizer_owner
if server_args.load_balance_method.lower() in _LOAD_AWARE_METHODS:
return caller == "dp_controller"
return caller == "tokenizer"
return caller == "DataParallelController"
return caller == tokenizer_owner
# ---------------------------------------------------------------------------
@@ -619,6 +639,22 @@ class ZmqShmLoadSnapshotReader:
"load snapshot shm write failed for rank %d: %s", dp_rank, e
)
def fileno(self) -> int:
"""Edge-triggered fd that becomes readable when zmq messages arrive.
Lets an owner process register the reader with an event loop and drain
it via ``poll()`` instead of polling on a timer.
"""
return self._socket.getsockopt(self._zmq.FD)
def poll(self) -> None:
"""Drain the zmq PULL socket into SHM.
Public entry point so an owner process (e.g. MultiTokenizerRouter) can
keep SHM fresh without touching internals.
"""
self._poll()
def read(self, dp_rank: int) -> Optional[LoadSnapshot]:
self._poll()
return self._shm_reader.read(dp_rank)
@@ -683,8 +719,9 @@ def create_load_snapshot_reader(server_args, port_args, caller: str):
"""Create a load snapshot reader.
Args:
caller: ``"dp_controller"`` or ``"tokenizer"`` -- determines who
binds the zmq PULL socket when zmq mode is active.
caller: ``"DataParallelController"``, ``"TokenizerManager"``, or
``"MultiTokenizerRouter"`` -- determines who binds the zmq PULL
socket when zmq mode is active.
"""
dp_size = server_args.dp_size
if zmq_reader_owner(server_args, caller):
@@ -50,6 +50,10 @@ from sglang.srt.managers.io_struct import (
PauseGenerationReqInput,
TokenizerWorkerRegistration,
)
from sglang.srt.managers.load_snapshot import (
create_load_snapshot_reader,
zmq_reader_owner,
)
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import (
@@ -381,6 +385,19 @@ class MultiTokenizerRouter:
self._handle_task = asyncio.run_coroutine_threadsafe(
print_exception_wrapper(self.handle_loop), self._loop
)
# In multi-tokenizer mode the N TokenizerWorker processes cannot each
# bind the zmq PULL socket used for load snapshots, so the single
# MultiTokenizerRouter process owns it (zmq -> SHM) and the workers
# read SHM only. Drain it event-driven via the socket's fd instead of
# polling on a timer.
self.load_snapshot_reader = None
if zmq_reader_owner(server_args, "MultiTokenizerRouter"):
self.load_snapshot_reader = create_load_snapshot_reader(
server_args, port_args, caller="MultiTokenizerRouter"
)
self._loop.call_soon_threadsafe(self._register_load_snapshot_reader)
self.disaggregation_bootstrap_server = start_disagg_service(self.server_args)
# Worker IPC names for pause/continue broadcasting
@@ -391,6 +408,20 @@ class MultiTokenizerRouter:
def _run_loop(self):
self._loop.run_forever()
def _register_load_snapshot_reader(self):
"""Drain zmq load snapshots into SHM whenever the PULL socket is readable.
zmq exposes an edge-triggered fd; ``poll()`` drains it until empty, which
also re-arms the fd, so TokenizerWorkers reading SHM stay up to date
without any timer.
"""
assert self.load_snapshot_reader is not None
self._loop.add_reader(
self.load_snapshot_reader.fileno(), self.load_snapshot_reader.poll
)
# Drain anything already queued before the fd was registered.
self.load_snapshot_reader.poll()
async def router_worker_obj(self):
"""Forward path: workers → scheduler, with pause/continue broadcast."""
while True:
@@ -385,7 +385,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.load_snapshot_reader = create_load_snapshot_reader(
self.server_args,
port_args,
caller="tokenizer",
caller="TokenizerManager",
)
def init_running_status(self):