TensorFlow循环创建多模型时OOM错误的解决方法
解决Keras循环训练模型的GPU内存不足问题
针对循环训练模型时出现的OOM错误,只调用keras.backend.clear_session()不够,得结合多步内存清理操作,同时配合显存优化设置,才能稳定跑完100个模型的测试,具体方案如下:
完整的内存清理流程:训练完每个模型后,按顺序执行三步操作,确保Python和Keras层面的内存都被释放:
- 删除模型变量:
del model,先断掉Python对模型对象的引用 - 清理Keras会话:
keras.backend.clear_session(),重置Keras的内部状态,释放会话占用的显存 - 强制垃圾回收:
import gc; gc.collect(),让Python主动回收未被引用的内存块
- 删除模型变量:
显存优化设置:提前配置TensorFlow的显存增长模式,避免程序启动就占满所有GPU显存,给后续模型训练留足空间:
import tensorflow as tf gpus = tf.config.experimental.list_physical_devices('GPU') if gpus: try: for gpu in gpus: tf.config.experimental.set_memory_growth(gpu, True) except RuntimeError as e: print(e)训练过程减载:
- 调小
model.fit()的batch_size,比如从32降到16,降低单步训练的显存占用 - 关闭不必要的训练回调,比如不要保存每一轮的checkpoint,只保留最优模型;如果用了TensorBoard,及时清理日志文件或限制保存的变量数量
- 确保训练时没有把模型、训练历史、中间数据存入全局列表/变量,每轮循环后所有临时数据都要清理
- 调小
修改后的完整伪代码:
import gc import keras.backend as K import tensorflow as tf # 提前配置显存增长 gpus = tf.config.experimental.list_physical_devices('GPU') if gpus: try: for gpu in gpus: tf.config.experimental.set_memory_growth(gpu, True) except RuntimeError as e: print(e) for hyperparameter in hyperparameters: model = createModel(hyperparameter) # 执行训练 model.fit(...) # 清理内存 del model K.clear_session() gc.collect()
内容的提问来源于stack exchange,提问作者ThomasCodesThings
相关产品推荐
相关产品推荐

