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:
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -16,6 +16,7 @@ from sglang.srt.managers.load_snapshot import (
|
||||
create_load_snapshot_reader,
|
||||
create_load_snapshot_writer,
|
||||
should_use_zmq,
|
||||
zmq_reader_owner,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
|
||||
@@ -217,13 +218,16 @@ class TestFactoryFunctions(CustomTestCase):
|
||||
dp_size=1,
|
||||
load_balance_method="round_robin",
|
||||
node_rank=0,
|
||||
tokenizer_worker_num=1,
|
||||
)
|
||||
port_args = SimpleNamespace(instance_id="test_shm_factory")
|
||||
writer = create_load_snapshot_writer(
|
||||
server_args, port_args, dp_size=1, dp_rank=0
|
||||
)
|
||||
self.assertIsInstance(writer, ShmLoadSnapshotWriter)
|
||||
reader = create_load_snapshot_reader(server_args, port_args, caller="tokenizer")
|
||||
reader = create_load_snapshot_reader(
|
||||
server_args, port_args, caller="TokenizerManager"
|
||||
)
|
||||
self.assertIsInstance(reader, ShmLoadSnapshotReader)
|
||||
reader.close()
|
||||
writer.close()
|
||||
@@ -240,6 +244,7 @@ class TestFactoryFunctions(CustomTestCase):
|
||||
dp_size=1,
|
||||
load_balance_method="round_robin",
|
||||
node_rank=0,
|
||||
tokenizer_worker_num=1,
|
||||
)
|
||||
port_args = SimpleNamespace(instance_id="test_zmq_factory")
|
||||
os.environ["SGLANG_LOAD_SNAPSHOT_USE_ZMQ"] = "1"
|
||||
@@ -249,7 +254,7 @@ class TestFactoryFunctions(CustomTestCase):
|
||||
)
|
||||
self.assertIsInstance(writer, ZmqLoadSnapshotWriter)
|
||||
reader = create_load_snapshot_reader(
|
||||
server_args, port_args, caller="tokenizer"
|
||||
server_args, port_args, caller="TokenizerManager"
|
||||
)
|
||||
self.assertIsInstance(reader, ZmqShmLoadSnapshotReader)
|
||||
reader.close()
|
||||
@@ -270,6 +275,71 @@ class TestFactoryFunctions(CustomTestCase):
|
||||
self.assertFalse(should_use_zmq(args))
|
||||
|
||||
|
||||
class TestZmqReaderOwner(CustomTestCase):
|
||||
"""At most one process binds the zmq PULL socket across all callers."""
|
||||
|
||||
CALLERS = ("TokenizerManager", "MultiTokenizerRouter", "DataParallelController")
|
||||
|
||||
@staticmethod
|
||||
def _args(**overrides):
|
||||
base = dict(
|
||||
enable_dp_attention=True,
|
||||
nnodes=2,
|
||||
node_rank=0,
|
||||
dp_size=1,
|
||||
load_balance_method="round_robin",
|
||||
tokenizer_worker_num=1,
|
||||
)
|
||||
base.update(overrides)
|
||||
return SimpleNamespace(**base)
|
||||
|
||||
def _owners(self, args):
|
||||
return {c for c in self.CALLERS if zmq_reader_owner(args, c)}
|
||||
|
||||
def test_zmq_disabled_no_owner(self):
|
||||
args = self._args(enable_dp_attention=False, nnodes=1)
|
||||
self.assertEqual(self._owners(args), set())
|
||||
|
||||
def test_non_zero_node_rank_no_owner(self):
|
||||
args = self._args(node_rank=1, dp_size=4, tokenizer_worker_num=8)
|
||||
self.assertEqual(self._owners(args), set())
|
||||
|
||||
def test_tokenizer_manager_owns_when_dp1(self):
|
||||
self.assertEqual(self._owners(self._args(dp_size=1)), {"TokenizerManager"})
|
||||
|
||||
def test_multi_tokenizer_router_owns_in_multi_tokenizer_dp1(self):
|
||||
args = self._args(dp_size=1, tokenizer_worker_num=8)
|
||||
self.assertEqual(self._owners(args), {"MultiTokenizerRouter"})
|
||||
|
||||
def test_multi_tokenizer_router_owns_in_multi_tokenizer_round_robin(self):
|
||||
args = self._args(dp_size=4, tokenizer_worker_num=8)
|
||||
self.assertEqual(self._owners(args), {"MultiTokenizerRouter"})
|
||||
|
||||
def test_data_parallel_controller_owns_load_aware(self):
|
||||
for method in ("total_tokens", "total_requests"):
|
||||
args = self._args(
|
||||
dp_size=4, tokenizer_worker_num=8, load_balance_method=method
|
||||
)
|
||||
self.assertEqual(self._owners(args), {"DataParallelController"})
|
||||
|
||||
def test_tokenizer_manager_owns_dp4_round_robin(self):
|
||||
args = self._args(dp_size=4, tokenizer_worker_num=1)
|
||||
self.assertEqual(self._owners(args), {"TokenizerManager"})
|
||||
|
||||
def test_at_most_one_owner_across_configs(self):
|
||||
for dp_size in (1, 4):
|
||||
for tw in (1, 8):
|
||||
for method in ("round_robin", "total_tokens", "total_requests"):
|
||||
for node_rank in (0, 1):
|
||||
args = self._args(
|
||||
dp_size=dp_size,
|
||||
tokenizer_worker_num=tw,
|
||||
load_balance_method=method,
|
||||
node_rank=node_rank,
|
||||
)
|
||||
self.assertLessEqual(len(self._owners(args)), 1, args)
|
||||
|
||||
|
||||
class TestZmqAddr(CustomTestCase):
|
||||
def test_ipc_for_single_node(self):
|
||||
port_args = SimpleNamespace(instance_id="myinstance")
|
||||
|
||||
Reference in New Issue
Block a user