使用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
相关产品推荐
相关产品推荐

