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

Keras加载已保存神经网络模型时报缺少projection_dim参数错误如何解决

问题成因

  1. 自定义层PatchEncoder的__init__方法要求传入必填参数projection_dim,但你的get_config方法中只存储了重命名后的属性prj_dim。Keras加载模型时会将get_config返回的字典键值对直接传入__init__方法作为入参,键名和入参名不匹配就会导致projection_dim参数缺失。
  2. get_config方法中不应存储projection、position_embedding这类层实例对象,配置文件仅支持存储可序列化的基础参数(如数值、字符串等),层实例会在类初始化时自动创建,写入配置反而会引发序列化异常。
  3. 原有代码中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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 22:15:03