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

如何在TensorFlow 2.10中以历史最优损失为基准重启模型训练

解决TensorFlow续训时ModelCheckpoint以历史最优损失为基准的问题

当加载已保存的最优权重续训时,ModelCheckpoint默认会把基准损失设为+inf(mode='min'时),导致第一次验证损失只要不是无穷大就会触发保存,这不符合以历史最优为基准的需求。可以通过手动修改回调的_best属性实现目标,具体步骤如下:

步骤1:加载权重后评估验证集获取历史最优损失

加载权重后先编译模型,再在验证集上评估,得到当前权重对应的验证损失(即之前训练的最优损失)。如果是首次训练,则将基准初始化为inf。

步骤2:手动设置ModelCheckpoint的基准损失

创建ModelCheckpoint回调后,将其_best属性设置为刚才得到的验证损失,续训时就会以此为基准判断是否保存新的最优模型。

修改后的完整代码

model = generate_model(lstm_size, conv_size, num_variables, num_timesteps)
if os.path.isfile(checkpoint_filepath):
    model.load_weights(checkpoint_filepath)
    # 编译模型以支持评估
    model.compile(
        optimizer=tf.keras.optimizers.Adam(learning_rate=learning_rate),
        loss=focal_loss(),
        metrics=[tf.keras.metrics.Recall(), tf.keras.metrics.Precision()]
    )
    # 在验证集上评估,获取当前最优权重对应的val_loss
    val_loss, _, _ = model.evaluate(test_dataset, verbose=0)
else:
    # 首次训练,初始基准设为无穷大
    model.compile(
        optimizer=tf.keras.optimizers.Adam(learning_rate=learning_rate),
        loss=focal_loss(),
        metrics=[tf.keras.metrics.Recall(), tf.keras.metrics.Precision()]
    )
    val_loss = float('inf')

# 初始化模型保存回调
save_model_callback = tf.keras.callbacks.ModelCheckpoint(
    filepath=checkpoint_filepath,
    save_best_only=True,
    monitor='val_loss',
    mode='min',
    save_weights_only=True,
    verbose=1
)

# 手动设置回调的基准损失为历史最优值
save_model_callback._best = val_loss

# 启动训练
model.fit(
    train_dataset,
    epochs=num_epochs,
    validation_data=test_dataset,
    callbacks=[save_model_callback]
)

注意事项

  • 必须先编译模型才能调用evaluate,编译逻辑要和权重加载流程对应处理。
  • 若之前保存的是完整模型而非仅权重,也可以加载模型后从model.history提取历史最优损失,但仅保存权重时,直接评估验证集是最可靠的方式。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 00:45:21