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

含自定义层的Keras模型加载时报RandomNormal错误如何解决?

问题原因

  • 你在__init__方法中重复调用了两次父类构造函数,第一次传入的name参数会被第二次调用覆盖,属于逻辑缺陷
  • 你在get_config方法中直接存入了初始化器实例self.init,Keras序列化时仅会将该实例转换为对应的类名字符串RandomNormal,反序列化时无法自动还原该初始化器实例,因此抛出找不到配置项的错误

修复方案

方案一(适合固定初始化器的场景,最简单)

因为你的初始化器固定为normal,不需要将其放入序列化配置中,仅需保留必要的自定义参数即可,修改后的自定义层代码如下:

from keras.engine.base_layer import Layer
from keras import initializers

class AttentionLayer(Layer):
  def __init__(self, attention_dim, **kwargs):
    # 移除重复的父类构造调用
    super(AttentionLayer, self).__init__(**kwargs)
    self.init = initializers.get("normal")
    self.supports_masking = True
    self.attention_dim = attention_dim

  def get_config(self):
    # 仅序列化需要自定义传入的参数
    config = {
      "attention_dim": self.attention_dim,
    }
    base_config = super(AttentionLayer, self).get_config()
    return dict(list(base_config.items()) + list(config.items()))

方案二(适合初始化器可配置的场景)

如果需要支持不同的初始化器传入,可通过Keras自带的序列化/反序列化工具处理初始化器参数:

from keras.engine.base_layer import Layer
from keras import initializers
from keras.utils import serialize_keras_object, deserialize_keras_object

class AttentionLayer(Layer):
  def __init__(self, attention_dim, init="normal", **kwargs):
    super(AttentionLayer, self).__init__(**kwargs)
    self.init = initializers.get(init)
    self.supports_masking = True
    self.attention_dim = attention_dim

  def get_config(self):
    config = {
      "init": serialize_keras_object(self.init),
      "attention_dim": self.attention_dim,
    }
    base_config = super(AttentionLayer, self).get_config()
    return dict(list(base_config.items()) + list(config.items()))
  
  @classmethod
  def from_config(cls, config):
    config["init"] = deserialize_keras_object(config["init"], module_objects=globals(), custom_objects={})
    return cls(**config)

加载模型

修改自定义层代码后,重新训练保存模型,再用原加载代码即可正常加载:

model = keras.models.load_model("model.h5", 
                custom_objects={"AttentionLayer": AttentionLayer})

注意:训练和加载模型时要保证Keras相关模块的导入路径完全一致,例如统一使用keras或统一使用tensorflow.keras开头的导入,否则也可能出现序列化不匹配的问题


内容的提问来源于stack exchange,提问作者Rishik Mani

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 02:06:00