为何Keras model_from_json()加载模型配置返回字符串而非模型实例?
问题原因与解决方法
核心原因
model_from_json() 仅能识别Keras内置的模型类(如Sequential、Functional模型)。你的配置文件中指定的是自定义模型MyModel,但该类未注册到Keras的序列化注册表中,导致函数无法实例化模型,只能返回配置里的class_name字符串。
解决步骤
确保自定义模型类实现规范
你的MyModel必须继承tf.keras.Model基类,并且正确实现get_config()方法(配置文件中的config字段就是该方法的返回值),示例如下:import tensorflow as tf class MyModel(tf.keras.Model): def __init__(self, hidden_layers=2, units=64, activation="swish", **kwargs): super().__init__(**kwargs) self.hidden_layers = hidden_layers self.units = units self.activation = activation # 此处添加你的模型层定义逻辑 def get_config(self): config = super().get_config() config.update({ "hidden_layers": self.hidden_layers, "units": self.units, "activation": self.activation }) return config注册自定义模型类
在调用model_from_json()前,将MyModel注册到Keras的自定义对象注册表中,有两种方式:- 方式一:手动添加到自定义对象字典
from tensorflow.keras.models import model_from_json tf.keras.utils.get_custom_objects()["MyModel"] = MyModel # 加载模型 with open("your_model_config.json", "r") as f: model_json = f.read() model = model_from_json(model_json) - 方式二:使用装饰器自动注册
在定义MyModel时添加@tf.keras.utils.register_keras_serializable()装饰器,无需手动注册:@tf.keras.utils.register_keras_serializable() class MyModel(tf.keras.Model): # 类定义同上
- 方式一:手动添加到自定义对象字典
补充说明
Keras的序列化系统依赖注册表匹配模型类名与构造方法,未注册的自定义类无法被识别,因此只会返回配置中的类名字符串,而非实例化的模型对象。
内容的提问来源于stack exchange,提问作者the man
相关产品推荐
相关产品推荐

