TensorFlow 2.17自定义子类模型无法保存与加载问题求助
TensorFlow 2.17自定义子类模型无法保存与加载问题求助
嘿,我来帮你排查这个问题!你已经给自定义模型加了@tf.keras.utils.register_keras_serializable()装饰器,但子类化模型的序列化还需要额外的小处理——得重写get_config和from_config方法,让Keras能正确序列化、反序列化你的自定义模型类。
下面是修改后的可运行代码,你可以直接测试:
import numpy as np import tensorflow as tf # 设置随机种子,保证结果可复现 np.random.seed(42) tf.random.set_seed(42) x = np.random.random((1000, 32)) y = np.random.random((1000, 1)) @tf.keras.utils.register_keras_serializable() class MModel(tf.keras.Model): def __init__(self): super().__init__() self.dense1 = tf.keras.layers.Dense(10) self.dense2 = tf.keras.layers.Dense(1) def call(self, inputs): x = self.dense1(inputs) return self.dense2(x) # 重写get_config,返回模型初始化所需参数(这里__init__无额外参数,直接调用父类方法) def get_config(self): return super().get_config() # 重写from_config,从配置字典重建模型实例 @classmethod def from_config(cls, config): return cls(**config) model = MModel() model.compile( optimizer="adam", loss="mse", metrics=["mae"]) model.fit( x, y, epochs=5) model.save("save1.keras") model_ = tf.keras.models.load_model("save1.keras") print() print(f"{model.evaluate(x, y,verbose=0) = }") print(f"{model_.evaluate(x, y,verbose=0) = }") # 验证权重是否完全一致 for m, m_ in zip(model.weights, model_.weights): np.testing.assert_allclose(m.numpy(), m_.numpy(), rtol=1e-6) print("所有权重完全匹配!问题解决啦~")
问题原因解释:
子类化的tf.keras.Model默认不会自动序列化初始化逻辑,哪怕加了序列化装饰器,Keras也需要明确知道怎么从配置字典重建你的模型。重写这两个方法后,加载模型时就能正确初始化MModel实例,完整恢复训练后的权重。
另外我加了随机种子,这样你每次运行的结果都是可复现的,方便验证问题是否真的解决。如果还有问题,可以试试在compile前手动调用一次model(x),确保模型的输入形状被正确记录下来。
备注:内容来源于stack exchange,提问作者u2gilles
相关产品推荐
相关产品推荐

