You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用Slurm提交任务时jax.distributed.initialize挂起问题排查

问题

尝试使用Slurm在多台主机上分布式运行JAX脚本,初始化代码如下:

import os
import jax

jax.distributed.initialize(
    coordinator_address=f"{os.environ['COORDINATOR_ADDR']}:{os.environ['COORDINATOR_PORT']}",
    num_processes=int(os.environ['SLURM_NTASKS']),
    process_id=int(os.environ['SLURM_PROCID'])
)

Slurm提交脚本配置如下:

#!/bin/bash
#SBATCH -N 2 -n 32
#SBATCH --job-name=deepmd-jax
#SBATCH --partition=gpu
#SBATCH --gres=gpu:4
#SBATCH --mem=512G

set -e

export JAX_PLATFORM_NAME=gpu
export TF_GPU_ALLOCATOR=cuda_malloc_async
MASTER_NODE=$(scontrol show hostnames $SLURM_JOB_NODELIST | head -n 1)
export COORDINATOR_ADDR=$(ssh $MASTER_NODE "hostname -I" | awk '{print $1}')
export COORDINATOR_PORT=12345 

# Note: srun command is missing in the original snippet

脚本在初始化阶段无限挂起,该如何解决?

解决方法

1. 补全srun启动命令

Slurm分布式任务必须通过srun启动所有进程,否则任务无法在多节点上正确分发,这是初始化挂起的核心原因。在提交脚本末尾添加:

srun --export=ALL python your_jax_script.py
  • --export=ALL确保所有环境变量(包括COORDINATOR_ADDR、COORDINATOR_PORT等)传递给每个进程
  • 替换your_jax_script.py为你的实际JAX脚本文件名

2. 简化主节点地址获取逻辑

通过ssh获取主节点IP可能引发权限或网络延迟问题,直接使用主节点主机名即可(JAX会自动解析),修改COORDINATOR_ADDR的设置:

MASTER_NODE=$(scontrol show hostnames $SLURM_JOB_NODELIST | head -n 1)
export COORDINATOR_ADDR=$MASTER_NODE

如果必须使用IP,可改用更可靠的方式(避免ssh):

export COORDINATOR_ADDR=$(hostname -I | awk '{print $1}' | head -n 1)

(注:此命令在主节点执行时会获取自身IP,而srun会确保主节点的进程先启动作为协调器)

3. 验证节点间网络连通性

挂起可能是因为工作节点无法访问主节点的协调端口,在提交脚本中添加前置测试:

# 测试主节点端口可达性
srun --nodes=$SLURM_NNODES --ntasks-per-node=1 ping -c 2 $MASTER_NODE
srun --nodes=$SLURM_NNODES --ntasks-per-node=1 telnet $MASTER_NODE $COORDINATOR_PORT || echo "Port unreachable"

如果测试失败,联系集群管理员开放节点间的12345端口,或更换一个未被占用的端口(比如29500,JAX常用默认端口)

4. 匹配任务数与GPU资源配置

当前配置-N 2 -n 32(2节点共32任务)搭配--gres=gpu:4(每节点4GPU),意味着每个GPU将绑定8个进程,可能导致资源竞争。建议调整任务数与GPU的绑定关系,比如:

#SBATCH -N 2 --ntasks-per-node=8 --gres=gpu:4

这样每个GPU绑定2个进程,更符合常规分布式训练的资源分配逻辑,同时在srun中添加GPU绑定参数:

srun --export=ALL --gres=gpu:4 --ntasks-per-node=8 python your_jax_script.py

5. 增加JAX初始化调试日志

在Python脚本中添加调试日志,确认每个进程的参数是否正确:

import os
import jax
import logging

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

logger.info(f"Process {os.environ['SLURM_PROCID']} connecting to coordinator {os.environ['COORDINATOR_ADDR']}:{os.environ['COORDINATOR_PORT']}")
logger.info(f"Total processes: {os.environ['SLURM_NTASKS']}")

jax.distributed.initialize(
    coordinator_address=f"{os.environ['COORDINATOR_ADDR']}:{os.environ['COORDINATOR_PORT']}",
    num_processes=int(os.environ['SLURM_NTASKS']),
    process_id=int(os.environ['SLURM_PROCID'])
)

logger.info(f"Process {os.environ['SLURM_PROCID']} initialized successfully")

通过日志可以定位是哪个进程无法连接协调器,或协调器未正常启动。

内容的提问来源于stack exchange,提问作者link89

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.11 17:12:38