TensorFlow自定义模型保存加载报错求助:RuntimeError与ValueError
问题原因分析与正确实现方案
错误原因拆解
第一种方法报错原因
@tf.function装饰的函数运行在TensorFlow的静态图模式下,而self.agent.model.predict()是Keras高层API,底层依赖旧的Session会话机制,还会在内部隐式创建会话。在图函数里直接调用这类依赖会话的动态图操作,就会触发"Cannot get session inside Tensorflow graph function"错误——图模式不允许这种跨模式的调用。
第二种方法报错原因
tf.numpy_function包装的Python自定义函数(比如你写的custom_predict、convert_to_numpy)无法被SavedModel序列化保存。SavedModel只能存储TensorFlow原生的图操作,加载时找不到被序列化的Python回调函数,就会报"callback pyfunc_0 is not found"。
正确实现思路
核心是把预测逻辑改成纯TensorFlow图兼容的操作,同时确保模型的所有可训练参数和子模块被tf.Module正确跟踪:
- 放弃使用
model.predict(),改用模型的__call__方法(直接调用模型实例),因为__call__原生支持图模式。 - 将
self.agent.model注册为tf.Module的子模块,让SavedModel自动序列化它的参数和图结构。 - 用
@tf.function装饰预测函数时,全程使用TensorFlow张量操作,避免混合NumPy操作。
完整代码示例
import tensorflow as tf from rl.agents.dqn import DQNAgent class CustomModel(tf.Module): def __init__(self, agent: DQNAgent): super().__init__() # 关键:将agent的模型注册为子模块,让tf.Module跟踪它 self.model = agent.model @tf.function(input_signature=[tf.TensorSpec(shape=(None, 94), dtype=tf.float32)]) def predict(self, s): # 用TensorFlow操作调整输入形状,替代NumPy的np.newaxis input_tensor = tf.expand_dims(s, axis=1) # 直接调用模型__call__方法,training=False关闭训练模式 q_values = self.model(input_tensor, training=False) scores = q_values[:, 1] return scores # 保存模型 def save_model(custom_model: CustomModel, save_path): tf.saved_model.save(custom_model, save_path) # 加载并调用 def load_and_predict(save_path, test_X): loaded_model = tf.saved_model.load(save_path) q_values = loaded_model.predict(test_X) return q_values
注意事项
- 确保
agent.model是标准的tf.keras.Model实例,keras-rl2的DQNAgent的model属性通常满足这个要求。 - 调用模型时设置
training=False,避免触发Dropout、BatchNorm等训练专属行为,保证预测结果稳定。 - 输入形状调整用
tf.expand_dims替代NumPy操作,确保全程在图模式下执行。
内容的提问来源于stack exchange,提问作者khouloud
相关产品推荐
相关产品推荐

