From ae62279071fefa0936041bca0526009e7a170766 Mon Sep 17 00:00:00 2001 From: Lianmin Zheng Date: Fri, 7 Nov 2025 12:13:21 -0800 Subject: [PATCH] Fix data parallel controller launch for num nodes > 2 (#12822) --- python/sglang/srt/entrypoints/engine.py | 16 +++--- .../srt/managers/data_parallel_controller.py | 52 +++++++++++-------- python/sglang/srt/managers/scheduler.py | 4 +- python/sglang/srt/server_args.py | 2 +- 4 files changed, 41 insertions(+), 33 deletions(-) diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index a065d665a..5c65cc12f 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -788,24 +788,26 @@ def _launch_subprocesses( scheduler_procs = [] if server_args.dp_size == 1: + # Launch tensor parallel scheduler processes memory_saver_adapter = TorchMemorySaverAdapter.create( enable=server_args.enable_memory_saver ) scheduler_pipe_readers = [] - nnodes_per_tp_group = max(server_args.nnodes // server_args.pp_size, 1) + pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1) + nnodes_per_pp_rank = max(server_args.nnodes // server_args.pp_size, 1) + pp_rank_range = range( + pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank), + pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank + 1), + ) + + nnodes_per_tp_group = nnodes_per_pp_rank tp_size_per_node = server_args.tp_size // nnodes_per_tp_group tp_rank_range = range( tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group), tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1), ) - pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1) - pp_rank_range = range( - pp_size_per_node * (server_args.node_rank // nnodes_per_tp_group), - pp_size_per_node * (server_args.node_rank // nnodes_per_tp_group + 1), - ) - for pp_rank in pp_rank_range: for tp_rank in tp_rank_range: reader, writer = mp.Pipe(duplex=False) diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 92437aa80..cb897a643 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -49,8 +49,9 @@ from sglang.srt.tracing.trace import ( trace_slice_end, trace_slice_start, ) -from sglang.srt.utils import ( +from sglang.srt.utils.common import ( bind_port, + configure_ipv6, configure_logger, get_zmq_socket, kill_itself_when_parent_died, @@ -118,16 +119,16 @@ class DataParallelController: """A controller that dispatches requests to multiple data parallel workers.""" def __init__(self, server_args: ServerArgs, port_args: PortArgs) -> None: - # for dp balance - self.global_balance_id = 0 - # Parse args - self.max_total_num_tokens = None self.server_args = server_args self.port_args = port_args self.load_balance_method = LoadBalanceMethod.from_str( server_args.load_balance_method ) + self.run_scheduler_process = run_scheduler_process + + # For DP balance + self.global_balance_id = 0 # Init inter-process communication self.context = zmq.Context(1 + server_args.dp_size) @@ -162,8 +163,6 @@ class DataParallelController: self.launch_dp_schedulers(server_args, port_args) self.control_message_step = 1 - self.max_req_input_len = None - self.init_dispatcher() def send_to_all_workers(self, obj): @@ -281,8 +280,12 @@ class DataParallelController: # Determine the endpoint for inter-node communication if server_args.dist_init_addr is None: endpoint = f"tcp://127.0.0.1:{server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA}" + elif server_args.dist_init_addr.startswith("["): # ipv6 address + port, host = configure_ipv6(server_args.dist_init_addr) + endpoint = f"tcp://{host}:{int(port) + DP_ATTENTION_HANDSHAKE_PORT_DELTA}" else: - endpoint = f"tcp://{server_args.dist_init_addr}" + host, port = server_args.dist_init_addr.split(":") + endpoint = f"tcp://{host}:{int(port) + DP_ATTENTION_HANDSHAKE_PORT_DELTA}" if server_args.node_rank == 0: # Node 0: Broadcast worker ports to all other nodes @@ -326,8 +329,8 @@ class DataParallelController: logger.debug(f"Connecting to node 0 to receive worker ports") req_socket = get_zmq_socket(self.context, zmq.REQ, endpoint, False) - req_socket.setsockopt(zmq.RCVTIMEO, 60 * 1000) # 1 minute timeout - req_socket.setsockopt(zmq.SNDTIMEO, 60 * 1000) + req_socket.setsockopt(zmq.RCVTIMEO, 600 * 1000) # 10 minute timeout + req_socket.setsockopt(zmq.SNDTIMEO, 600 * 1000) try: # Send handshake with our node rank @@ -381,19 +384,20 @@ class DataParallelController: scheduler_pipe_readers = [] - nnodes_per_tp_group = max(server_args.nnodes // server_args.pp_size, 1) + pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1) + nnodes_per_pp_rank = max(server_args.nnodes // server_args.pp_size, 1) + pp_rank_range = range( + pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank), + pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank + 1), + ) + + nnodes_per_tp_group = nnodes_per_pp_rank tp_size_per_node = server_args.tp_size // nnodes_per_tp_group tp_rank_range = range( tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group), tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1), ) - pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1) - pp_rank_range = range( - pp_size_per_node * (server_args.node_rank // nnodes_per_tp_group), - pp_size_per_node * (server_args.node_rank // nnodes_per_tp_group + 1), - ) - for pp_rank in pp_rank_range: for tp_rank in tp_rank_range: rank_port_args = port_args @@ -424,7 +428,7 @@ class DataParallelController: moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size) with self.env_lock, maybe_reindex_device_id(gpu_id) as gpu_id: proc = mp.Process( - target=run_scheduler_process, + target=self.run_scheduler_process, args=( server_args, rank_port_args, @@ -504,8 +508,14 @@ def run_data_parallel_controller_process( server_args: ServerArgs, port_args: PortArgs, pipe_writer, + data_parallel_controller_class=DataParallelController, ): + setproctitle.setproctitle("sglang::data_parallel_controller") + faulthandler.enable() kill_itself_when_parent_died() + parent_process = psutil.Process().parent() + + configure_logger(server_args) if server_args.enable_trace: process_tracing_init(server_args.otlp_traces_endpoint, "sglang") thread_label = "DP Controller" @@ -514,13 +524,9 @@ def run_data_parallel_controller_process( elif server_args.disaggregation_mode == "decode": thread_label = "Decode DP Controller" trace_set_thread_info(thread_label) - setproctitle.setproctitle("sglang::data_parallel_controller") - faulthandler.enable() - configure_logger(server_args) - parent_process = psutil.Process().parent() try: - controller = DataParallelController(server_args, port_args) + controller = data_parallel_controller_class(server_args, port_args) pipe_writer.send( { "status": "ready", diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 2a7afdb4e..92f4e1025 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2762,12 +2762,12 @@ def run_scheduler_process( dp_rank = int(os.environ["SGLANG_DP_RANK"]) if dp_rank is not None: prefix += f" DP{dp_rank}" + if server_args.pp_size > 1: + prefix += f" PP{pp_rank}" if server_args.tp_size > 1: prefix += f" TP{tp_rank}" if server_args.ep_size > 1: prefix += f" EP{moe_ep_rank}" - if server_args.pp_size > 1: - prefix += f" PP{pp_rank}" # Config the process setproctitle.setproctitle(f"sglang::scheduler{prefix.replace(' ', '_')}") diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 0b410d4ba..4367d729c 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -4026,7 +4026,7 @@ def prepare_server_args(argv: List[str]) -> ServerArgs: ZMQ_TCP_PORT_DELTA = 233 -DP_ATTENTION_HANDSHAKE_PORT_DELTA = 5 +DP_ATTENTION_HANDSHAKE_PORT_DELTA = 13 @dataclasses.dataclass