[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:
co-authored by
Shangming Cai
parent
ebcb74abd4
commit
6a046fad09
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user