[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,
|
PerformanceLogger,
|
||||||
capture_memory_snapshot,
|
capture_memory_snapshot,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -106,7 +107,9 @@ class GPUWorker:
|
|||||||
ring_degree=self.server_args.ring_degree,
|
ring_degree=self.server_args.ring_degree,
|
||||||
sp_size=self.server_args.sp_degree,
|
sp_size=self.server_args.sp_degree,
|
||||||
dp_size=self.server_args.dp_size,
|
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,
|
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 (
|
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
||||||
MooncakeTransferEngine,
|
MooncakeTransferEngine,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from memfabric_hybrid import TransferEngine
|
from memfabric_hybrid import TransferEngine
|
||||||
@@ -47,7 +48,9 @@ class AscendTransferEngine(MooncakeTransferEngine):
|
|||||||
else:
|
else:
|
||||||
logger.error(f"Unsupported DisaggregationMode: {disaggregation_mode}")
|
logger.error(f"Unsupported DisaggregationMode: {disaggregation_mode}")
|
||||||
raise ValueError(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()
|
self.initialize()
|
||||||
|
|
||||||
def initialize(self) -> None:
|
def initialize(self) -> None:
|
||||||
|
|||||||
@@ -238,7 +238,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
"""Register prefill server info to bootstrap server via HTTP POST."""
|
"""Register prefill server info to bootstrap server via HTTP POST."""
|
||||||
if self.dist_init_addr:
|
if self.dist_init_addr:
|
||||||
# Multi-node case: bootstrap server's host is 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:
|
else:
|
||||||
# Single-node case: bootstrap server's host is the same as http server's host
|
# Single-node case: bootstrap server's host is the same as http server's host
|
||||||
host = self.bootstrap_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.managers.schedule_batch import Modality
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
from sglang.srt.utils import random_uuid
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
SGLangEncoderServicer = sglang_encoder_pb2_grpc.SglangEncoderServicer
|
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)
|
port_args = PortArgs.init_new(server_args)
|
||||||
|
|
||||||
if server_args.dist_init_addr:
|
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:
|
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] = []
|
send_sockets: List[zmq.Socket] = []
|
||||||
for rank in range(1, server_args.tp_size):
|
for rank in range(1, server_args.tp_size):
|
||||||
|
|||||||
@@ -49,7 +49,12 @@ from sglang.srt.utils import (
|
|||||||
load_video,
|
load_video,
|
||||||
random_uuid,
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -1002,11 +1007,10 @@ class MMEncoder:
|
|||||||
mm_data.embedding = None
|
mm_data.embedding = None
|
||||||
|
|
||||||
# Send ack/data
|
# Send ack/data
|
||||||
endpoint = (
|
if url is not None:
|
||||||
f"tcp://{url}"
|
endpoint = NetworkAddress.parse(url).to_tcp()
|
||||||
if url is not None
|
else:
|
||||||
else f"tcp://{prefill_host}:{embedding_port}"
|
endpoint = NetworkAddress(prefill_host, embedding_port).to_tcp()
|
||||||
)
|
|
||||||
logger.info(f"{endpoint = }")
|
logger.info(f"{endpoint = }")
|
||||||
|
|
||||||
# Serialize data
|
# Serialize data
|
||||||
@@ -1303,9 +1307,12 @@ def launch_server(server_args: ServerArgs):
|
|||||||
ipc_path_prefix = random_uuid()
|
ipc_path_prefix = random_uuid()
|
||||||
port_args = PortArgs.init_new(server_args)
|
port_args = PortArgs.init_new(server_args)
|
||||||
if server_args.dist_init_addr:
|
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:
|
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):
|
for rank in range(1, server_args.tp_size):
|
||||||
schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}"
|
schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}"
|
||||||
send_sockets.append(
|
send_sockets.append(
|
||||||
|
|||||||
@@ -288,7 +288,8 @@ class DataParallelController:
|
|||||||
# Determine the endpoint for inter-node communication
|
# Determine the endpoint for inter-node communication
|
||||||
if server_args.dist_init_addr is None:
|
if server_args.dist_init_addr is None:
|
||||||
na = NetworkAddress(
|
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:
|
else:
|
||||||
na = NetworkAddress.parse(server_args.dist_init_addr)
|
na = NetworkAddress.parse(server_args.dist_init_addr)
|
||||||
|
|||||||
@@ -176,7 +176,7 @@ from sglang.srt.utils import (
|
|||||||
set_cuda_arch,
|
set_cuda_arch,
|
||||||
slow_rank_detector,
|
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.nvtx_pytorch_hooks import PytHooks
|
||||||
from sglang.srt.utils.offloader import (
|
from sglang.srt.utils.offloader import (
|
||||||
create_offloader_from_server_args,
|
create_offloader_from_server_args,
|
||||||
@@ -672,9 +672,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.remote_instance_transfer_engine.initialize(
|
self.remote_instance_transfer_engine.initialize(
|
||||||
local_ip, "P2PHANDSHAKE", "rdma", envs.MOONCAKE_DEVICE.get()
|
local_ip, "P2PHANDSHAKE", "rdma", envs.MOONCAKE_DEVICE.get()
|
||||||
)
|
)
|
||||||
self.remote_instance_transfer_engine_session_id = (
|
self.remote_instance_transfer_engine_session_id = NetworkAddress(
|
||||||
f"{local_ip}:{self.remote_instance_transfer_engine.get_rpc_port()}"
|
local_ip, self.remote_instance_transfer_engine.get_rpc_port()
|
||||||
)
|
).to_host_port_str()
|
||||||
|
|
||||||
def model_specific_adjustment(self):
|
def model_specific_adjustment(self):
|
||||||
server_args = self.server_args
|
server_args = self.server_args
|
||||||
@@ -774,9 +774,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
if dist_init_method_override:
|
if dist_init_method_override:
|
||||||
dist_init_method = dist_init_method_override
|
dist_init_method = dist_init_method_override
|
||||||
elif self.server_args.dist_init_addr:
|
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:
|
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_custom_all_reduce(not self.server_args.disable_custom_all_reduce)
|
||||||
set_mscclpp_all_reduce(self.server_args.enable_mscclpp)
|
set_mscclpp_all_reduce(self.server_args.enable_mscclpp)
|
||||||
set_torch_symm_mem_all_reduce(self.server_args.enable_torch_symm_mem)
|
set_torch_symm_mem_all_reduce(self.server_args.enable_torch_symm_mem)
|
||||||
@@ -956,7 +959,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
== RemoteInstanceWeightLoaderBackend.NCCL
|
== RemoteInstanceWeightLoaderBackend.NCCL
|
||||||
):
|
):
|
||||||
if self.tp_rank == 0:
|
if self.tp_rank == 0:
|
||||||
instance_ip = socket.gethostbyname(socket.gethostname())
|
instance_ip = NetworkAddress.resolve_host(socket.gethostname())
|
||||||
t = threading.Thread(
|
t = threading.Thread(
|
||||||
target=trigger_init_weights_send_group_for_remote_instance_request,
|
target=trigger_init_weights_send_group_for_remote_instance_request,
|
||||||
args=(
|
args=(
|
||||||
@@ -1234,9 +1237,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
success = False
|
success = False
|
||||||
message = ""
|
message = ""
|
||||||
try:
|
try:
|
||||||
|
na = NetworkAddress(master_address, group_port)
|
||||||
self._weights_send_group[group_name] = init_custom_process_group(
|
self._weights_send_group[group_name] = init_custom_process_group(
|
||||||
backend=backend,
|
backend=backend,
|
||||||
init_method=f"tcp://{master_address}:{group_port}",
|
init_method=na.to_tcp(),
|
||||||
world_size=world_size,
|
world_size=world_size,
|
||||||
rank=group_rank,
|
rank=group_rank,
|
||||||
group_name=group_name,
|
group_name=group_name,
|
||||||
@@ -1244,9 +1248,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
dist.barrier(group=self._weights_send_group[group_name])
|
dist.barrier(group=self._weights_send_group[group_name])
|
||||||
success = True
|
success = True
|
||||||
message = (
|
message = f"Succeeded to init group through {na.to_host_port_str()} group."
|
||||||
f"Succeeded to init group through {master_address}:{group_port} group."
|
|
||||||
)
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
message = f"Failed to init group: {e}."
|
message = f"Failed to init group: {e}."
|
||||||
logger.error(message)
|
logger.error(message)
|
||||||
@@ -1281,6 +1283,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
|
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
success = False
|
success = False
|
||||||
|
na = NetworkAddress(master_address, group_port)
|
||||||
message = ""
|
message = ""
|
||||||
try:
|
try:
|
||||||
for _, weights in self.model.named_parameters():
|
for _, weights in self.model.named_parameters():
|
||||||
@@ -1290,7 +1293,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
group=send_group,
|
group=send_group,
|
||||||
)
|
)
|
||||||
success = True
|
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:
|
except Exception as e:
|
||||||
message = f"Failed to send weights: {e}."
|
message = f"Failed to send weights: {e}."
|
||||||
logger.error(message)
|
logger.error(message)
|
||||||
@@ -1333,9 +1336,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
na = NetworkAddress(master_address, master_port)
|
||||||
self._model_update_group[group_name] = init_custom_process_group(
|
self._model_update_group[group_name] = init_custom_process_group(
|
||||||
backend=backend,
|
backend=backend,
|
||||||
init_method=f"tcp://{master_address}:{master_port}",
|
init_method=na.to_tcp(),
|
||||||
world_size=world_size,
|
world_size=world_size,
|
||||||
rank=rank,
|
rank=rank,
|
||||||
group_name=group_name,
|
group_name=group_name,
|
||||||
|
|||||||
@@ -416,6 +416,11 @@ class NetworkAddress:
|
|||||||
host: str
|
host: str
|
||||||
port: int
|
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
|
@property
|
||||||
def is_ipv6(self) -> bool:
|
def is_ipv6(self) -> bool:
|
||||||
return _is_ipv6(self.host)
|
return _is_ipv6(self.host)
|
||||||
@@ -436,6 +441,26 @@ class NetworkAddress:
|
|||||||
"""``host:port`` string for gRPC listen address, session IDs, logs."""
|
"""``host:port`` string for gRPC listen address, session IDs, logs."""
|
||||||
return f"{_wrap(self.host)}:{self.port}"
|
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]:
|
def to_bind_tuple(self) -> Tuple[str, int]:
|
||||||
"""Raw ``(host, port)`` tuple for ``socket.bind()`` / ``socket.connect()``.
|
"""Raw ``(host, port)`` tuple for ``socket.bind()`` / ``socket.connect()``.
|
||||||
|
|
||||||
@@ -492,17 +517,6 @@ class NetworkAddress:
|
|||||||
)
|
)
|
||||||
return NetworkAddress(host, _parse_port(port_str))
|
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:
|
def __str__(self) -> str:
|
||||||
return self.to_host_port_str()
|
return self.to_host_port_str()
|
||||||
|
|
||||||
|
|||||||
@@ -179,23 +179,23 @@ class TestNetworkAddressParseErrors(unittest.TestCase):
|
|||||||
NetworkAddress.parse(":8000")
|
NetworkAddress.parse(":8000")
|
||||||
|
|
||||||
|
|
||||||
class TestNetworkAddressFromParts(unittest.TestCase):
|
class TestNetworkAddressBracketStripping(unittest.TestCase):
|
||||||
def test_strip_brackets(self):
|
def test_strip_brackets(self):
|
||||||
na = NetworkAddress.from_parts("[::1]", 8000)
|
na = NetworkAddress("[::1]", 8000)
|
||||||
self.assertEqual(na.host, "::1")
|
self.assertEqual(na.host, "::1")
|
||||||
self.assertTrue(na.is_ipv6)
|
self.assertTrue(na.is_ipv6)
|
||||||
|
|
||||||
def test_no_brackets(self):
|
def test_no_brackets(self):
|
||||||
na = NetworkAddress.from_parts("::1", 8000)
|
na = NetworkAddress("::1", 8000)
|
||||||
self.assertEqual(na.host, "::1")
|
self.assertEqual(na.host, "::1")
|
||||||
|
|
||||||
def test_ipv4_passthrough(self):
|
def test_ipv4_passthrough(self):
|
||||||
na = NetworkAddress.from_parts("127.0.0.1", 30000)
|
na = NetworkAddress("127.0.0.1", 30000)
|
||||||
self.assertEqual(na.host, "127.0.0.1")
|
self.assertEqual(na.host, "127.0.0.1")
|
||||||
self.assertFalse(na.is_ipv6)
|
self.assertFalse(na.is_ipv6)
|
||||||
|
|
||||||
def test_hostname_passthrough(self):
|
def test_hostname_passthrough(self):
|
||||||
na = NetworkAddress.from_parts("myhost", 30000)
|
na = NetworkAddress("myhost", 30000)
|
||||||
self.assertEqual(na.host, "myhost")
|
self.assertEqual(na.host, "myhost")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user