fix(tcp-port): replace bind_server_socket to get_zmq_socket(Port conflict) (#11961)
Co-authored-by: wangchao <wcsjtu@163.com>
This commit is contained in:
@@ -32,8 +32,8 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
format_tcp_address,
|
format_tcp_address,
|
||||||
get_free_port,
|
|
||||||
get_local_ip_auto,
|
get_local_ip_auto,
|
||||||
|
get_zmq_socket_on_host,
|
||||||
is_valid_ipv6_address,
|
is_valid_ipv6_address,
|
||||||
maybe_wrap_ipv6_address,
|
maybe_wrap_ipv6_address,
|
||||||
)
|
)
|
||||||
@@ -68,11 +68,16 @@ class CommonKVManager(BaseKVManager):
|
|||||||
)
|
)
|
||||||
self.pp_size = server_args.pp_size
|
self.pp_size = server_args.pp_size
|
||||||
self.pp_rank = self.kv_args.pp_rank
|
self.pp_rank = self.kv_args.pp_rank
|
||||||
self.rank_port = get_free_port()
|
|
||||||
self.local_ip = get_local_ip_auto()
|
self.local_ip = get_local_ip_auto()
|
||||||
self.server_socket = zmq.Context().socket(zmq.PULL)
|
|
||||||
if is_valid_ipv6_address(self.local_ip):
|
# bind zmq socket
|
||||||
self.server_socket.setsockopt(zmq.IPV6, 1)
|
context = zmq.Context()
|
||||||
|
zmq_bind_host = maybe_wrap_ipv6_address(self.local_ip)
|
||||||
|
self.rank_port, self.server_socket = get_zmq_socket_on_host(
|
||||||
|
context, zmq.PULL, host=zmq_bind_host
|
||||||
|
)
|
||||||
|
logger.debug(f"kv manager bind to {zmq_bind_host}:{self.rank_port}")
|
||||||
|
|
||||||
self.request_status: Dict[int, KVPoll] = {}
|
self.request_status: Dict[int, KVPoll] = {}
|
||||||
|
|
||||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||||
@@ -92,9 +97,6 @@ class CommonKVManager(BaseKVManager):
|
|||||||
f"Unsupported DisaggregationMode: {self.disaggregation_mode}"
|
f"Unsupported DisaggregationMode: {self.disaggregation_mode}"
|
||||||
)
|
)
|
||||||
|
|
||||||
def _bind_server_socket(self):
|
|
||||||
self.server_socket.bind(format_tcp_address(self.local_ip, self.rank_port))
|
|
||||||
|
|
||||||
def _register_to_bootstrap(self):
|
def _register_to_bootstrap(self):
|
||||||
"""Register KVSender to bootstrap server via HTTP POST."""
|
"""Register KVSender to bootstrap server via HTTP POST."""
|
||||||
if self.dist_init_addr:
|
if self.dist_init_addr:
|
||||||
@@ -177,7 +179,6 @@ class CommonKVManager(BaseKVManager):
|
|||||||
|
|
||||||
|
|
||||||
class CommonKVSender(BaseKVSender):
|
class CommonKVSender(BaseKVSender):
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
mgr: BaseKVManager,
|
mgr: BaseKVManager,
|
||||||
|
|||||||
@@ -238,6 +238,9 @@ class ZmqEventPublisher(EventPublisher):
|
|||||||
or self._endpoint.startswith("ipc://")
|
or self._endpoint.startswith("ipc://")
|
||||||
or self._endpoint.startswith("inproc://")
|
or self._endpoint.startswith("inproc://")
|
||||||
):
|
):
|
||||||
|
logger.debug(
|
||||||
|
f"ZmqEventPublisher socket publisher_endpoint bind to {self._endpoint}"
|
||||||
|
)
|
||||||
self._pub.bind(self._endpoint)
|
self._pub.bind(self._endpoint)
|
||||||
else:
|
else:
|
||||||
self._pub.connect(self._endpoint)
|
self._pub.connect(self._endpoint)
|
||||||
@@ -248,6 +251,9 @@ class ZmqEventPublisher(EventPublisher):
|
|||||||
# 3) works in our non‑blocking poll loop alongside PUB
|
# 3) works in our non‑blocking poll loop alongside PUB
|
||||||
if self._replay_endpoint is not None:
|
if self._replay_endpoint is not None:
|
||||||
self._replay = self._ctx.socket(zmq.ROUTER)
|
self._replay = self._ctx.socket(zmq.ROUTER)
|
||||||
|
logger.debug(
|
||||||
|
f"ZmqEventPublisher socket replay_endpoint bind to {self._replay_endpoint}"
|
||||||
|
)
|
||||||
self._replay.bind(self._replay_endpoint)
|
self._replay.bind(self._replay_endpoint)
|
||||||
|
|
||||||
def _publisher_thread(self) -> None:
|
def _publisher_thread(self) -> None:
|
||||||
|
|||||||
@@ -849,8 +849,6 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def start_prefill_thread(self):
|
def start_prefill_thread(self):
|
||||||
self._bind_server_socket()
|
|
||||||
|
|
||||||
def bootstrap_thread():
|
def bootstrap_thread():
|
||||||
"""This thread recvs pre-alloc notification from the decode engine"""
|
"""This thread recvs pre-alloc notification from the decode engine"""
|
||||||
# KVPoll.Bootstrapping -> KVPoll.WaitingForInput
|
# KVPoll.Bootstrapping -> KVPoll.WaitingForInput
|
||||||
@@ -887,8 +885,6 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
threading.Thread(target=bootstrap_thread).start()
|
threading.Thread(target=bootstrap_thread).start()
|
||||||
|
|
||||||
def start_decode_thread(self):
|
def start_decode_thread(self):
|
||||||
self._bind_server_socket()
|
|
||||||
|
|
||||||
def decode_thread():
|
def decode_thread():
|
||||||
while True:
|
while True:
|
||||||
msg = self.server_socket.recv_multipart()
|
msg = self.server_socket.recv_multipart()
|
||||||
|
|||||||
@@ -83,8 +83,8 @@ class KVArgsRegisterInfo:
|
|||||||
dst_port=int(msg[2].decode("ascii")),
|
dst_port=int(msg[2].decode("ascii")),
|
||||||
agent_name=msg[3].decode("ascii"),
|
agent_name=msg[3].decode("ascii"),
|
||||||
agent_metadata=msg[4],
|
agent_metadata=msg[4],
|
||||||
dst_kv_ptrs=list(struct.unpack(f"{len(msg[5])//8}Q", msg[5])),
|
dst_kv_ptrs=list(struct.unpack(f"{len(msg[5]) // 8}Q", msg[5])),
|
||||||
dst_aux_ptrs=list(struct.unpack(f"{len(msg[6])//8}Q", msg[6])),
|
dst_aux_ptrs=list(struct.unpack(f"{len(msg[6]) // 8}Q", msg[6])),
|
||||||
gpu_id=int(msg[7].decode("ascii")),
|
gpu_id=int(msg[7].decode("ascii")),
|
||||||
decode_tp_size=int(msg[8].decode("ascii")),
|
decode_tp_size=int(msg[8].decode("ascii")),
|
||||||
decode_tp_rank=int(msg[9].decode("ascii")),
|
decode_tp_rank=int(msg[9].decode("ascii")),
|
||||||
@@ -647,8 +647,6 @@ class NixlKVManager(CommonKVManager):
|
|||||||
return self.transfer_statuses[room].is_done()
|
return self.transfer_statuses[room].is_done()
|
||||||
|
|
||||||
def _start_bootstrap_thread(self):
|
def _start_bootstrap_thread(self):
|
||||||
self._bind_server_socket()
|
|
||||||
|
|
||||||
def bootstrap_thread():
|
def bootstrap_thread():
|
||||||
"""This thread recvs transfer info from the decode engine"""
|
"""This thread recvs transfer info from the decode engine"""
|
||||||
while True:
|
while True:
|
||||||
@@ -687,7 +685,6 @@ class NixlKVManager(CommonKVManager):
|
|||||||
|
|
||||||
|
|
||||||
class NixlKVSender(CommonKVSender):
|
class NixlKVSender(CommonKVSender):
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
mgr: NixlKVManager,
|
mgr: NixlKVManager,
|
||||||
|
|||||||
@@ -215,7 +215,6 @@ class MessageQueue:
|
|||||||
socket_addr = f"tcp://127.0.0.1:{local_subscribe_port}"
|
socket_addr = f"tcp://127.0.0.1:{local_subscribe_port}"
|
||||||
logger.debug("Binding to %s", socket_addr)
|
logger.debug("Binding to %s", socket_addr)
|
||||||
self.local_socket.bind(socket_addr)
|
self.local_socket.bind(socket_addr)
|
||||||
|
|
||||||
self.current_idx = 0
|
self.current_idx = 0
|
||||||
|
|
||||||
else:
|
else:
|
||||||
@@ -232,9 +231,9 @@ class MessageQueue:
|
|||||||
remote_subscribe_port = get_open_port()
|
remote_subscribe_port = get_open_port()
|
||||||
if is_valid_ipv6_address(connect_ip):
|
if is_valid_ipv6_address(connect_ip):
|
||||||
self.remote_socket.setsockopt(IPV6, 1)
|
self.remote_socket.setsockopt(IPV6, 1)
|
||||||
self.remote_socket.bind(
|
address = format_tcp_address(connect_ip, remote_subscribe_port)
|
||||||
format_tcp_address(connect_ip, remote_subscribe_port)
|
logger.debug(f"class MessageQueue: Binding remote socket to {address=}")
|
||||||
)
|
self.remote_socket.bind(address)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
remote_subscribe_port = None
|
remote_subscribe_port = None
|
||||||
|
|||||||
@@ -1359,6 +1359,29 @@ def get_zmq_socket(
|
|||||||
return socket
|
return socket
|
||||||
|
|
||||||
|
|
||||||
|
def get_zmq_socket_on_host(
|
||||||
|
context: zmq.Context,
|
||||||
|
socket_type: zmq.SocketType,
|
||||||
|
host: Optional[str] = None,
|
||||||
|
) -> Tuple[int, zmq.Socket]:
|
||||||
|
"""Create and configure a ZeroMQ socket.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
context: ZeroMQ context to create the socket from.
|
||||||
|
socket_type: Type of ZeroMQ socket to create.
|
||||||
|
host: Optional host to bind/connect to, without "tcp://" prefix. If None, binds to "tcp://*".
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (port, socket) where port is the randomly assigned TCP port.
|
||||||
|
"""
|
||||||
|
socket = context.socket(socket_type)
|
||||||
|
# Bind to random TCP port
|
||||||
|
config_socket(socket, socket_type)
|
||||||
|
bind_host = f"tcp://{host}" if host else "tcp://*"
|
||||||
|
port = socket.bind_to_random_port(bind_host)
|
||||||
|
return port, socket
|
||||||
|
|
||||||
|
|
||||||
def config_socket(socket, socket_type: zmq.SocketType):
|
def config_socket(socket, socket_type: zmq.SocketType):
|
||||||
mem = psutil.virtual_memory()
|
mem = psutil.virtual_memory()
|
||||||
total_mem = mem.total / 1024**3
|
total_mem = mem.total / 1024**3
|
||||||
|
|||||||
Reference in New Issue
Block a user