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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 10:09:03