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

如何用model.save()保存TensorFlow自定义模型的方法与属性?

问题描述

我创建了一个带有自定义方法new_method和自定义属性testing的Keras模型,尝试用model.save()保存,但加载后无法访问这些自定义内容。

实现代码如下:

@tf.keras.utils.register_keras_serializable()
class GreatClass(tf.keras.Model):
    def __init__(self, **kwargs):
        super().__init__(**kwargs)
        self.testing = 3424
        self.dense = tf.keras.layers.Dense(100)
        
    def get_config(self):
        config = super().get_config()
        config['testing'] = self.testing
        config['dense'] = self.dense
        return config
    
    def new_method(self):
        print('hello world')
    
    def call(self, inputs):
        return self.dense(inputs)

保存模型的代码:

model = GreatClass()
model.compile()

array = np.array([100,10])
model.predict(array)

model.save('testing')

加载后调用自定义方法和属性时报错:

loaded_model = tf.keras.models.load_model("testing")
loaded_model.new_method()

AttributeError: 'Custom>GreatClass' object has no attribute 'new_method'

loaded_model.testing

AttributeError: 'Custom>GreatClass' object has no attribute 'testing'

请问能否通过model.save()保存自定义方法与属性?

解决方案

可以保存自定义属性,但自定义方法无法通过model.save()直接序列化,具体原因和修复方案如下:

修复自定义属性的问题

你的get_config()中有冗余代码:Keras会自动处理层对象的序列化,不需要手动将self.dense加入配置。修正后就能正确恢复testing属性:

@tf.keras.utils.register_keras_serializable()
class GreatClass(tf.keras.Model):
    def __init__(self, **kwargs):
        super().__init__(**kwargs)
        self.testing = 3424
        self.dense = tf.keras.layers.Dense(100)
        
    def get_config(self):
        config = super().get_config()
        # 仅需加入自定义非层属性
        config['testing'] = self.testing
        return config
    
    def new_method(self):
        print('hello world')
    
    def call(self, inputs):
        return self.dense(inputs)

关键前提:加载模型时,GreatClass的完整定义必须在当前环境可访问,否则Keras会生成一个代理类Custom>GreatClass,无法恢复自定义属性。

恢复自定义方法的问题

Keras的model.save()只序列化模型的结构、权重和配置,不会保存自定义方法的代码逻辑。要恢复方法,只需确保加载模型时,GreatClass的定义已存在:

  1. 直接导入类定义:在加载模型的代码中,提前导入或定义GreatClass。
  2. 指定自定义对象加载:如果类在其他模块,可通过custom_objects参数声明:
loaded_model = tf.keras.models.load_model("testing", custom_objects={"GreatClass": GreatClass})

完成后,加载的模型会是GreatClass的真实实例,可正常调用new_method。

注意事项

  • 不要手动在get_config()中添加层对象,避免序列化冲突。
  • 自定义类必须用@tf.keras.utils.register_keras_serializable()装饰,或在加载时通过custom_objects指定,否则会生成代理类导致无法访问自定义内容。

内容的提问来源于stack exchange,提问作者Ottpocket

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 05:03:27