自定义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
相关产品推荐
相关产品推荐

