如何用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的定义已存在:
- 直接导入类定义:在加载模型的代码中,提前导入或定义
GreatClass。 - 指定自定义对象加载:如果类在其他模块,可通过
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
相关产品推荐
相关产品推荐

