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
关键修改点说明
- 将
l2reg设为__init__的可选参数并赋予默认值,确保序列化后能正确初始化实例。 get_config仅返回num_features和l2reg这两个基础数据类型参数,避免Tensor序列化问题。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
相关产品推荐
相关产品推荐

