训练保存正常但加载自定义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)
验证步骤
- 使用修改后的代码重新训练模型(或直接重新实例化模型后加载已训练权重,再保存)
- 执行
model.save("vedio_model.keras")保存模型 - 执行
loaded_model = tf.keras.models.load_model("vedio_model.keras")加载模型,此时应无报错
内容的提问来源于stack exchange,提问作者Shivanand Garg
相关产品推荐
相关产品推荐

