SLURM环境下LSTM网格搜索交叉验证的OOM-kill问题求助
问题描述
在SLURM环境中使用MPI并行训练LSTM模型,同时执行带10折交叉验证的网格搜索(CPU训练)时,触发OOM错误:
slurmstepd: error: Detected 1 oom-kill event(s) in StepId=778254.0 cgroup. Some of your processes may have been killed by the cgroup out-of-memory handler.
已通过以下命令监控内存,结果显示内存占用随时间持续增长:
vmstat -SM 300 120 > $SCRATCH/memoryusage_$SLURM_JOB_ID.out &
每个网格搜索组合的交叉验证流程末尾,已执行内存清理操作,但问题仍存在:
for model in Models + Models_inv: del model K.clear_session() gc.collect()
任务由主节点向从节点迭代分配,上述清理代码会被从节点多次执行。SBATCH脚本如下:
#!/bin/bash ##----------------------- Start job description ----------------------- #SBATCH --partition=standard #SBATCH --job-name=multi_RNN #SBATCH --nodes=1 #SBATCH --ntasks=40 #SBATCH --cpus-per-task=1 #SBATCH --mem-per-cpu=4096 #SBATCH --time=48:00:00 #SBATCH --mail-type=ALL #SBATCH --output=out-%j.log #SBATCH --error=err-%j.log #SBATCH --chdir=/home/v597/v597521/multirnn/ ##------------------------ End job description ------------------------ module load Anaconda3/2024.02-1 source ~/.bashrc conda activate lstm export SCRATCH="/home/v597/v597521/multirnn" mkdir -p $SCRATCH/$SLURM_JOB_ID vmstat -SM 300 120 > $SCRATCH/memoryusage_$SLURM_JOB_ID.out & srun --ntasks=$SLURM_NTASKS python /home/v597/v597521/multirnn/multirnn_par_mpi.py rm -rf $SCRATCH/$SLURM_JOB_ID
解决方案
1. 修复模型与数据的引用泄漏
- 清空模型列表本身:仅删除列表内的model对象,列表仍会持有空引用,需主动清空列表:
for model in Models + Models_inv: del model # 清空列表,释放所有引用 Models.clear() Models_inv.clear() K.clear_session() gc.collect() - 清理训练数据变量:每次交叉验证后删除训练/验证数据集变量,避免重复加载累积内存:
# 训练完成后立即删除数据变量 del X_train, y_train, X_val, y_val gc.collect()
2. 优化TensorFlow/keras的CPU内存配置
- 禁用TF预分配内存,限制线程数避免内存竞争:
import tensorflow as tf import keras.backend as K # 限制线程数,避免CPU内存过度占用 tf.config.threading.set_intra_op_parallelism_threads(1) tf.config.threading.set_inter_op_parallelism_threads(1) # 启用CPU内存动态增长,禁止预分配 physical_devices = tf.config.list_physical_devices('CPU') tf.config.experimental.set_memory_growth(physical_devices[0], True) - 补充TF图重置操作:在
K.clear_session()后添加图重置,彻底清理TF内部资源:K.clear_session() tf.compat.v1.reset_default_graph() # 兼容TF2.x版本 gc.collect()
3. 调整SLURM资源分配策略
- 减少单节点任务数或提升单进程内存配额:当前40进程×4G的配置可能存在内存竞争,可尝试:
#SBATCH --ntasks=20 #SBATCH --mem-per-cpu=8192 - 添加
--exclusive选项,确保进程独占CPU核心与内存:srun --ntasks=$SLURM_NTASKS --exclusive python /home/v597/v597521/multirnn/multirnn_par_mpi.py
4. 进程级隔离内存泄漏
将单个网格搜索任务封装为独立子进程,利用进程退出自动回收内存:
from multiprocessing import Process def run_single_task(params): # 此处放入单个网格搜索+10折交叉验证的完整逻辑 # 无需手动清理内存,进程结束后操作系统自动释放所有资源 pass # 在MPI从节点的任务循环中 for params in task_queue: p = Process(target=run_single_task, args=(params,)) p.start() p.join()
5. 定位具体泄漏点
- 监控单个MPI进程的内存变化:替换
vmstat为ps,跟踪每个Python进程的内存占用:while true; do ps aux | grep python | grep -v grep >> $SCRATCH/process_memory_$SLURM_JOB_ID.out; sleep 300; done & - 用
tracemalloc跟踪代码内的内存分配:import tracemalloc tracemalloc.start() # 执行一次完整的网格搜索+交叉验证 run_single_task(test_params) # 生成内存快照,输出Top10内存占用点 snapshot = tracemalloc.take_snapshot() top_stats = snapshot.statistics('lineno') print("[Top 10 Memory Leak Points]") for stat in top_stats[:10]: print(stat)
内容的提问来源于stack exchange,提问作者user26458368
相关产品推荐
相关产品推荐

