TensorFlow Keras SavedModel两次保存加载后抛出TypeError问题
Keras自定义层模型多次保存加载报错解决方案
问题现象
当创建包含一个或多个自定义层的Keras模型时,可使用model.save()方法以TensorFlow SavedModel格式持久化存储模型。
使用tf.keras.models.load_model()函数从文件系统加载模型后可再次保存,但第二次加载该SavedModel时会触发如下异常:
TypeError: f(inputs, training, training, training, training, *, training, training) missing 1 required argument: training
问题复现代码
import tensorflow as tf class CustomLayer(tf.keras.layers.Layer): def call(self, inputs, *args, **kwargs): return inputs model1 = tf.keras.Sequential([ CustomLayer() ]) model1.build((None, 1)) model1.compile() model1.save("model1") model2 = tf.keras.models.load_model("model1") model2.save("model2") # 该行代码会抛出TypeError model3 = tf.keras.models.load_model("model2")
问题原因
自定义层的call方法没有显式声明training参数,Keras在序列化、反序列化过程中会自动向call方法注入training参数,多次保存后会出现参数签名重复、参数匹配失败的问题。
解决方案
- 方案1:自定义层实现时显式在
call方法中声明training参数,不需要使用该参数也可以保留定义:
class CustomLayer(tf.keras.layers.Layer): def call(self, inputs, training=None, *args, **kwargs): return inputs
- 方案2:加载模型时显式传入自定义类声明,避免Keras自动解析签名出错:
model3 = tf.keras.models.load_model("model2", custom_objects={"CustomLayer": CustomLayer})
- 方案3:如果加载模型仅用于推理,不需要保留训练配置,可以添加
compile=False参数跳过编译步骤:
model3 = tf.keras.models.load_model("model2", compile=False)
内容的提问来源于stack exchange,提问作者James Mishra
相关产品推荐
相关产品推荐

