[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
@@ -1260,10 +1260,13 @@ class CommonKVReceiver(BaseKVReceiver):
return
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
self._register_kv_args()
# Register kv_args only once to prefill KVManager according to the info fetched
# 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:
self.bootstrap_infos = self.kv_mgr.connection_pool[bootstrap_key]
@@ -1322,6 +1325,11 @@ class CommonKVReceiver(BaseKVReceiver):
if is_ipv6:
sock.setsockopt(zmq.IPV6, 1)
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)
cls._socket_cache[endpoint] = sock
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)
return sock, lock
def _register_kv_args(self):
pass
def _register_kv_args(self) -> bool:
return True
def send_metadata(
self,
@@ -12,6 +12,7 @@ from typing import List, Optional, Tuple, Union
import numpy as np
import numpy.typing as npt
import zmq
from prometheus_client import Counter
from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll, StateType
@@ -1901,7 +1902,7 @@ class MooncakeKVReceiver(CommonKVReceiver):
self.init_time = None
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:
packed_kv_data_ptrs = b"".join(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs
@@ -1936,25 +1937,35 @@ class MooncakeKVReceiver(CommonKVReceiver):
staging_total_size_str = b""
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
with lock:
sock.send_multipart(
[
"None".encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
self.session_id.encode("ascii"),
packed_kv_data_ptrs,
packed_aux_data_ptrs,
packed_state_data_ptrs,
dst_tp_rank,
dst_attn_tp_size,
dst_kv_item_len,
packed_state_item_lens,
packed_state_dim_per_tensor,
packed_staging_base_ptr,
staging_total_size_str,
]
try:
with lock:
sock.send_multipart(
[
"None".encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
self.session_id.encode("ascii"),
packed_kv_data_ptrs,
packed_aux_data_ptrs,
packed_state_data_ptrs,
dst_tp_rank,
dst_attn_tp_size,
dst_kv_item_len,
packed_state_item_lens,
packed_state_dim_per_tensor,
packed_staging_base_ptr,
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(
self,
@@ -1983,25 +1994,33 @@ class MooncakeKVReceiver(CommonKVReceiver):
for bootstrap_info in self.bootstrap_infos:
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
is_dummy = bootstrap_info["is_dummy"]
with lock:
sock.send_multipart(
[
str(self.bootstrap_room).encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
self.session_id.encode("ascii"),
kv_indices.tobytes() if not is_dummy else b"",
str(aux_index).encode("ascii") if not is_dummy else b"",
(
pack_int_lists(state_indices, "i")
if not is_dummy and state_indices
else b""
),
str(self.required_dst_info_num).encode("ascii"),
str(decode_prefix_len or 0).encode("ascii"),
]
try:
with lock:
sock.send_multipart(
[
str(self.bootstrap_room).encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
self.session_id.encode("ascii"),
kv_indices.tobytes() if not is_dummy else b"",
str(aux_index).encode("ascii") if not is_dummy else b"",
(
pack_int_lists(state_indices, "i")
if not is_dummy and state_indices
else b""
),
str(self.required_dst_info_num).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()
def poll(self) -> KVPoll:
+53 -34
View File
@@ -1658,9 +1658,9 @@ class MoriKVReceiver(CommonKVReceiver):
return
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:
return
return False
engine_desc_blob = self.kv_mgr.engine_desc.pack()
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)
@@ -1678,25 +1678,35 @@ class MoriKVReceiver(CommonKVReceiver):
for bootstrap_info in self.bootstrap_infos:
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
with lock:
sock.send_multipart(
[
MORI_GUARD,
"None".encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
engine_desc_blob,
packed_kv_descs,
packed_aux_descs,
packed_state_descs,
gpu_id,
decode_tp_size,
decode_tp_rank,
kv_item_len,
packed_state_item_lens,
packed_state_dim_per_tensor,
]
try:
with lock:
sock.send_multipart(
[
MORI_GUARD,
"None".encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
engine_desc_blob,
packed_kv_descs,
packed_aux_descs,
packed_state_descs,
gpu_id,
decode_tp_size,
decode_tp_rank,
kv_item_len,
packed_state_item_lens,
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(
self,
@@ -1727,21 +1737,30 @@ class MoriKVReceiver(CommonKVReceiver):
state_bytes = _pack_state_indices(normalized_state)
else:
state_bytes = b""
with lock:
sock.send_multipart(
[
MORI_GUARD,
str(self.bootstrap_room).encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
self.kv_mgr.engine_desc.key.encode("ascii"),
kv_indices_bytes if not is_dummy else b"",
aux_bytes if not is_dummy else b"",
state_bytes,
str(self.required_dst_info_num).encode("ascii"),
decode_prefix_bytes,
]
try:
with lock:
sock.send_multipart(
[
MORI_GUARD,
str(self.bootstrap_room).encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
self.kv_mgr.engine_desc.key.encode("ascii"),
kv_indices_bytes if not is_dummy else b"",
aux_bytes if not is_dummy else b"",
state_bytes,
str(self.required_dst_info_num).encode("ascii"),
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()
def poll(self) -> KVPoll:
+60 -40
View File
@@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple
import numpy as np
import numpy.typing as npt
import zmq
if TYPE_CHECKING:
from sglang.srt.disaggregation.common.staging_handler import StagingTransferInfo
@@ -2637,21 +2638,30 @@ class NixlKVReceiver(CommonKVReceiver):
if not is_dummy and state_indices is not None
else b""
)
with lock:
sock.send_multipart(
[
GUARD,
str(self.bootstrap_room).encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
self.kv_mgr.agent.name.encode("ascii"),
kv_indices.tobytes() if not is_dummy else b"",
str(aux_index).encode("ascii"),
str(self.required_dst_info_num).encode("ascii"),
packed_state_indices,
str(decode_prefix_len or 0).encode("ascii"),
]
try:
with lock:
sock.send_multipart(
[
GUARD,
str(self.bootstrap_room).encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
self.kv_mgr.agent.name.encode("ascii"),
kv_indices.tobytes() if not is_dummy else b"",
str(aux_index).encode("ascii"),
str(self.required_dst_info_num).encode("ascii"),
packed_state_indices,
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.
# 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 KVPoll.WaitingForInput # type: ignore
def _register_kv_args(self):
def _register_kv_args(self) -> bool:
for bootstrap_info in self.bootstrap_infos:
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
packed_kv_data_ptrs = b"".join(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs
)
@@ -2729,31 +2738,42 @@ class NixlKVReceiver(CommonKVReceiver):
// self.kv_mgr.kv_args.kv_item_lens[0]
)
with lock:
sock.send_multipart(
[
GUARD,
"None".encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
self.kv_mgr.agent.name.encode("ascii"),
self.kv_mgr.agent.get_agent_metadata(),
packed_kv_data_ptrs,
packed_aux_data_ptrs,
packed_state_data_ptrs,
str(self.kv_mgr.kv_args.gpu_id).encode("ascii"),
str(self.kv_mgr.attn_tp_size).encode("ascii"),
str(self.kv_mgr.kv_args.engine_rank).encode("ascii"),
str(self.kv_mgr.kv_args.kv_item_lens[0]).encode("ascii"),
packed_state_item_lens,
packed_state_dim_per_tensor,
packed_staging_base_ptr,
staging_total_size_str,
str(dst_num_slots).encode("ascii"),
packed_kv_data_mem_kinds,
packed_kv_item_lens,
]
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
try:
with lock:
sock.send_multipart(
[
GUARD,
"None".encode("ascii"),
self.kv_mgr.local_ip.encode("ascii"),
str(self.kv_mgr.rank_port).encode("ascii"),
self.kv_mgr.agent.name.encode("ascii"),
self.kv_mgr.agent.get_agent_metadata(),
packed_kv_data_ptrs,
packed_aux_data_ptrs,
packed_state_data_ptrs,
str(self.kv_mgr.kv_args.gpu_id).encode("ascii"),
str(self.kv_mgr.attn_tp_size).encode("ascii"),
str(self.kv_mgr.kv_args.engine_rank).encode("ascii"),
str(self.kv_mgr.kv_args.kv_item_lens[0]).encode("ascii"),
packed_state_item_lens,
packed_state_dim_per_tensor,
packed_staging_base_ptr,
staging_total_size_str,
str(dst_num_slots).encode("ascii"),
packed_kv_data_mem_kinds,
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):
with self.kv_mgr.failure_lock:
+1
View File
@@ -392,6 +392,7 @@ class Envs:
SGLANG_DISAGGREGATION_THREAD_POOL_SIZE = EnvInt(None)
SGLANG_DISAGGREGATION_QUEUE_SIZE = EnvInt(4)
SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT = EnvInt(300)
SGLANG_DISAGGREGATION_ZMQ_SEND_TIMEOUT = EnvInt(1)
SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL = EnvFloat(5.0)
SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2)
SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300)