Fix data parallel controller launch for num nodes > 2 (#12822)

This commit is contained in:
Lianmin Zheng
2025-11-07 12:13:21 -08:00
committed by GitHub
parent 5c9273c032
commit ae62279071
4 changed files with 41 additions and 33 deletions

View File

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