使用datasets.load_dataset在Slurm集群节点加载数据集时出现挂起问题
我在使用datasets.load_dataset加载数据时遇到了问题:在头节点上运行完全正常,但提交到Slurm节点后就会挂起。我已经在conda环境中安装了datasets库。
头节点上的正常运行命令
激活conda环境后,这条命令可以顺利执行:
python -c "from datasets import load_dataset; d=load_dataset(\"json\", data_files={\"train\": \"/scratch/train/shard1.jsonl\"}); print(d)"
Slurm集群上的挂起情况
当我提交作业到集群时,操作会挂起,我用的命令如下:
salloc --nodes 1 --qos interactive --time 00:15:00 --constraint gpu --account=my_account --mem=1G --gres=gpu:1 srun --nodes=1 --ntasks-per-node=1 --constraint=gpu --account=my_account --gres=gpu:1 \ bash -c ' source /global/homes/my_username/miniconda3/etc/profile.d/conda.sh && conda activate my_env && python -c "from datasets import load_dataset; load_dataset(\"json\", data_files={\"train\": \"/scratch/my_username/train/shard1.jsonl\"})" '
用sbatch提交也会出现类似的挂起情况。我用来测试的是一个极小的JSONL文件,内容如下:
{"text": "ACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGTACGT"}
排查与解决建议
看起来这个问题大概率和Slurm节点的环境配置或者数据集加载时的隐性依赖有关,我给你几个实用的排查方向:
确认conda环境是否正确激活:有时候在srun的bash脚本里,conda激活会因为环境变量没完全加载出问题。你可以在激活conda后加一句
conda info --envs,确认当前环境确实是my_env;也可以直接用环境的绝对路径运行Python,比如/global/homes/my_username/miniconda3/envs/my_env/bin/python,绕开激活脚本的潜在问题。检查数据集文件的权限与路径:虽然用了绝对路径,但Slurm节点对
/scratch目录的访问权限可能和头节点有差异?可以在srun脚本里先加一句ls -l /scratch/my_username/train/shard1.jsonl,确认文件存在且当前用户有读权限。禁用数据集缓存机制:
datasets默认会缓存加载的数据集,集群环境下缓存目录的读写权限或磁盘问题可能导致挂起。你可以在load_dataset中添加cache_dir=None参数,或者临时设置环境变量禁用缓存:python -c "import os; os.environ['HF_DATASETS_CACHE']='/dev/null'; from datasets import load_dataset; load_dataset(\"json\", data_files={\"train\": \"/scratch/my_username/train/shard1.jsonl\"})"增加日志输出定位卡点:在Python代码中添加打印语句或开启DEBUG日志,能帮你看到加载过程卡在了哪个阶段:
python -c "import datasets; datasets.logging.set_verbosity_debug(); from datasets import load_dataset; print('Starting load...'); d=load_dataset(\"json\", data_files={\"train\": \"/scratch/my_username/train/shard1.jsonl\"}); print('Load completed!')"核对节点间的依赖版本与网络:检查Slurm节点上的
datasets版本和头节点是否一致(两边都运行pip show datasets),避免版本差异导致的问题;另外,虽然用的是本地文件,但datasets可能会拉取少量元数据,确认Slurm节点是否有正常的网络访问权限。
备注:内容来源于stack exchange,提问作者ate50eggs

