序列化重载Keras模型后自定义predict方法失效问题求助
解决Keras序列化模型后自定义predict方法失效的问题
问题原因
Keras的SavedModel序列化机制默认仅保存模型的层结构、权重和训练配置,不会完整保留自定义子类模型的实例方法(比如你重载的predict)。重新加载后得到的是tf.keras.Model的通用实例,而非你定义的MyModel类实例,因此自定义的predict方法会被忽略,仅执行默认的predict逻辑(内部调用call方法)。
解决方案
方案一:加载时指定自定义类
要让加载后的模型仍是你的自定义类实例,需要在加载时通过custom_objects参数传入你的模型类,同时建议实现get_config方法以确保配置正确序列化:
- 完善自定义模型类:
import tensorflow as tf import numpy as np class MyModel(tf.keras.Model): def __init__(self): super(MyModel, self).__init__() self.dense1 = tf.keras.layers.Dense(1) def call(self, inputs): tf.print('Hello from call') return self.dense1(inputs) def predict(self, inputs): tf.print('Hello from predict') return self(inputs) def get_config(self): # 继承父类的配置序列化逻辑 return super().get_config()
- 加载模型时传入自定义类:
# 保存模型(使用默认的SavedModel格式) model.save("my_model") # 加载时指定custom_objects reconstructed_model = tf.keras.models.load_model("my_model", custom_objects={"MyModel": MyModel})
此时调用reconstructed_model.predict(test_input)会执行你自定义的predict方法。
方案二:避免重载predict,改用自定义推断方法
Keras的predict方法内置了批量处理、设备分发等逻辑,直接重载可能破坏这些功能。更稳健的方式是定义独立的自定义推断方法:
class MyModel(tf.keras.Model): def __init__(self): super(MyModel, self).__init__() self.dense1 = tf.keras.layers.Dense(1) def call(self, inputs): tf.print('Hello from call') return self.dense1(inputs) def custom_predict(self, inputs): tf.print('Hello from custom predict') return self(inputs)
加载模型时同样传入custom_objects,之后直接调用reconstructed_model.custom_predict(test_input)即可。
方案三:手动复制权重到自定义模型实例
如果加载后得到的是通用模型实例,可以手动将权重复制到新创建的自定义模型中:
# 加载通用模型 loaded_model = tf.keras.models.load_model("my_model") # 初始化自定义模型实例 reconstructed_model = MyModel() # 构建模型(需要输入样本来初始化层权重) reconstructed_model(test_input) # 复制权重 reconstructed_model.set_weights(loaded_model.get_weights())
这样reconstructed_model就是你的自定义类实例,调用predict会执行自定义逻辑。
验证效果
使用方案一加载后,运行代码会输出:
Normal model... Hello from predict Hello from call Serialized model... Hello from predict Hello from call
说明自定义predict方法已生效。
内容的提问来源于stack exchange,提问作者Salman
相关产品推荐
相关产品推荐

