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

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正确跟踪:

  1. 放弃使用model.predict(),改用模型的__call__方法(直接调用模型实例),因为__call__原生支持图模式。
  2. 将self.agent.model注册为tf.Module的子模块,让SavedModel自动序列化它的参数和图结构。
  3. 用@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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 02:52:07