Jupyter中Keras模型重复训练速度变慢问题求助
我之前也碰到过几乎一模一样的问题!在Jupyter Notebook里反复构建Keras/TensorFlow模型时,训练速度会越来越慢,重启内核就立刻恢复正常,当时折腾了好一阵子才找到几个有效的解决办法,分享给你:
先说说可能的核心原因
虽然你调用了limit_mem(),但TensorFlow在交互式环境(比如Jupyter)里,很容易残留未彻底清理的计算图节点、变量缓存或者会话资源——每次构建新模型,这些残留都会在默认计算图里叠加,导致计算图越来越庞大,训练时的调度开销自然越来越高。另外,BatchNormalization层的移动均值/方差这类状态变量,也可能因为旧会话没有完全销毁而残留,间接影响新模型的初始化和计算效率。
几个实测有效的解决方案
1. 彻底升级你的内存清理函数
原来的limit_mem()不够彻底,试试换成这个版本,同时清理会话、计算图和Keras的残留资源:
def limit_mem(): """彻底清理GPU内存、会话和计算图""" import tensorflow as tf from keras import backend as K K.clear_session() # 销毁所有Keras创建的张量和模型 tf.reset_default_graph() # 清空默认计算图的所有节点 # 重新配置GPU内存增长策略 cfg = tf.ConfigProto() cfg.gpu_options.allow_growth = True K.set_session(tf.Session(config=cfg))
每次构建新模型前调用这个函数,能最大程度减少资源残留。
2. 给每个模型分配独立的计算图
每次构建模型时,主动创建一个新的计算图上下文,确保模型的所有节点都和之前的模型完全隔离:
import tensorflow as tf from keras import backend as K from keras.models import Sequential from keras.layers import Dense, BatchNormalization, Dropout from keras.optimizers import Adam def build_model(dropout, dense1, dense2, dense3, lr): # 新建独立计算图 with tf.Graph().as_default(): # 为当前图配置专属会话 cfg = tf.ConfigProto() cfg.gpu_options.allow_growth = True sess = tf.Session(config=cfg) K.set_session(sess) # 在这里正常构建你的MLP模型 model = Sequential() model.add(Dense(dense1, activation='relu', input_shape=(6000,))) model.add(BatchNormalization()) model.add(Dropout(dropout)) model.add(Dense(dense2, activation='relu')) model.add(BatchNormalization()) model.add(Dropout(dropout)) model.add(Dense(dense3, activation='relu')) model.add(BatchNormalization()) model.add(Dropout(dropout)) model.add(Dense(1, activation='sigmoid')) model.compile(optimizer=Adam(lr=lr), loss='binary_crossentropy') return model
这样每次训练完的模型资源,会随着计算图的销毁被垃圾回收,不会累积。
3. 检查自定义回调函数的内存泄漏
你自定义的Metrics回调函数有没有可能持有模型的引用,或者保存了大量未清理的临时数据?比如如果回调里长期缓存训练日志、引用模型内部张量,会导致旧模型无法被彻底回收,进而占用资源。可以先暂时去掉自定义回调,测试训练速度是否还会变慢——如果恢复正常,就需要优化回调函数,确保不会残留不必要的对象引用。
4. 切换到TensorFlow 2.x的原生Keras
如果你还在使用TensorFlow 1.x版本的Keras,强烈建议升级到TensorFlow 2.x,用tf.keras来构建模型。TF2.x默认的即时执行模式(Eager Execution)不再依赖全局计算图,每次构建模型都是完全独立的,资源管理也更智能,基本不会出现这种逐次变慢的问题。
另外你提到每次重启后的训练散点图一致,说明模型的初始化和训练逻辑完全没问题,问题确实出在资源累积导致的速度下降,不是模型本身的bug。
内容的提问来源于stack exchange,提问作者B. Verhoeff

