Keras自定义AlexNet模型保存后加载时报缺少input_shape、num_classes参数错误
问题原因
这个报错的核心是:Keras的Sequential类自带的from_config反序列化逻辑,不会自动把get_config返回的参数传递给自定义子类的__init__方法,而是会先尝试无参数初始化你的AlexNet类,再逐层加载层配置,自然就会缺失你要求必传的input_shape和num_classes参数。
你之前硬写get_config的返回值、直接给参数加None默认值会报错,是因为初始化的时候没有正确传入参数,也没有兼容None值的逻辑。
修复方案
直接修改你的AlexNet类代码即可,完整修改后的代码如下:
# Define the AlexNet model class AlexNet(Sequential): # 给必填参数加默认值,兼容无参数初始化场景,默认值可按你的实际需求调整 def __init__(self, input_shape=(256,256,3), num_classes=3, **kwargs): # 先把参数存为实例属性,方便后续get_config调用 self.input_shape = input_shape self.num_classes = num_classes super().__init__(**kwargs) self.add(Conv2D(96, kernel_size=(11,11), strides= 4, padding= 'valid', activation= 'relu', input_shape= input_shape, kernel_initializer= 'he_normal')) self.add(BatchNormalization()) self.add(MaxPooling2D(pool_size=(3,3), strides= (2,2), padding= 'valid', data_format= None)) self.add(Conv2D(256, kernel_size=(5,5), strides= 1, padding= 'same', activation= 'relu', kernel_initializer= 'he_normal')) self.add(BatchNormalization()) self.add(MaxPooling2D(pool_size=(3,3), strides= (2,2), padding= 'valid', data_format= None)) self.add(Conv2D(384, kernel_size=(3,3), strides= 1, padding= 'same', activation= 'relu', kernel_initializer= 'he_normal')) self.add(BatchNormalization()) self.add(Conv2D(384, kernel_size=(3,3), strides= 1, padding= 'same', activation= 'relu', kernel_initializer= 'he_normal')) self.add(BatchNormalization()) self.add(Conv2D(256, kernel_size=(3,3), strides= 1, padding= 'same', activation= 'relu', kernel_initializer= 'he_normal')) self.add(BatchNormalization()) self.add(MaxPooling2D(pool_size=(3,3), strides= (2,2), padding= 'valid', data_format= None)) self.add(Flatten()) self.add(Dense(num_classes, activation= 'sigmoid')) self.compile(optimizer= tf.keras.optimizers.Adam(learning_rate=lr_schedule), loss='binary_crossentropy', metrics=['accuracy']) def get_config(self): # 从实例属性取参数,不要硬写死,适配不同参数初始化的模型 config = super().get_config() config.update({ "input_shape": self.input_shape, "num_classes": self.num_classes, }) return config # 手动实现from_config类方法,把配置参数传给__init__ @classmethod def from_config(cls, config): return cls(**config)
验证步骤
修改完类代码后,重新训练保存模型,再用原来的加载代码即可正常加载:
# Save the model model.save('./alexnet_model.hdf5') # Load the model alexnet_model = tf.keras.models.load_model('./alexnet_model.hdf5', custom_objects={'AlexNet': AlexNet})
- 如果是已经保存好的旧模型不想重新训练,在加载前用上面修改后的类定义即可直接加载,不需要重新训练。
- 注意如果你的
lr_schedule是自定义的学习率调度器,加载模型的时候也要把它加入custom_objects参数里,或者加载后重新编译模型传入优化器配置。
内容的提问来源于stack exchange,提问作者Yogesh Riyat
相关产品推荐
相关产品推荐

