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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 04:25:02