Keras加载已保存神经网络模型时报缺少projection_dim参数错误如何解决
问题成因
- 自定义层
PatchEncoder的__init__方法要求传入必填参数projection_dim,但你的get_config方法中只存储了重命名后的属性prj_dim。Keras加载模型时会将get_config返回的字典键值对直接传入__init__方法作为入参,键名和入参名不匹配就会导致projection_dim参数缺失。 get_config方法中不应存储projection、position_embedding这类层实例对象,配置文件仅支持存储可序列化的基础参数(如数值、字符串等),层实例会在类初始化时自动创建,写入配置反而会引发序列化异常。- 原有代码中
super().__init__()没有传入**kwargs,会导致父类的默认参数无法正常初始化,也可能引发序列化问题。
修复方案
第一步:修正PatchEncoder类定义
class PatchEncoder(layers.Layer): def __init__(self, num_patches, projection_dim, **kwargs): # 给父类传入kwargs参数 super(PatchEncoder, self).__init__(**kwargs) self.num_patches = num_patches self.prj_dim = projection_dim self.projection = layers.Dense(units=projection_dim) self.position_embedding = layers.Embedding( input_dim=num_patches, output_dim=projection_dim ) def call(self, patch): positions = tf.range(start=0, limit=self.num_patches, delta=1) encoded = self.projection(patch) + self.position_embedding(positions) return encoded def get_config(self): config = super().get_config().copy() config.update({ 'num_patches': self.num_patches, # 键名和__init__的入参名保持一致,仅存储基础数值参数 'projection_dim': self.prj_dim }) return config
第二步:正确加载模型
加载模型时需要通过custom_objects参数告知Keras你的自定义层定义:
model = tf.keras.models.load_model("你的模型存储路径", custom_objects={"PatchEncoder": PatchEncoder})
如果不想每次加载都传custom_objects,可以在定义类时加上Keras序列化装饰器:
@tf.keras.utils.register_keras_serializable() class PatchEncoder(layers.Layer): # 剩余类代码和上方一致
内容的提问来源于stack exchange,提问作者hoai lan
相关产品推荐
相关产品推荐

