Keras网络使用多进程计算样本损失时报The Session graph is empty错误
错误原因
TensorFlow/Keras的计算图、会话上下文是和进程强绑定的,不能直接跨进程传递已构造的模型对象:主进程训练好的模型通过多进程池序列化传递给子进程时,不会同步主进程的TensorFlow运行时状态,子进程调用model.evaluate时找不到对应的计算图,就会抛出The Session graph is empty错误。
解决方案
优先推荐批量运算方案,多进程引入的跨进程通信开销反而会降低运算效率:
方案1:直接批量计算单样本损失(最高效)
不需要拆分单样本循环,直接通过批量运算一步得到所有样本的损失值,运算效率远高于多进程方案:
# 若使用自带损失函数,关闭默认的均值聚合,返回每个样本的损失 loss_fn = tf.keras.losses.get(model.loss, reduction=tf.keras.losses.Reduction.NONE) # 批量预测所有样本输出 y_pred = model.predict(X_measured, batch_size=batch_size) # 直接得到所有样本的损失数组 per_sample_loss = loss_fn(y_train, y_pred).numpy()
方案2:必须使用多进程的修正方法
如果确实需要多进程处理,不要在主进程传递模型对象,改为子进程内部构造模型、加载权重:
# 主进程先保存训练好的权重 model.save_weights("temp_model_weights.h5") # 子进程运算逻辑 def get_loss(i, X_measured, y_train, weight_path): # 子进程内重新构造和训练时结构完全一致的模型 model = build_your_model() # 替换为你的模型结构定义代码 # 加载训练好的权重 model.load_weights(weight_path) # 执行单样本评估 return model.evaluate(x=X_measured[i:i+1,:],y=y_train[i:i+1,:],batch_size=None,verbose=0,steps=1) # 启动多进程池运算 pool = mp.Pool(mp.cpu_count()) samples_loss = pool.starmap(get_loss, [(j, X_measured, y_train, "temp_model_weights.h5") for j in range(X_measured.shape[0])]) pool.close() # 清理临时权重文件 import os os.remove("temp_model_weights.h5")
注意:如果使用GPU版本的TensorFlow,建议在子进程开头添加
tf.config.set_visible_devices([], 'GPU'),避免子进程占用GPU显存。
内容的提问来源于stack exchange,提问作者Xiaoqiao Chen
相关产品推荐
相关产品推荐

