[PD] Prevent decode scheduler from blocking on ZMQ sends to a stalled prefill peer (#31144)

Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
SovietPower
2026-07-25 12:53:41 +08:00
committed by GitHub
co-authored by Shangming Cai
parent ebcb74abd4
commit 6a046fad09
6 changed files with 188 additions and 116 deletions
@@ -1428,6 +1428,11 @@ SGLang supports various environment variables that can be used to configure its
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Consecutive heartbeat failures tolerated before a peer is considered dead.</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Consecutive heartbeat failures tolerated before a peer is considered dead.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>2</code></td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>2</code></td>
</tr> </tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_DISAGGREGATION_ZMQ_SEND_TIMEOUT</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Send timeout (seconds) for decode's ZMQ sockets to prefill peers. Raise if healthy sends exceed the default.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>1</code></td>
</tr>
<tr> <tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL</code></td> <td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Interval (seconds) for cleaning up stale bootstrap entries.</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Interval (seconds) for cleaning up stale bootstrap entries.</td>
@@ -1260,10 +1260,13 @@ class CommonKVReceiver(BaseKVReceiver):
return return
self.bootstrap_infos = bootstrap_infos self.bootstrap_infos = bootstrap_infos
self.kv_mgr.connection_pool[bootstrap_key] = self.bootstrap_infos
# Register kv_args only once to prefill KVManager according to the info fetched from the bootstrap server # Register kv_args only once to prefill KVManager according to the info fetched
self._register_kv_args() # from the bootstrap server. Do this before caching in connection_pool so a failed
# registration does not leave a stale entry that later requests would reuse.
if not self._register_kv_args():
return
self.kv_mgr.connection_pool[bootstrap_key] = self.bootstrap_infos
else: else:
self.bootstrap_infos = self.kv_mgr.connection_pool[bootstrap_key] self.bootstrap_infos = self.kv_mgr.connection_pool[bootstrap_key]
@@ -1322,6 +1325,11 @@ class CommonKVReceiver(BaseKVReceiver):
if is_ipv6: if is_ipv6:
sock.setsockopt(zmq.IPV6, 1) sock.setsockopt(zmq.IPV6, 1)
sock.setsockopt(zmq.LINGER, 0) sock.setsockopt(zmq.LINGER, 0)
# Bound send so a dead peer cannot block the scheduler forever.
sock.setsockopt(
zmq.SNDTIMEO,
envs.SGLANG_DISAGGREGATION_ZMQ_SEND_TIMEOUT.get() * 1000,
)
sock.connect(endpoint) sock.connect(endpoint)
cls._socket_cache[endpoint] = sock cls._socket_cache[endpoint] = sock
cls._socket_locks[endpoint] = threading.Lock() cls._socket_locks[endpoint] = threading.Lock()
@@ -1348,8 +1356,8 @@ class CommonKVReceiver(BaseKVReceiver):
sock, lock = cls._connect(na.to_tcp(), is_ipv6=na.is_ipv6) sock, lock = cls._connect(na.to_tcp(), is_ipv6=na.is_ipv6)
return sock, lock return sock, lock
def _register_kv_args(self): def _register_kv_args(self) -> bool:
pass return True
def send_metadata( def send_metadata(
self, self,
@@ -12,6 +12,7 @@ from typing import List, Optional, Tuple, Union
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
import zmq
from prometheus_client import Counter from prometheus_client import Counter
from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll, StateType from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll, StateType
@@ -1901,7 +1902,7 @@ class MooncakeKVReceiver(CommonKVReceiver):
self.init_time = None self.init_time = None
super().__init__(mgr, bootstrap_addr, bootstrap_room) super().__init__(mgr, bootstrap_addr, bootstrap_room)
def _register_kv_args(self): def _register_kv_args(self) -> bool:
for bootstrap_info in self.bootstrap_infos: for bootstrap_info in self.bootstrap_infos:
packed_kv_data_ptrs = b"".join( packed_kv_data_ptrs = b"".join(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs
@@ -1936,6 +1937,7 @@ class MooncakeKVReceiver(CommonKVReceiver):
staging_total_size_str = b"" staging_total_size_str = b""
sock, lock = self._connect_to_bootstrap_server(bootstrap_info) sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
try:
with lock: with lock:
sock.send_multipart( sock.send_multipart(
[ [
@@ -1955,6 +1957,15 @@ class MooncakeKVReceiver(CommonKVReceiver):
staging_total_size_str, staging_total_size_str,
] ]
) )
except zmq.ZMQError:
self.kv_mgr.record_failure(
self.bootstrap_room,
f"_register_kv_args to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
self.conclude_state = KVPoll.Failed
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
return False
return True
def send_metadata( def send_metadata(
self, self,
@@ -1983,7 +1994,7 @@ class MooncakeKVReceiver(CommonKVReceiver):
for bootstrap_info in self.bootstrap_infos: for bootstrap_info in self.bootstrap_infos:
sock, lock = self._connect_to_bootstrap_server(bootstrap_info) sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
is_dummy = bootstrap_info["is_dummy"] is_dummy = bootstrap_info["is_dummy"]
try:
with lock: with lock:
sock.send_multipart( sock.send_multipart(
[ [
@@ -2002,6 +2013,14 @@ class MooncakeKVReceiver(CommonKVReceiver):
str(decode_prefix_len or 0).encode("ascii"), str(decode_prefix_len or 0).encode("ascii"),
] ]
) )
except zmq.ZMQError:
self.kv_mgr.record_failure(
self.bootstrap_room,
f"send_metadata to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
self.conclude_state = KVPoll.Failed
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
return
self.init_time = time.time() self.init_time = time.time()
def poll(self) -> KVPoll: def poll(self) -> KVPoll:
+21 -2
View File
@@ -1658,9 +1658,9 @@ class MoriKVReceiver(CommonKVReceiver):
return return
self.kv_mgr.room_to_bootstrap_addr[self.bootstrap_room] = self.bootstrap_addr self.kv_mgr.room_to_bootstrap_addr[self.bootstrap_room] = self.bootstrap_addr
def _register_kv_args(self): def _register_kv_args(self) -> bool:
if self.bootstrap_infos is None: if self.bootstrap_infos is None:
return return False
engine_desc_blob = self.kv_mgr.engine_desc.pack() engine_desc_blob = self.kv_mgr.engine_desc.pack()
packed_kv_descs = _pack_mem_desc_list(self.kv_mgr.kv_mem_descs) packed_kv_descs = _pack_mem_desc_list(self.kv_mgr.kv_mem_descs)
packed_aux_descs = _pack_mem_desc_list(self.kv_mgr.aux_mem_descs) packed_aux_descs = _pack_mem_desc_list(self.kv_mgr.aux_mem_descs)
@@ -1678,6 +1678,7 @@ class MoriKVReceiver(CommonKVReceiver):
for bootstrap_info in self.bootstrap_infos: for bootstrap_info in self.bootstrap_infos:
sock, lock = self._connect_to_bootstrap_server(bootstrap_info) sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
try:
with lock: with lock:
sock.send_multipart( sock.send_multipart(
[ [
@@ -1697,6 +1698,15 @@ class MoriKVReceiver(CommonKVReceiver):
packed_state_dim_per_tensor, packed_state_dim_per_tensor,
] ]
) )
except zmq.ZMQError:
self.kv_mgr.record_failure(
self.bootstrap_room,
f"_register_kv_args to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
self.conclude_state = KVPoll.Failed
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
return False
return True
def send_metadata( def send_metadata(
self, self,
@@ -1727,6 +1737,7 @@ class MoriKVReceiver(CommonKVReceiver):
state_bytes = _pack_state_indices(normalized_state) state_bytes = _pack_state_indices(normalized_state)
else: else:
state_bytes = b"" state_bytes = b""
try:
with lock: with lock:
sock.send_multipart( sock.send_multipart(
[ [
@@ -1742,6 +1753,14 @@ class MoriKVReceiver(CommonKVReceiver):
decode_prefix_bytes, decode_prefix_bytes,
] ]
) )
except zmq.ZMQError:
self.kv_mgr.record_failure(
self.bootstrap_room,
f"send_metadata to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
self.conclude_state = KVPoll.Failed
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
return
self.init_time = time.time() self.init_time = time.time()
def poll(self) -> KVPoll: def poll(self) -> KVPoll:
+22 -2
View File
@@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
import zmq
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.disaggregation.common.staging_handler import StagingTransferInfo from sglang.srt.disaggregation.common.staging_handler import StagingTransferInfo
@@ -2637,6 +2638,7 @@ class NixlKVReceiver(CommonKVReceiver):
if not is_dummy and state_indices is not None if not is_dummy and state_indices is not None
else b"" else b""
) )
try:
with lock: with lock:
sock.send_multipart( sock.send_multipart(
[ [
@@ -2652,6 +2654,14 @@ class NixlKVReceiver(CommonKVReceiver):
str(decode_prefix_len or 0).encode("ascii"), str(decode_prefix_len or 0).encode("ascii"),
] ]
) )
except zmq.ZMQError:
self.kv_mgr.record_failure(
self.bootstrap_room,
f"send_metadata to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
self.conclude_state = KVPoll.Failed
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
return
# Mark that we expect state data if state_indices was provided. # Mark that we expect state data if state_indices was provided.
# Match the prefill-side truthy check: an empty list means the # Match the prefill-side truthy check: an empty list means the
@@ -2687,9 +2697,8 @@ class NixlKVReceiver(CommonKVReceiver):
return self.conclude_state # type: ignore return self.conclude_state # type: ignore
return KVPoll.WaitingForInput # type: ignore return KVPoll.WaitingForInput # type: ignore
def _register_kv_args(self): def _register_kv_args(self) -> bool:
for bootstrap_info in self.bootstrap_infos: for bootstrap_info in self.bootstrap_infos:
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
packed_kv_data_ptrs = b"".join( packed_kv_data_ptrs = b"".join(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs
) )
@@ -2729,6 +2738,8 @@ class NixlKVReceiver(CommonKVReceiver):
// self.kv_mgr.kv_args.kv_item_lens[0] // self.kv_mgr.kv_args.kv_item_lens[0]
) )
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
try:
with lock: with lock:
sock.send_multipart( sock.send_multipart(
[ [
@@ -2754,6 +2765,15 @@ class NixlKVReceiver(CommonKVReceiver):
packed_kv_item_lens, packed_kv_item_lens,
] ]
) )
except zmq.ZMQError:
self.kv_mgr.record_failure(
self.bootstrap_room,
f"_register_kv_args to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
self.conclude_state = KVPoll.Failed
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
return False
return True
def failure_exception(self): def failure_exception(self):
with self.kv_mgr.failure_lock: with self.kv_mgr.failure_lock:
+1
View File
@@ -392,6 +392,7 @@ class Envs:
SGLANG_DISAGGREGATION_THREAD_POOL_SIZE = EnvInt(None) SGLANG_DISAGGREGATION_THREAD_POOL_SIZE = EnvInt(None)
SGLANG_DISAGGREGATION_QUEUE_SIZE = EnvInt(4) SGLANG_DISAGGREGATION_QUEUE_SIZE = EnvInt(4)
SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT = EnvInt(300) SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT = EnvInt(300)
SGLANG_DISAGGREGATION_ZMQ_SEND_TIMEOUT = EnvInt(1)
SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL = EnvFloat(5.0) SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL = EnvFloat(5.0)
SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2) SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2)
SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300) SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300)