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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 10:13:18