[Network] Use NetworkAddress for dist_init_method and loopback fallbacks (#20657)

Co-authored-by: hnyls2002 <lsyincs@gmail.com>
This commit is contained in:
psaab
2026-03-16 19:59:49 -07:00
committed by GitHub
parent 71a54c1c42
commit 9f56b471aa
9 changed files with 79 additions and 44 deletions

View File

@@ -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,
)

View File

@@ -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:

View File

@@ -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

View File

@@ -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):

View File

@@ -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(

View File

@@ -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)

View File

@@ -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,

View File

@@ -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()