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

Keras加载含自定义OrthogonalRegularizer的PointNet模型失败求助

解决Keras自定义OrthogonalRegularizer序列化异常导致模型加载失败的问题

问题根源

你的问题出在自定义正则化器OrthogonalRegularizer的序列化逻辑上:

  • get_config方法中错误地将self.eye(Tensor对象)加入了配置字典,而Tensor无法被Keras正确序列化和反序列化。
  • 加载模型时,Keras会尝试用包含Tensor的配置初始化正则化器,但你的__init__方法并不接受eye参数,且Tensor无法直接作为初始化参数传入,从而触发object.__init__() takes exactly one argument错误。

解决方案

修改OrthogonalRegularizer类,调整get_config和__init__逻辑,只保存可序列化的参数,eye在初始化时通过num_features重新生成:

@keras.saving.register_keras_serializable('OrthogonalRegularizer')
class OrthogonalRegularizer(keras.regularizers.Regularizer):

    def __init__(self, num_features, l2reg=0.001, **kwargs):
        super().__init__(**kwargs)
        self.num_features = num_features
        self.l2reg = l2reg
        # 初始化时通过num_features生成eye,不再保存到配置
        self.eye = tf.eye(num_features)

    def __call__(self, x):
        x = tf.reshape(x, (-1, self.num_features, self.num_features))
        xxt = tf.tensordot(x, x, axes=(2, 2))
        xxt = tf.reshape(xxt, (-1, self.num_features, self.num_features))
        return tf.math.reduce_sum(self.l2reg * tf.square(xxt - self.eye))

    def get_config(self):
        # 只保存可序列化的num_features和l2reg,去掉eye
        config = super().get_config()
        config.update({
            "num_features": self.num_features,
            "l2reg": self.l2reg
        })
        return config

关键修改点说明

  1. 将l2reg设为__init__的可选参数并赋予默认值,确保序列化后能正确初始化实例。
  2. get_config仅返回num_features和l2reg这两个基础数据类型参数,避免Tensor序列化问题。
  3. eye在__init__中通过num_features重新生成,无需从配置中加载,保证实例初始化逻辑一致。

验证修改

修改完成后,使用原有的保存加载代码即可正常工作:

model.save('my_model.h5')
model = keras.models.load_model('my_model.h5', custom_objects={'OrthogonalRegularizer': OrthogonalRegularizer})

内容的提问来源于stack exchange,提问作者Musabbir Arrafi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 19:24:53