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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 16:39:03