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

训练保存正常但加载自定义Keras Model时报错:Could not locate class 'Model'

解决自定义Keras Model加载时的"Could not locate class 'Model'"错误

问题场景

自定义了继承自tf.keras.Model的视频二分类模型,训练完成后通过model.save("vedio_model.keras")保存模型,执行tf.keras.models.load_model("vedio_model.keras")加载时触发TypeError,提示无法找到'Model'类,要求使用@keras.saving.register_keras_serializable()装饰自定义类。训练与保存阶段无报错,仅加载时出现问题。

原模型代码:

import tensorflow as tf
from tensorflow.keras.applications import ResNet50
from tensorflow.keras.layers import LSTM, Dense, GlobalAveragePooling2D, Dropout
from tensorflow.keras.activations import relu

class Model(tf.keras.Model):
    def __init__(self, num_classes, latent_dim=2048, lstm_layers=1, hidden_dim=2048, bidirectional=False):
        super(Model, self).__init__()
        self.base_model = tf.keras.Sequential([
            ResNet50(include_top=False, weights=None, input_shape=(None, None, 3)),
        ], name='resnet50_base')
        self.base_model.layers[0].name = 'resnet50_base'  # 设置ResNet50的唯一名称
        self.base_model.layers[0].trainable = False  # 冻结ResNet50权重
        self.lstm = LSTM(hidden_dim, return_sequences=True, return_state=True, name='lstm_layer')
        self.relu = relu
        self.dp = Dropout(0.4, name='dropout_layer')
        self.linear1 = Dense(num_classes, name='output_layer')
        self.avgpool = GlobalAveragePooling2D(name='global_avg_pooling')
        self.built = True  # 标记模型已构建

    def call(self, inputs):
        batch_size, seq_length, h, w, c = inputs.shape.as_list()
        inputs = tf.reshape(inputs, (batch_size * seq_length, h, w, c))
        fmap = self.base_model(inputs)
        x = self.avgpool(fmap)
        x = tf.reshape(x, (batch_size, seq_length, -1))
        x_lstm, _, _ = self.lstm(x)
        x = tf.reduce_mean(x_lstm, axis=1)
        x = self.dp(x)
        x = self.linear1(x)
        return fmap, x

model = Model(num_classes=2)  # 二分类:真实/伪造

报错信息:

TypeError                                 Traceback (most recent call last)
Cell In[34], line 1
----> 1 loaded_model = tf.keras.models.load_model("vedio_model.keras")

File /usr/local/lib/python3.11/dist-packages/keras/src/saving/saving_api.py:176, in load_model(filepath, custom_objects, compile, safe_mode)
173         is_keras_zip = True
175 if is_keras_zip:
--> 176     return saving_lib.load_model(
177         filepath,
178         custom_objects=custom_objects,
179         compile=compile,
180         safe_mode=safe_mode,
181     )
182 if str(filepath).endswith((".h5", ".hdf5")):
183     return legacy_h5_format.load_model_from_hdf5(filepath)

File /usr/local/lib/python3.11/dist-packages/keras/src/saving/saving_lib.py:155, in load_model(filepath, custom_objects, compile, safe_mode)
153 # Construct the model from the configuration file in the archive.
154 with ObjectSharingScope():
--> 155     model = deserialize_keras_object(
156         config_dict, custom_objects, safe_mode=safe_mode
157     )
159 all_filenames = zf.namelist()
160 if _VARS_FNAME + ".h5" in all_filenames:

File /usr/local/lib/python3.11/dist-packages/keras/src/saving/serialization_lib.py:687, in deserialize_keras_object(config, custom_objects, safe_mode, **kwargs)
684     if obj is not None:
685         return obj
--> 687 cls = _retrieve_class_or_fn(
688     class_name,
689     registered_name,
690     module,
691     obj_type="class",
692     full_config=config,
693     custom_objects=custom_objects,
694 )
696 if isinstance(cls, types.FunctionType):
697     return cls

File /usr/local/lib/python3.11/dist-packages/keras/src/saving/serialization_lib.py:805, in _retrieve_class_or_fn(name, registered_name, module, obj_type, full_config, custom_objects)
    802     if obj is not None:
    803         return obj
--> 805 raise TypeError(
806     f"Could not locate {obj_type} '{name}'. "
807     "Make sure custom classes are decorated with "
808     "`@keras.saving.register_keras_serializable()` "
809     f"Full object config: {full_config}"
810 )

TypeError: Could not locate class 'Model'. Make sure custom classes are decorated with      `@keras.saving.register_keras_serializable()`. Full object config: {'module': None, 'class_name':    'Model', 'config': {'trainable': True, 'dtype': 'float32'}, 'registered_name': 'Model'}

原因分析

Keras在序列化自定义模型时,需要将类的定义信息注册到Keras的序列化系统中。如果自定义类没有通过@keras.saving.register_keras_serializable()装饰,保存的模型文件中只会记录类名,加载时无法找到对应的类定义,从而抛出找不到类的错误。

解决方案

给自定义的Model类添加@keras.saving.register_keras_serializable()装饰器,同时移除手动设置的self.built = True(Keras会自动处理模型构建流程,手动设置可能引发潜在问题)。

修改后的完整代码:

import tensorflow as tf
from tensorflow.keras.applications import ResNet50
from tensorflow.keras.layers import LSTM, Dense, GlobalAveragePooling2D, Dropout
from tensorflow.keras.activations import relu

# 添加序列化装饰器
@tf.keras.saving.register_keras_serializable()
class Model(tf.keras.Model):
    def __init__(self, num_classes, latent_dim=2048, lstm_layers=1, hidden_dim=2048, bidirectional=False):
        super(Model, self).__init__()
        self.base_model = tf.keras.Sequential([
            ResNet50(include_top=False, weights=None, input_shape=(None, None, 3)),
        ], name='resnet50_base')
        self.base_model.layers[0].name = 'resnet50_base'
        self.base_model.layers[0].trainable = False
        self.lstm = LSTM(hidden_dim, return_sequences=True, return_state=True, name='lstm_layer')
        self.relu = relu
        self.dp = Dropout(0.4, name='dropout_layer')
        self.linear1 = Dense(num_classes, name='output_layer')
        self.avgpool = GlobalAveragePooling2D(name='global_avg_pooling')
        # 移除手动设置的self.built = True

    def call(self, inputs):
        batch_size, seq_length, h, w, c = inputs.shape.as_list()
        inputs = tf.reshape(inputs, (batch_size * seq_length, h, w, c))
        fmap = self.base_model(inputs)
        x = self.avgpool(fmap)
        x = tf.reshape(x, (batch_size, seq_length, -1))
        x_lstm, _, _ = self.lstm(x)
        x = tf.reduce_mean(x_lstm, axis=1)
        x = self.dp(x)
        x = self.linear1(x)
        return fmap, x

model = Model(num_classes=2)

验证步骤

  1. 使用修改后的代码重新训练模型(或直接重新实例化模型后加载已训练权重,再保存)
  2. 执行model.save("vedio_model.keras")保存模型
  3. 执行loaded_model = tf.keras.models.load_model("vedio_model.keras")加载模型,此时应无报错

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 21:30:55