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

仅用于预测的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 17:23:00