如何在Keras Sequential模型中保存并加载自定义属性?
解决Keras Sequential模型自定义属性保存与加载问题
方法1:手动单独存储自定义属性(快速实现)
Keras默认不会保存手动通过setattr添加的自定义属性,最简单的方式是把属性单独存成文件,加载模型后再赋值回去:
保存代码
import json from keras.saving import save_model # 保存模型到指定路径 save_model(model, "./my_sequential_model") # 将自定义属性保存为json文件(list[str]适合用json存储) with open("./my_sequential_model/custom_attr.json", "w") as f: json.dump(model.custom_attr, f)
加载代码
import json from keras.saving import load_model # 加载模型 loaded_model = load_model("./my_sequential_model") # 读取并赋值自定义属性 with open("./my_sequential_model/custom_attr.json", "r") as f: loaded_model.custom_attr = json.load(f)
方法2:自定义Sequential子类(规范集成Keras序列化流程)
如果希望把自定义属性纳入Keras的原生保存/加载流程,可以继承Sequential类,重写序列化相关方法:
定义自定义Sequential类
from keras.models import Sequential class AttrSequential(Sequential): def __init__(self, custom_attr=None, **kwargs): super().__init__(**kwargs) # 初始化自定义属性,默认空列表 self.custom_attr = custom_attr or [] def get_config(self): # 先获取父类的配置,再加入自定义属性 config = super().get_config() config["custom_attr"] = self.custom_attr return config @classmethod def from_config(cls, config): # 从配置中取出自定义属性,再创建实例 custom_attr = config.pop("custom_attr") return cls(custom_attr=custom_attr, **config)
使用自定义类保存加载
from keras.saving import save_model, load_model # 创建自定义Sequential模型,添加层并设置属性 model = AttrSequential() model.add(...) # 添加你的网络层 model.custom_attr = ["one", "two", "three"] # 保存模型 save_model(model, "./my_attr_model") # 加载时指定自定义类 loaded_model = load_model("./my_attr_model", custom_objects={"AttrSequential": AttrSequential}) # 此时loaded_model.custom_attr已自动加载
注意事项
- 方法1适合快速解决问题,不需要修改原有模型结构;方法2更适合长期维护的项目,让属性管理更规范。
- 若自定义属性是复杂对象(而非list[str]),可以用
pickle替代json进行存储,但要注意pickle的安全性。
内容的提问来源于stack exchange,提问作者leqo
相关产品推荐
相关产品推荐

