Use ephemeral nccl port via get_free_port() (#18009)

Co-authored-by: Chanh Nguyen <cnguyen@linkedin.com>
This commit is contained in:
Chanh Nguyen
2026-02-15 08:32:47 -08:00
committed by GitHub
parent 7a607c4900
commit 597d17dd18
2 changed files with 45 additions and 74 deletions

View File

@@ -43,6 +43,7 @@ from sglang.srt.utils.common import (
get_device_memory_capacity,
get_device_name,
get_device_sm,
get_free_port,
get_int_env_var,
get_quantization_config,
is_blackwell_supported,
@@ -52,7 +53,6 @@ from sglang.srt.utils.common import (
is_hopper_with_cuda_12_3,
is_no_spec_infer_or_topk_one,
is_npu,
is_port_available,
is_remote_url,
is_sm90_supported,
is_sm100_supported,
@@ -1135,9 +1135,9 @@ class ServerArgs:
# suggest them to be explicit about kv_cache_dtype to avoid surprises
if (user_set_prefill or user_set_decode) and self.kv_cache_dtype == "auto":
logger.warning(
f"When specifying --nsa-prefill-backend or --nsa-decode-backend, "
f"you should also explicitly set --kv-cache-dtype (e.g., 'fp8_e4m3' or 'bfloat16'). "
f"DeepSeek V3.2 defaults to FP8 KV cache which may not be compatible with all backends."
"When specifying --nsa-prefill-backend or --nsa-decode-backend, "
"you should also explicitly set --kv-cache-dtype (e.g., 'fp8_e4m3' or 'bfloat16'). "
"DeepSeek V3.2 defaults to FP8 KV cache which may not be compatible with all backends."
)
if self.kv_cache_dtype == "auto":
@@ -1216,7 +1216,7 @@ class ServerArgs:
if not is_npu(): # CUDA or ROCm GPU
if self.enable_nsa_prefill_context_parallel:
logger.warning(
f"Context parallel feature is still under experiment. It has only been verified on Hopper platform."
"Context parallel feature is still under experiment. It has only been verified on Hopper platform."
)
if self.nsa_prefill_cp_mode == "in-seq-split":
# TODO Supports moe_dense_tp_size != 1, kv cache dtype = "fp8",moe_a2a_backend non-deepep and cross-machine operation .
@@ -1226,7 +1226,7 @@ class ServerArgs:
self.ep_size = self.tp_size
self.kv_cache_dtype = "bf16"
logger.warning(
f"For in-seq split mode, we have the following restrictions: moe_dense_tp_size == 1, moe_a2a_backend == deepep, ep_size == tp_size, kv_cache_dtype == bf16, batch_size == 1"
"For in-seq split mode, we have the following restrictions: moe_dense_tp_size == 1, moe_a2a_backend == deepep, ep_size == tp_size, kv_cache_dtype == bf16, batch_size == 1"
)
else:
self.enable_dp_attention = True
@@ -1663,7 +1663,7 @@ class ServerArgs:
if envs.SGLANG_EMBEDDINGS_SPARSE_HEAD.is_set():
self.disable_overlap_schedule = True
logger.warning(
f"Overlap scheduler is disabled when using sparse head for embedding model."
"Overlap scheduler is disabled when using sparse head for embedding model."
)
# TRTLLM AllReduce Fusion supports SM90/100, enable it by default
@@ -5500,7 +5500,7 @@ class ServerArgs:
def validate_transfer_engine(self):
if importlib.util.find_spec("mooncake.engine") is None:
logger.warning(
f"Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend."
"Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend."
)
return False
elif self.enable_memory_saver:
@@ -5603,14 +5603,7 @@ class PortArgs:
worker_ports: Optional[List[int]] = None,
) -> PortArgs:
if server_args.nccl_port is None:
nccl_port = server_args.port + random.randint(100, 1000)
while True:
if is_port_available(nccl_port):
break
if nccl_port < 60000:
nccl_port += 42
else:
nccl_port -= 43
nccl_port = get_free_port()
else:
nccl_port = server_args.nccl_port
@@ -5671,7 +5664,7 @@ class PortArgs:
# Skip check when using worker_ports since the port is already bound by our ZMQ socket
if dp_rank is None or worker_ports is None:
wait_port_available(scheduler_input_port, "scheduler_input_port")
except ValueError as e:
except ValueError:
logger.exception(
f"Port is already in use. {dist_init_port=} {port_base=} {detokenizer_port=} {nccl_port=} {scheduler_input_port=}"
)

View File

@@ -1,6 +1,6 @@
import json
import unittest
from unittest.mock import patch
from unittest.mock import MagicMock, patch
from sglang.srt.server_args import PortArgs, ServerArgs, prepare_server_args
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
@@ -44,10 +44,29 @@ class TestLoadBalanceMethod(unittest.TestCase):
class TestPortArgs(unittest.TestCase):
@patch("sglang.srt.server_args.is_port_available")
@patch("sglang.srt.server_args.get_free_port")
@patch("sglang.srt.server_args.tempfile.NamedTemporaryFile")
def test_init_new_standard_case(self, mock_temp_file, mock_is_port_available):
mock_is_port_available.return_value = True
def test_init_new_with_nccl_port_none(self, mock_temp_file, mock_get_free_port):
"""Test that get_free_port() is called when nccl_port is None"""
mock_temp_file.return_value.name = "temp_file"
mock_get_free_port.return_value = 45678 # Mock ephemeral port
# Use MagicMock here to verify get_free_port is called
server_args = MagicMock()
server_args.nccl_port = None
server_args.enable_dp_attention = False
server_args.tokenizer_worker_num = 1
port_args = PortArgs.init_new(server_args)
# Verify get_free_port was called
mock_get_free_port.assert_called_once()
# Verify the returned port is used
self.assertEqual(port_args.nccl_port, 45678)
@patch("sglang.srt.server_args.tempfile.NamedTemporaryFile")
def test_init_new_standard_case(self, mock_temp_file):
mock_temp_file.return_value.name = "temp_file"
server_args = ServerArgs(model_path="dummy")
@@ -62,9 +81,7 @@ class TestPortArgs(unittest.TestCase):
self.assertTrue(port_args.detokenizer_ipc_name.startswith("ipc://"))
self.assertIsInstance(port_args.nccl_port, int)
@patch("sglang.srt.server_args.is_port_available")
def test_init_new_with_single_node_dp_attention(self, mock_is_port_available):
mock_is_port_available.return_value = True
def test_init_new_with_single_node_dp_attention(self):
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
@@ -82,10 +99,7 @@ class TestPortArgs(unittest.TestCase):
self.assertTrue(port_args.detokenizer_ipc_name.startswith("tcp://127.0.0.1:"))
self.assertIsInstance(port_args.nccl_port, int)
@patch("sglang.srt.server_args.is_port_available")
def test_init_new_with_dp_rank(self, mock_is_port_available):
mock_is_port_available.return_value = True
def test_init_new_with_dp_rank(self):
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None
@@ -102,13 +116,9 @@ class TestPortArgs(unittest.TestCase):
self.assertTrue(port_args.detokenizer_ipc_name.startswith("tcp://192.168.1.1:"))
self.assertIsInstance(port_args.nccl_port, int)
@patch("sglang.srt.server_args.is_port_available")
def test_init_new_with_ipv4_address(self, mock_is_port_available):
mock_is_port_available.return_value = True
def test_init_new_with_ipv4_address(self):
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None
server_args.enable_dp_attention = True
@@ -124,10 +134,7 @@ class TestPortArgs(unittest.TestCase):
self.assertTrue(port_args.detokenizer_ipc_name.startswith("tcp://192.168.1.1:"))
self.assertIsInstance(port_args.nccl_port, int)
@patch("sglang.srt.server_args.is_port_available")
def test_init_new_with_malformed_ipv4_address(self, mock_is_port_available):
mock_is_port_available.return_value = True
def test_init_new_with_malformed_ipv4_address(self):
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None
@@ -143,12 +150,7 @@ class TestPortArgs(unittest.TestCase):
"please provide --dist-init-addr as host:port", str(context.exception)
)
@patch("sglang.srt.server_args.is_port_available")
def test_init_new_with_malformed_ipv4_address_invalid_port(
self, mock_is_port_available
):
mock_is_port_available.return_value = True
def test_init_new_with_malformed_ipv4_address_invalid_port(self):
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None
@@ -157,16 +159,11 @@ class TestPortArgs(unittest.TestCase):
server_args.nnodes = 2
server_args.dist_init_addr = "192.168.1.1:abc"
with self.assertRaises(ValueError) as context:
with self.assertRaises(ValueError):
PortArgs.init_new(server_args)
@patch("sglang.srt.server_args.is_port_available")
@patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True)
def test_init_new_with_ipv6_address(
self, mock_is_valid_ipv6, mock_is_port_available
):
mock_is_port_available.return_value = True
def test_init_new_with_ipv6_address(self, mock_is_valid_ipv6):
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None
@@ -186,13 +183,8 @@ class TestPortArgs(unittest.TestCase):
)
self.assertIsInstance(port_args.nccl_port, int)
@patch("sglang.srt.server_args.is_port_available")
@patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=False)
def test_init_new_with_invalid_ipv6_address(
self, mock_is_valid_ipv6, mock_is_port_available
):
mock_is_port_available.return_value = True
def test_init_new_with_invalid_ipv6_address(self, mock_is_valid_ipv6):
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None
@@ -206,12 +198,7 @@ class TestPortArgs(unittest.TestCase):
self.assertIn("invalid IPv6 address", str(context.exception))
@patch("sglang.srt.server_args.is_port_available")
def test_init_new_with_malformed_ipv6_address_missing_bracket(
self, mock_is_port_available
):
mock_is_port_available.return_value = True
def test_init_new_with_malformed_ipv6_address_missing_bracket(self):
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None
@@ -225,13 +212,10 @@ class TestPortArgs(unittest.TestCase):
self.assertIn("invalid IPv6 address format", str(context.exception))
@patch("sglang.srt.server_args.is_port_available")
@patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True)
def test_init_new_with_malformed_ipv6_address_missing_port(
self, mock_is_valid_ipv6, mock_is_port_available
self, mock_is_valid_ipv6
):
mock_is_port_available.return_value = True
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None
@@ -247,13 +231,10 @@ class TestPortArgs(unittest.TestCase):
"a port must be specified in IPv6 address", str(context.exception)
)
@patch("sglang.srt.server_args.is_port_available")
@patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True)
def test_init_new_with_malformed_ipv6_address_invalid_port(
self, mock_is_valid_ipv6, mock_is_port_available
self, mock_is_valid_ipv6
):
mock_is_port_available.return_value = True
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None
@@ -267,13 +248,10 @@ class TestPortArgs(unittest.TestCase):
self.assertIn("invalid port in IPv6 address", str(context.exception))
@patch("sglang.srt.server_args.is_port_available")
@patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True)
def test_init_new_with_malformed_ipv6_address_wrong_separator(
self, mock_is_valid_ipv6, mock_is_port_available
self, mock_is_valid_ipv6
):
mock_is_port_available.return_value = True
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None