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

自定义ActorCritic模型调用build后仍提示未设置输入形状无法保存如何解决

问题原因

你使用的是Keras子类化方式实现的自定义ActorCritic模型,这类模型的元信息跟踪逻辑和序列式/函数式模型不同,仅调用build()方法只能设置模型顶层的输入形状,不会完整触发内部所有层的形状关联注册,保存时的形状校验会判定输入形状未设置。

解决方法

下面提供3种可行的解决方案,按需选择即可:

  • 方案1:触发一次前向传播(最简单通用)
    在build之后、save之前,传入符合输入形状的假张量跑一次前向推理,触发全链路形状初始化:
self.global_model = ActorCritic(self.action_size, self.state_size)
self.global_model.build((None, *self.state_size))
# 新增以下一行即可
_ = self.global_model(tf.random.normal((1, *self.state_size)))
self.global_model.save(self.save_path)
  • 方案2:保存时显式指定输入签名
    不需要跑前向推理,直接在调用save时手动指定模型的输入签名即可通过校验:
self.global_model.save(
    self.save_path,
    signatures=self.global_model.call.get_concrete_function(
        tf.TensorSpec(shape=(None, *self.state_size), dtype=tf.float32, name="state_input")
    )
)
  • 方案3:仅保存模型权重
    如果你不需要保存完整的模型结构,只需要保存训练好的参数,用save_weights方法不会触发形状校验,更轻量:
# 保存权重代码
self.global_model.save_weights(os.path.join(self.save_path, "actor_critic_weights"))

# 后续加载权重的示例代码
loaded_model = ActorCritic(action_size=2, state_size=(9,23))
loaded_model.build((None, 9,23))
loaded_model.load_weights(os.path.join(os.getcwd(), 'model', 'actor_critic_weights'))
注意事项

如果选择保存完整的SavedModel格式,加载时必须传入custom_objects参数指定自定义的模型类,否则会报错:

loaded_model = tf.keras.models.load_model(
    "model",
    custom_objects={"ActorCritic": ActorCritic}
)

内容的提问来源于stack exchange,提问作者Ne-al

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 14:12:04