[EPD] fix: zmq PUSH socket reconnect-aware connection management with tcp keepalive (#27039)
This commit is contained in:
@@ -6,7 +6,6 @@ import logging
|
|||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from functools import cache
|
|
||||||
from typing import Dict, List, Optional, Set, Tuple, Union
|
from typing import Dict, List, Optional, Set, Tuple, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -144,13 +143,16 @@ class CommonKVManager(BaseKVManager):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# bind zmq socket
|
# bind zmq socket
|
||||||
context = zmq.Context()
|
self._zmq_ctx = zmq.Context()
|
||||||
self.rank_port, self.server_socket = get_zmq_socket_on_host(
|
self.rank_port, self.server_socket = get_zmq_socket_on_host(
|
||||||
context, zmq.PULL, host=self.local_ip
|
self._zmq_ctx, zmq.PULL, host=self.local_ip
|
||||||
)
|
)
|
||||||
logger.debug(f"kv manager bind to {self.local_ip}:{self.rank_port}")
|
logger.debug(f"kv manager bind to {self.local_ip}:{self.rank_port}")
|
||||||
|
|
||||||
self.request_status: Dict[int, KVPoll] = {}
|
self.request_status: Dict[int, KVPoll] = {}
|
||||||
|
self._socket_cache: Dict[str, zmq.Socket] = {}
|
||||||
|
self._monitor_cache: Dict[str, zmq.Socket] = {}
|
||||||
|
self._socket_lock = threading.Lock()
|
||||||
self.failure_records: Dict[int, str] = {}
|
self.failure_records: Dict[int, str] = {}
|
||||||
self.failure_lock = threading.Lock()
|
self.failure_lock = threading.Lock()
|
||||||
|
|
||||||
@@ -446,13 +448,44 @@ class CommonKVManager(BaseKVManager):
|
|||||||
f"Prefill instance failed to register to bootstrap server after {max_retries} retries"
|
f"Prefill instance failed to register to bootstrap server after {max_retries} retries"
|
||||||
)
|
)
|
||||||
|
|
||||||
@cache
|
|
||||||
def _connect(self, endpoint: str, is_ipv6: bool = False):
|
def _connect(self, endpoint: str, is_ipv6: bool = False):
|
||||||
socket = zmq.Context().socket(zmq.PUSH)
|
with self._socket_lock:
|
||||||
|
sock = self._socket_cache.get(endpoint)
|
||||||
|
if sock is not None:
|
||||||
|
monitor = self._monitor_cache.get(endpoint)
|
||||||
|
disconnected = False
|
||||||
|
if monitor is not None:
|
||||||
|
try:
|
||||||
|
monitor.recv_multipart(zmq.NOBLOCK)
|
||||||
|
disconnected = True
|
||||||
|
except zmq.Again:
|
||||||
|
pass
|
||||||
|
except zmq.ZMQError:
|
||||||
|
disconnected = True
|
||||||
|
if not disconnected:
|
||||||
|
return sock
|
||||||
|
sock.close(linger=0)
|
||||||
|
if monitor is not None:
|
||||||
|
monitor.close()
|
||||||
|
self._socket_cache.pop(endpoint, None)
|
||||||
|
self._monitor_cache.pop(endpoint, None)
|
||||||
|
|
||||||
|
sock = self._zmq_ctx.socket(zmq.PUSH)
|
||||||
if is_ipv6:
|
if is_ipv6:
|
||||||
socket.setsockopt(zmq.IPV6, 1)
|
sock.setsockopt(zmq.IPV6, 1)
|
||||||
socket.connect(endpoint)
|
sock.setsockopt(zmq.RECONNECT_IVL, -1)
|
||||||
return socket
|
sock.setsockopt(zmq.SNDTIMEO, 30000)
|
||||||
|
sock.setsockopt(zmq.LINGER, 0)
|
||||||
|
sock.setsockopt(zmq.TCP_KEEPALIVE, 1)
|
||||||
|
sock.setsockopt(zmq.TCP_KEEPALIVE_IDLE, 30)
|
||||||
|
sock.setsockopt(zmq.TCP_KEEPALIVE_INTVL, 5)
|
||||||
|
sock.setsockopt(zmq.TCP_KEEPALIVE_CNT, 3)
|
||||||
|
sock.connect(endpoint)
|
||||||
|
self._socket_cache[endpoint] = sock
|
||||||
|
self._monitor_cache[endpoint] = sock.get_monitor_socket(
|
||||||
|
zmq.EVENT_DISCONNECTED
|
||||||
|
)
|
||||||
|
return sock
|
||||||
|
|
||||||
def get_mha_kv_ptrs_with_pp(
|
def get_mha_kv_ptrs_with_pp(
|
||||||
self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int]
|
self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int]
|
||||||
|
|||||||
@@ -1639,7 +1639,7 @@ class MMEncoder:
|
|||||||
else:
|
else:
|
||||||
sock.send_multipart([serialized_data], copy=False)
|
sock.send_multipart([serialized_data], copy=False)
|
||||||
finally:
|
finally:
|
||||||
sock.close()
|
sock.close(linger=5000)
|
||||||
|
|
||||||
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)
|
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user