[Network] Use NetworkAddress for dist_init_method and loopback fallbacks (#20657)
Co-authored-by: hnyls2002 <lsyincs@gmail.com>
This commit is contained in:
@@ -57,6 +57,7 @@ from sglang.multimodal_gen.runtime.utils.perf_logger import (
|
||||
PerformanceLogger,
|
||||
capture_memory_snapshot,
|
||||
)
|
||||
from sglang.srt.utils.network import NetworkAddress
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -106,7 +107,9 @@ class GPUWorker:
|
||||
ring_degree=self.server_args.ring_degree,
|
||||
sp_size=self.server_args.sp_degree,
|
||||
dp_size=self.server_args.dp_size,
|
||||
distributed_init_method=f"tcp://127.0.0.1:{self.master_port}",
|
||||
distributed_init_method=NetworkAddress(
|
||||
"127.0.0.1", self.master_port
|
||||
).to_tcp(),
|
||||
dist_timeout=self.server_args.dist_timeout,
|
||||
)
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
||||
MooncakeTransferEngine,
|
||||
)
|
||||
from sglang.srt.utils.network import NetworkAddress
|
||||
|
||||
try:
|
||||
from memfabric_hybrid import TransferEngine
|
||||
@@ -47,7 +48,9 @@ class AscendTransferEngine(MooncakeTransferEngine):
|
||||
else:
|
||||
logger.error(f"Unsupported DisaggregationMode: {disaggregation_mode}")
|
||||
raise ValueError(f"Unsupported DisaggregationMode: {disaggregation_mode}")
|
||||
self.session_id = f"{self.hostname}:{self.engine.get_rpc_port()}"
|
||||
self.session_id = NetworkAddress(
|
||||
self.hostname, self.engine.get_rpc_port()
|
||||
).to_host_port_str()
|
||||
self.initialize()
|
||||
|
||||
def initialize(self) -> None:
|
||||
|
||||
@@ -238,7 +238,7 @@ class CommonKVManager(BaseKVManager):
|
||||
"""Register prefill server info to bootstrap server via HTTP POST."""
|
||||
if self.dist_init_addr:
|
||||
# Multi-node case: bootstrap server's host is dist_init_addr
|
||||
host = NetworkAddress.parse(self.dist_init_addr).host
|
||||
host = NetworkAddress.parse(self.dist_init_addr).resolved().host
|
||||
else:
|
||||
# Single-node case: bootstrap server's host is the same as http server's host
|
||||
host = self.bootstrap_host
|
||||
|
||||
@@ -29,7 +29,7 @@ from sglang.srt.disaggregation.encode_server import (
|
||||
from sglang.srt.managers.schedule_batch import Modality
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import random_uuid
|
||||
from sglang.srt.utils.network import get_zmq_socket
|
||||
from sglang.srt.utils.network import NetworkAddress, get_zmq_socket
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
SGLangEncoderServicer = sglang_encoder_pb2_grpc.SglangEncoderServicer
|
||||
@@ -212,9 +212,12 @@ async def serve_grpc_encoder(server_args: ServerArgs):
|
||||
port_args = PortArgs.init_new(server_args)
|
||||
|
||||
if server_args.dist_init_addr:
|
||||
dist_init_method = f"tcp://{server_args.dist_init_addr}"
|
||||
na = NetworkAddress.parse(server_args.dist_init_addr)
|
||||
dist_init_method = na.to_tcp()
|
||||
else:
|
||||
dist_init_method = f"tcp://127.0.0.1:{port_args.nccl_port}"
|
||||
dist_init_method = NetworkAddress(
|
||||
server_args.host or "127.0.0.1", port_args.nccl_port
|
||||
).to_tcp()
|
||||
|
||||
send_sockets: List[zmq.Socket] = []
|
||||
for rank in range(1, server_args.tp_size):
|
||||
|
||||
@@ -49,7 +49,12 @@ from sglang.srt.utils import (
|
||||
load_video,
|
||||
random_uuid,
|
||||
)
|
||||
from sglang.srt.utils.network import config_socket, get_local_ip_auto, get_zmq_socket
|
||||
from sglang.srt.utils.network import (
|
||||
NetworkAddress,
|
||||
config_socket,
|
||||
get_local_ip_auto,
|
||||
get_zmq_socket,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -1002,11 +1007,10 @@ class MMEncoder:
|
||||
mm_data.embedding = None
|
||||
|
||||
# Send ack/data
|
||||
endpoint = (
|
||||
f"tcp://{url}"
|
||||
if url is not None
|
||||
else f"tcp://{prefill_host}:{embedding_port}"
|
||||
)
|
||||
if url is not None:
|
||||
endpoint = NetworkAddress.parse(url).to_tcp()
|
||||
else:
|
||||
endpoint = NetworkAddress(prefill_host, embedding_port).to_tcp()
|
||||
logger.info(f"{endpoint = }")
|
||||
|
||||
# Serialize data
|
||||
@@ -1303,9 +1307,12 @@ def launch_server(server_args: ServerArgs):
|
||||
ipc_path_prefix = random_uuid()
|
||||
port_args = PortArgs.init_new(server_args)
|
||||
if server_args.dist_init_addr:
|
||||
dist_init_method = f"tcp://{server_args.dist_init_addr}"
|
||||
na = NetworkAddress.parse(server_args.dist_init_addr)
|
||||
dist_init_method = na.to_tcp()
|
||||
else:
|
||||
dist_init_method = f"tcp://127.0.0.1:{port_args.nccl_port}"
|
||||
dist_init_method = NetworkAddress(
|
||||
server_args.host or "127.0.0.1", port_args.nccl_port
|
||||
).to_tcp()
|
||||
for rank in range(1, server_args.tp_size):
|
||||
schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}"
|
||||
send_sockets.append(
|
||||
|
||||
@@ -288,7 +288,8 @@ class DataParallelController:
|
||||
# Determine the endpoint for inter-node communication
|
||||
if server_args.dist_init_addr is None:
|
||||
na = NetworkAddress(
|
||||
"127.0.0.1", server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA
|
||||
server_args.host or "127.0.0.1",
|
||||
server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
||||
)
|
||||
else:
|
||||
na = NetworkAddress.parse(server_args.dist_init_addr)
|
||||
|
||||
@@ -176,7 +176,7 @@ from sglang.srt.utils import (
|
||||
set_cuda_arch,
|
||||
slow_rank_detector,
|
||||
)
|
||||
from sglang.srt.utils.network import get_local_ip_auto
|
||||
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
||||
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
|
||||
from sglang.srt.utils.offloader import (
|
||||
create_offloader_from_server_args,
|
||||
@@ -672,9 +672,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
self.remote_instance_transfer_engine.initialize(
|
||||
local_ip, "P2PHANDSHAKE", "rdma", envs.MOONCAKE_DEVICE.get()
|
||||
)
|
||||
self.remote_instance_transfer_engine_session_id = (
|
||||
f"{local_ip}:{self.remote_instance_transfer_engine.get_rpc_port()}"
|
||||
)
|
||||
self.remote_instance_transfer_engine_session_id = NetworkAddress(
|
||||
local_ip, self.remote_instance_transfer_engine.get_rpc_port()
|
||||
).to_host_port_str()
|
||||
|
||||
def model_specific_adjustment(self):
|
||||
server_args = self.server_args
|
||||
@@ -774,9 +774,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
if dist_init_method_override:
|
||||
dist_init_method = dist_init_method_override
|
||||
elif self.server_args.dist_init_addr:
|
||||
dist_init_method = f"tcp://{self.server_args.dist_init_addr}"
|
||||
na = NetworkAddress.parse(self.server_args.dist_init_addr)
|
||||
dist_init_method = na.to_tcp()
|
||||
else:
|
||||
dist_init_method = f"tcp://127.0.0.1:{self.dist_port}"
|
||||
dist_init_method = NetworkAddress(
|
||||
self.server_args.host or "127.0.0.1", self.dist_port
|
||||
).to_tcp()
|
||||
set_custom_all_reduce(not self.server_args.disable_custom_all_reduce)
|
||||
set_mscclpp_all_reduce(self.server_args.enable_mscclpp)
|
||||
set_torch_symm_mem_all_reduce(self.server_args.enable_torch_symm_mem)
|
||||
@@ -956,7 +959,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
== RemoteInstanceWeightLoaderBackend.NCCL
|
||||
):
|
||||
if self.tp_rank == 0:
|
||||
instance_ip = socket.gethostbyname(socket.gethostname())
|
||||
instance_ip = NetworkAddress.resolve_host(socket.gethostname())
|
||||
t = threading.Thread(
|
||||
target=trigger_init_weights_send_group_for_remote_instance_request,
|
||||
args=(
|
||||
@@ -1234,9 +1237,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
success = False
|
||||
message = ""
|
||||
try:
|
||||
na = NetworkAddress(master_address, group_port)
|
||||
self._weights_send_group[group_name] = init_custom_process_group(
|
||||
backend=backend,
|
||||
init_method=f"tcp://{master_address}:{group_port}",
|
||||
init_method=na.to_tcp(),
|
||||
world_size=world_size,
|
||||
rank=group_rank,
|
||||
group_name=group_name,
|
||||
@@ -1244,9 +1248,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
)
|
||||
dist.barrier(group=self._weights_send_group[group_name])
|
||||
success = True
|
||||
message = (
|
||||
f"Succeeded to init group through {master_address}:{group_port} group."
|
||||
)
|
||||
message = f"Succeeded to init group through {na.to_host_port_str()} group."
|
||||
except Exception as e:
|
||||
message = f"Failed to init group: {e}."
|
||||
logger.error(message)
|
||||
@@ -1281,6 +1283,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
success = False
|
||||
na = NetworkAddress(master_address, group_port)
|
||||
message = ""
|
||||
try:
|
||||
for _, weights in self.model.named_parameters():
|
||||
@@ -1290,7 +1293,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
group=send_group,
|
||||
)
|
||||
success = True
|
||||
message = f"Succeeded to send weights through {master_address}:{group_port} {group_name}."
|
||||
message = f"Succeeded to send weights through {na.to_host_port_str()} {group_name}."
|
||||
except Exception as e:
|
||||
message = f"Failed to send weights: {e}."
|
||||
logger.error(message)
|
||||
@@ -1333,9 +1336,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
)
|
||||
|
||||
try:
|
||||
na = NetworkAddress(master_address, master_port)
|
||||
self._model_update_group[group_name] = init_custom_process_group(
|
||||
backend=backend,
|
||||
init_method=f"tcp://{master_address}:{master_port}",
|
||||
init_method=na.to_tcp(),
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
group_name=group_name,
|
||||
|
||||
@@ -416,6 +416,11 @@ class NetworkAddress:
|
||||
host: str
|
||||
port: int
|
||||
|
||||
def __post_init__(self):
|
||||
# Auto-strip IPv6 brackets so callers can pass "[::1]" or "::1"
|
||||
if self.host.startswith("[") and self.host.endswith("]"):
|
||||
object.__setattr__(self, "host", self.host[1:-1])
|
||||
|
||||
@property
|
||||
def is_ipv6(self) -> bool:
|
||||
return _is_ipv6(self.host)
|
||||
@@ -436,6 +441,26 @@ class NetworkAddress:
|
||||
"""``host:port`` string for gRPC listen address, session IDs, logs."""
|
||||
return f"{_wrap(self.host)}:{self.port}"
|
||||
|
||||
@staticmethod
|
||||
def resolve_host(host: str) -> str:
|
||||
"""Return *host* as-is if it's an IP, otherwise DNS-resolve to one."""
|
||||
try:
|
||||
ipaddress.ip_address(host)
|
||||
return host
|
||||
except ValueError:
|
||||
pass
|
||||
try:
|
||||
return socket.getaddrinfo(
|
||||
host, None, socket.AF_UNSPEC, 0, 0, socket.AI_ADDRCONFIG
|
||||
)[0][4][0]
|
||||
except socket.gaierror as e:
|
||||
raise ValueError(f"Cannot resolve host {host!r}: {e}") from e
|
||||
|
||||
def resolved(self) -> NetworkAddress:
|
||||
"""DNS-resolve hostname to IP; return self if already an IP."""
|
||||
ip = self.resolve_host(self.host)
|
||||
return self if ip == self.host else NetworkAddress(ip, self.port)
|
||||
|
||||
def to_bind_tuple(self) -> Tuple[str, int]:
|
||||
"""Raw ``(host, port)`` tuple for ``socket.bind()`` / ``socket.connect()``.
|
||||
|
||||
@@ -492,17 +517,6 @@ class NetworkAddress:
|
||||
)
|
||||
return NetworkAddress(host, _parse_port(port_str))
|
||||
|
||||
@staticmethod
|
||||
def from_parts(host: str, port: int) -> NetworkAddress:
|
||||
"""Create from separate host and port, stripping brackets if present.
|
||||
|
||||
Useful when the host may come from user input that already has
|
||||
brackets (e.g. ``[::1]``).
|
||||
"""
|
||||
if host.startswith("[") and host.endswith("]"):
|
||||
host = host[1:-1]
|
||||
return NetworkAddress(host, port)
|
||||
|
||||
def __str__(self) -> str:
|
||||
return self.to_host_port_str()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user