【问题标题】:How to load the Keras model with custom layers from .h5 file correctly?如何正确加载带有 .h5 文件中的自定义层的 Keras 模型?
【发布时间】:2019-12-21 17:16:54
【问题描述】:

我用自定义层构建了一个 Keras 模型,并通过回调 ModelCheckPoint 将其保存到 .h5 文件中。 当我在训练后尝试加载此模型时,出现以下错误消息:

__init__() missing 1 required positional argument: 'pool_size'

这是自定义层的定义及其__init__方法:

class MyMeanPooling(Layer):
    def __init__(self, pool_size, axis=1, **kwargs):
        self.supports_masking = True
        self.pool_size = pool_size
        self.axis = axis
        self.y_shape = None
        self.y_mask = None
        super(MyMeanPooling, self).__init__(**kwargs)

这就是我将这一层添加到我的模型的方式:

x = MyMeanPooling(globalvars.pool_size)(x)

这是我加载模型的方式:

from keras.models import load_model

model = load_model(model_path, custom_objects={'MyMeanPooling': MyMeanPooling})

这些是完整的错误消息:

Traceback (most recent call last):
  File "D:/My Projects/Attention_BLSTM/script3.py", line 9, in <module>
    model = load_model(model_path, custom_objects={'MyMeanPooling': MyMeanPooling})
  File "D:\ProgramData\Anaconda3\envs\tf\lib\site-packages\keras\engine\saving.py", line 419, in load_model
    model = _deserialize_model(f, custom_objects, compile)
  File "D:\ProgramData\Anaconda3\envs\tf\lib\site-packages\keras\engine\saving.py", line 225, in _deserialize_model
    model = model_from_config(model_config, custom_objects=custom_objects)
  File "D:\ProgramData\Anaconda3\envs\tf\lib\site-packages\keras\engine\saving.py", line 458, in model_from_config
    return deserialize(config, custom_objects=custom_objects)
  File "D:\ProgramData\Anaconda3\envs\tf\lib\site-packages\keras\layers\__init__.py", line 55, in deserialize
    printable_module_name='layer')
  File "D:\ProgramData\Anaconda3\envs\tf\lib\site-packages\keras\utils\generic_utils.py", line 145, in deserialize_keras_object
    list(custom_objects.items())))
  File "D:\ProgramData\Anaconda3\envs\tf\lib\site-packages\keras\engine\network.py", line 1022, in from_config
    process_layer(layer_data)
  File "D:\ProgramData\Anaconda3\envs\tf\lib\site-packages\keras\engine\network.py", line 1008, in process_layer
    custom_objects=custom_objects)
  File "D:\ProgramData\Anaconda3\envs\tf\lib\site-packages\keras\layers\__init__.py", line 55, in deserialize
    printable_module_name='layer')
  File "D:\ProgramData\Anaconda3\envs\tf\lib\site-packages\keras\utils\generic_utils.py", line 147, in deserialize_keras_object
    return cls.from_config(config['config'])
  File "D:\ProgramData\Anaconda3\envs\tf\lib\site-packages\keras\engine\base_layer.py", line 1109, in from_config
    return cls(**config)
TypeError: __init__() missing 1 required positional argument: 'pool_size'

【问题讨论】:

  • 您在 Layer 子类中实现了哪些方法?
  • 这是因为 keras 正在调用层的构造函数,但它需要 1 个位置参数,即“pool_size”。 (keras没有提供这个参数)

标签: python-3.x keras keras-layer


【解决方案1】:

其实我不认为你可以加载这个模型。

最可能的问题是您没有在您的层中实现get_config() 方法。此方法返回应保存的配置值字典:

def get_config(self):
    config = {'pool_size': self.pool_size,
              'axis': self.axis}
    base_config = super(MyMeanPooling, self).get_config()
    return dict(list(base_config.items()) + list(config.items()))

将此方法添加到您的层后,您必须重新训练模型,因为之前保存的模型没有保存该层的配置。这就是您无法加载它的原因,它需要在进行此更改后重新训练。

【讨论】:

  • @waleema 当然,不客气,但您应该根据问题是否解决了您的问题或对您有用来投票和/或接受问题。
  • 我投了你一票,但我的投票似乎并没有改变公开显示的帖子分数,因为我的声誉低于15,但系统通知我我的投票已记录。再次感谢!
  • @waleema 您所做的是接受答案,这与赞成/反对投票是分开的,但这没关系,我们只是想让您知道系统在 SO 中的工作原理:)
  • 是的,我接受了你的回答,我也投票给了你。
【解决方案2】:

来自“LiamHe 于 2017 年 9 月 27 日发表评论”关于以下问题的回答:https://github.com/keras-team/keras/issues/4871

我今天遇到了同样的问题:** TypeError: init() missing 1 required positional arguments**。这是我解决问题的方法:(Keras 2.0.2)

  1. 为层的位置参数提供一些默认值
  2. 用类似的东西覆盖 get_config 函数到层
def get_config(self):
    config = super().get_config()
    config['pool_size'] = # say self._pool_size  if you store the argument in __init__
    return config
  1. 在加载模型时将图层类添加到 custom_objects。

【讨论】:

  • 非常感谢!你的回答很有帮助。
【解决方案3】:

如果您没有足够的时间以 Matias Valdenegro 的求解方式重新训练模型。您可以在类 MyMeanPooling 中设置 pool_size 的默认值,如下面的代码。注意 pool_size 的值应该和训练模型时的值一致。然后就可以加载模型了。

class MyMeanPooling(Layer):
    def __init__(self, pool_size, axis=1, **kwargs):
        self.supports_masking = True
        self.pool_size = 2  # The value should be consistent with the value while training the model
        self.axis = axis
        self.y_shape = None
        self.y_mask = None
        super(MyMeanPooling, self).__init__(**kwargs)

参考:https://www.jianshu.com/p/e97112c34e43

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2021-01-05
    • 2019-08-03
    • 1970-01-01
    • 2022-12-17
    • 1970-01-01
    • 1970-01-01
    • 2018-07-21
    • 1970-01-01
    相关资源
    最近更新 更多