仅用于预测的TensorFlow Keras自定义损失模型存加载方案问询
这个问题我之前也碰到过,针对你这种带参数lambda自定义损失的多输出模型,有几个靠谱的解决方案,不用重新指定损失就能顺利做预测:
方案1:只保存/加载模型权重(最稳妥的方式)
因为预测只需要模型的权重参数,完全不需要编译时的损失函数、优化器这些信息。你可以在训练完成后只保存权重,而不是整个模型:
训练脚本里:
# 训练完成后 model.save_weights("./model_weights.h5")
然后在预测脚本里,完全复刻训练时的模型结构(包括输入输出层、网络层的定义都要和训练时一模一样),然后加载权重直接预测:
# 先重建和训练时完全相同的模型结构 def build_model(): # 这里写你训练时的模型构建代码,比如: inputs = tf.keras.Input(shape=(your_input_shape,)) x = tf.keras.layers.Dense(64, activation='relu')(inputs) output1 = tf.keras.layers.Dense(1)(x) output2 = tf.keras.layers.Dense(1)(x) # ... 其他输出层 model = tf.keras.Model(inputs=inputs, outputs=[output1, output2, ...]) return model # 加载权重 model = build_model() model.load_weights("./model_weights.h5") # 直接预测,完全不需要编译 predictions = model.predict(your_test_data)
这个方法的好处是完全避开了损失函数序列化的问题,而且权重文件体积更小。唯一要注意的是必须保证预测脚本里的模型结构和训练时100%一致,包括层的参数、顺序、命名(如果有的话)。
方案2:把自定义损失封装成可序列化的Loss类
你的问题根源在于lambda函数无法被TensorFlow序列化保存,所以可以把带参数的损失函数改成继承tf.keras.losses.Loss的自定义类,这样模型保存时能把损失函数的信息一起序列化,加载时就不用加compile=False了。
比如针对你的weighted_mse_loss,可以写成这样:
class WeightedMSELoss(tf.keras.losses.Loss): def __init__(self, weight, name="weighted_mse_loss"): super().__init__(name=name) self.weight = weight # 把需要的参数传入初始化 def call(self, y_true, y_pred): # 这里实现你的损失计算逻辑 return util.weighted_mse_loss(y_true, y_pred, tf.square(self.weight)) # 同理处理pole_zero_loss class PoleZeroLoss(tf.keras.losses.Loss): def __init__(self, r_weight, w_weight, name="pole_zero_loss"): super().__init__(name=name) self.r_weight = r_weight self.w_weight = w_weight def call(self, y_true, y_pred): return util.pole_zero_loss(y_true, y_pred, self.r_weight, self.w_weight)
然后训练时用这些类的实例代替lambda:
losses = [ WeightedMSELoss(gain_weight), WeightedMSELoss(Rd_weight), PoleZeroLoss(r_weight, w_weight), PoleZeroLoss(r_weight, w_weight) ] model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=10E-4), loss=losses) # 训练完成后直接保存整个模型 model.save("./full_model.h5")
预测脚本里直接加载模型,不需要compile=False,可以直接predict:
model = tf.keras.models.load_model("./full_model.h5") predictions = model.predict(your_test_data)
这个方法更规范,适合需要保存完整模型结构的场景,而且后续加载模型时也能正常重新训练(如果需要的话)。
方案3:加载后手动设置占位损失(快速hack)
如果你不想改训练代码,也不想重建模型结构,可以在加载模型后,随便指定一个占位损失函数编译一下——因为预测时根本不会用到损失计算,只是模型实例需要有loss属性而已。
预测脚本里:
model = tf.keras.models.load_model(model_path, compile=False) # 随便指定一个损失,多输出的话对应每个输出给一个就行 model.compile(loss=["mse", "mse", "mse", "mse"]) # 现在可以正常predict了 predictions = model.predict(your_test_data)
这个方法最简单,但属于临时 workaround,不够优雅,而且如果后续需要重新训练模型的话,这个占位损失肯定不对,只适合纯预测的场景。
内容的提问来源于stack exchange,提问作者ruhig brauner

