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

TensorFlow 2训练中ReduceLROnPlateau回调重载模型权重相关问询

TensorFlow 2 实现ReduceLROnPlateau触发时重载最优权重

完全可以实现,通过自定义继承ReduceLROnPlateau的回调类即可满足需求,具体操作步骤如下:


步骤1:配置ModelCheckpoint回调

首先确保你配置的ModelCheckpoint已开启仅保存最优权重的配置,示例代码:

from tensorflow.keras.callbacks import ModelCheckpoint

checkpoint_cb = ModelCheckpoint(
    filepath="./best_weights.h5", # 权重存储路径,可自定义
    monitor="val_loss", # 和后续ReduceLROnPlateau监控的指标保持一致
    save_best_only=True,
    save_weights_only=True, # 仅保存权重,重载效率更高
    verbose=1
)

步骤2:自定义ReduceLROnPlateau回调

继承原生的ReduceLROnPlateau类,重写reduce_lr方法,在学习率调整完成后插入重载最优权重的逻辑:

from tensorflow.keras.callbacks import ReduceLROnPlateau

class ReduceLROnPlateauWithReloadBest(ReduceLROnPlateau):
    def reduce_lr(self, epoch):
        # 调用父类原生的学习率调整逻辑
        super().reduce_lr(epoch)
        # 调整学习率后重载当前最优权重
        print("触发学习率调整,重载验证集最优权重...")
        self.model.load_weights("./best_weights.h5")

步骤3:训练时传入回调

初始化自定义的学习率调整回调,和ModelCheckpoint一起传入model.fit的回调列表即可:

# 初始化自定义学习率调整回调
reduce_lr_cb = ReduceLROnPlateauWithReloadBest(
    monitor="val_loss", # 和ModelCheckpoint的监控指标保持一致
    factor=0.5, # 学习率衰减系数,可自定义
    patience=5, # 指标无进步多少轮后触发调整,可自定义
    min_lr=1e-7,
    verbose=1
)

# 训练时传入两个回调
model.fit(
    train_dataset,
    validation_data=val_dataset,
    epochs=100,
    callbacks=[checkpoint_cb, reduce_lr_cb]
)

注意事项

  • 如果你的ModelCheckpoint配置的是保存完整模型而非仅权重,需要将重载逻辑替换为对应加载完整模型的代码
  • 两个回调的monitor参数必须保持一致,避免监控指标不匹配导致加载的权重不符合预期
  • 如果使用了带动态变量的权重存储路径(比如带epoch、loss变量),要保证自定义回调中的加载路径和ModelCheckpoint的存储路径完全匹配

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 15:06:06