Keras自定义SkipCon层模型保存后加载报找不到匹配函数错误
报错核心原因
报错由三个代码问题共同导致:
get_config方法实现完全不符合规范:该方法需要返回自定义层__init__接口的所有入参,才能在加载时重建层实例。你当前的实现仅返回了{"Z": self.activation},既缺失了size/reduce/deep/skip_when等必填初始化参数,还将激活函数对象赋值给了和参数不匹配的键Z,重载时无法正确初始化层实例。- 激活函数未做序列化处理:你在
__init__中通过keras.activations.get(activation)将字符串格式的激活函数名转成了函数对象,这类对象无法直接被序列化存储,加载时无法识别。 - 内部子层未被Keras正确追踪:你将层内部嵌套的Dense、BatchNormalization子层存在普通Python列表
self.main_layers、self.skip_layers中,Keras不会自动追踪普通列表内的子层权重与结构,保存模型时这部分内容会丢失,加载时结构和权重不匹配直接触发函数匹配错误。
额外说明:你的模型构建代码存在一处结构笔误,第二个处理both_layer的SkipCon层错误传入了原始combined张量作为输入,没有接上前序Dropout层的输出,会导致模型结构不符合设计预期,该问题不直接触发加载报错,但会影响模型效果。
修复方案
按照以下步骤调整代码即可正常保存、重载模型:
- 修正自定义层的子层存储方式,将普通Python列表替换为
keras.layers.LayerList,让Keras可以自动追踪内部子层的权重、结构。 - 重写符合规范的
get_config方法:先调用父类的get_config获取通用层参数,再将当前层所有初始化参数写入配置,激活函数通过keras.activations.serialize()转成可序列化的字符串标识,不要直接存储函数对象。 - 修正模型构建时的张量传参笔误。
修正后的SkipCon层代码如下:
class SkipCon(keras.layers.Layer): def __init__(self, size, reduce = True, deep = 3, skip_when=0, activation="relu", **kwargs): super().__init__(**kwargs) # 存储初始化参数,供序列化使用 self.size = size self.reduce = reduce self.deep = deep self.skip_when = skip_when self.activation = keras.activations.get(activation) # 用LayerList替换普通列表,追踪内部子层 self.main_layers = keras.layers.LayerList() current_size = size for _ in range(deep): self.main_layers.append( keras.layers.Dense(current_size, activation=activation, use_bias=True) ) self.main_layers.append(keras.layers.BatchNormalization()) if reduce: current_size = current_size // 2 self.final_main_size = current_size self.skip_layers = keras.layers.LayerList() if skip_when > 0: skip_size = self.final_main_size * 2 if reduce else self.final_main_size self.skip_layers.append( keras.layers.Dense(skip_size, activation=activation, use_bias=True) ) self.skip_layers.append(keras.layers.BatchNormalization()) def call(self, inputs): Z = inputs for layer in self.main_layers: Z = layer(Z) if not self.skip_when: return self.activation(Z) skip_Z = inputs for layer in self.skip_layers: skip_Z = layer(skip_Z) return self.activation(Z + skip_Z) def get_config(self): config = super().get_config() config.update({ "size": self.size, "reduce": self.reduce, "deep": self.deep, "skip_when": self.skip_when, "activation": keras.activations.serialize(self.activation) }) return config
需要修正的模型构建代码段如下:
both_layer = SkipCon(size = 128, deep = 2, reduce = False, skip_when=1, activation="relu")(combined) both_layer = keras.layers.Dropout(0.5)(both_layer) # 原代码此处错误传入combined,改为传入上一层输出both_layer both_layer = SkipCon(size = 64, deep = 2, reduce = False, skip_when=1, activation="relu")(both_layer) both_layer = keras.layers.Dropout(0.5)(both_layer) both_layer = SkipCon(size = 16, deep = 2, reduce = False, skip_when=0, activation="relu")(both_layer)
调整完成后,原有的保存、加载逻辑不需要修改,即可正常重载模型。如果想要省略加载时传入custom_objects的步骤,可以在SkipCon类前加上@keras.utils.register_keras_serializable()装饰器,Keras会自动识别该自定义层。
内容的提问来源于stack exchange,提问作者Ibikunoluwa Adegbite
相关产品推荐
相关产品推荐

