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

序列化重载Keras模型后自定义predict方法失效问题求助

解决Keras序列化模型后自定义predict方法失效的问题

问题原因

Keras的SavedModel序列化机制默认仅保存模型的层结构、权重和训练配置,不会完整保留自定义子类模型的实例方法(比如你重载的predict)。重新加载后得到的是tf.keras.Model的通用实例,而非你定义的MyModel类实例,因此自定义的predict方法会被忽略,仅执行默认的predict逻辑(内部调用call方法)。

解决方案

方案一:加载时指定自定义类

要让加载后的模型仍是你的自定义类实例,需要在加载时通过custom_objects参数传入你的模型类,同时建议实现get_config方法以确保配置正确序列化:

  1. 完善自定义模型类:
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()
  1. 加载模型时传入自定义类:
# 保存模型(使用默认的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 16:05:31