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

如何保存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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.09 16:48:11