如何保存Call()方法含位置参数的Keras子类模型?
解决Keras子类模型保存时的多参数
call方法报错问题 这个问题我之前也碰到过,本质是Keras的SavedModel机制在处理子类模型的多参数call方法时,无法自动推断所有参数的默认值或输入签名,导致保存/加载时找不到参数的传递规则。下面给你几个可行的解决方案:
方法1:给call方法的额外参数设置默认值
最简单的方式就是给training、mask1等非核心输入参数设置默认值,让Keras明确这些参数的默认传递规则,这样保存和加载时就不会报错。修改后的代码如下:
import tensorflow as tf class MyModel(tf.keras.Model): def __init__(self): super(MyModel, self).__init__() self.dense1 = tf.keras.layers.Dense(4, activation=tf.nn.relu) self.dense2 = tf.keras.layers.Dense(5, activation=tf.nn.softmax) @tf.function # 给额外参数添加默认值 def call(self, enc_input, dec_input, training=False, mask1=None, mask2=None, mask3=None): x = self.dense1(enc_input) return self.dense2(x) x = tf.random.normal((10,20)) model = MyModel() y = model(x, x, False, None, None, None) # 现在可以正常保存完整模型 tf.keras.models.save_model(model, '/saved') # 加载测试 loaded_model = tf.keras.models.load_model('/saved') # 加载后可以用简化参数调用,也可以传全参数 loaded_y = loaded_model(x, x) print(tf.reduce_all(tf.equal(y, loaded_y))) # 输出True,验证结果一致
方法2:用Functional API包装子类模型
如果不想修改call方法的参数定义,可以用Keras的Functional API包装你的子类模型,明确输入的结构和数量,这样保存的模型会有清晰的输入签名:
import tensorflow as tf class MyModel(tf.keras.Model): def __init__(self): super(MyModel, self).__init__() self.dense1 = tf.keras.layers.Dense(4, activation=tf.nn.relu) self.dense2 = tf.keras.layers.Dense(5, activation=tf.nn.softmax) @tf.function def call(self, enc_input, dec_input, training, mask1, mask2, mask3): x = self.dense1(enc_input) return self.dense2(x) x = tf.random.normal((10,20)) model = MyModel() y = model(x, x, False, None, None, None) # 用Functional API定义输入和输出 enc_input = tf.keras.Input(shape=(20,)) dec_input = tf.keras.Input(shape=(20,)) # 调用子类模型时传入默认参数值 output = model(enc_input, dec_input, training=False, mask1=None, mask2=None, mask3=None) # 构建新的Functional模型 saved_model = tf.keras.Model(inputs=[enc_input, dec_input], outputs=output) # 保存这个Functional模型 tf.keras.models.save_model(saved_model, '/saved') # 加载测试 loaded_model = tf.keras.models.load_model('/saved') loaded_y = loaded_model([x, x]) print(tf.reduce_all(tf.equal(y, loaded_y))) # 输出True
方法3:保存时显式指定输入签名
如果需要严格保留原call方法的参数定义,也可以在保存时通过signatures参数显式指定输入张量的签名,告诉SavedModel如何处理这些参数:
import tensorflow as tf class MyModel(tf.keras.Model): def __init__(self): super(MyModel, self).__init__() self.dense1 = tf.keras.layers.Dense(4, activation=tf.nn.relu) self.dense2 = tf.keras.layers.Dense(5, activation=tf.nn.softmax) @tf.function def call(self, enc_input, dec_input, training, mask1, mask2, mask3): x = self.dense1(enc_input) return self.dense2(x) x = tf.random.normal((10,20)) model = MyModel() y = model(x, x, False, None, None, None) # 定义每个输入的张量签名 input_signature = [ tf.TensorSpec(shape=(None, 20), dtype=tf.float32, name='enc_input'), tf.TensorSpec(shape=(None, 20), dtype=tf.float32, name='dec_input'), tf.TensorSpec(shape=(), dtype=tf.bool, name='training'), tf.TensorSpec(shape=None, dtype=tf.float32, name='mask1'), tf.TensorSpec(shape=None, dtype=tf.float32, name='mask2'), tf.TensorSpec(shape=None, dtype=tf.float32, name='mask3') ] # 保存时指定签名 tf.keras.models.save_model( model, '/saved', signatures=model.call.get_concrete_function(input_signature) ) # 加载测试 loaded_model = tf.keras.models.load_model('/saved') # 需要按照签名传入所有参数 loaded_y = loaded_model(x, x, False, None, None, None) print(tf.reduce_all(tf.equal(y, loaded_y))) # 输出True
为什么原来的代码会报错?
Keras在保存子类模型时,会尝试自动推断模型的输入规范。你的call方法包含多个没有默认值的参数,SavedModel无法确定这些参数的默认传递方式,所以在加载或保存时会抛出参数缺失的错误。通过上述方法明确参数的默认值或输入签名,就能让Keras正确处理这些参数,实现完整模型的保存。
内容的提问来源于stack exchange,提问作者Uchiha Madara
相关产品推荐
相关产品推荐

