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

TensorFlow恢复训练异常:无法从中断位置继续训练求助

问题排查与修复方案

运行单元格4而非3的思路是对的,但代码存在几个关键问题,导致训练无法从中断处恢复,具体分析和修复如下:

1. 核心问题:未指定initial_epoch参数

loaded_model.fit()默认从epoch 0开始训练,这是“从头开始”的直接原因。你必须明确指定从上次中断的epoch继续训练,这需要先记录已完成的epoch数,再在恢复时传入initial_epoch参数。

2. 补充训练进度记录

原代码的ModelCheckpoint只保存了模型,没有记录已完成的epoch数。建议修改为按epoch保存模型,同时记录训练进度:

# 修改单元格2的代码
batch_size = 32
import json

# 自定义回调:保存当前完成的epoch数
class SaveEpochCallback(tf.keras.callbacks.Callback):
    def on_epoch_end(self, epoch, logs=None):
        # 保存下一个要启动的epoch编号(比如完成第3个epoch后,下次从第4个开始)
        with open('current_epoch.json', 'w') as f:
            json.dump({'current_epoch': epoch + 1}, f)

cp_callback = tf.keras.callbacks.ModelCheckpoint(
    'best_model', 
    verbose=1, 
    save_weights_only=False,
    save_freq='epoch')  # 改为按完整epoch保存,适配你的长周期训练场景

# 单元格3修改为:
history = model.fit(train_generator, steps_per_epoch= train_generator.n // 16, 
                    validation_data=valid_generator,
                    validation_steps= valid_generator.n // 16, 
                    callbacks=[cp_callback, SaveEpochCallback()])

3. 修复恢复训练的代码

单元格4存在拼写错误,且缺少initial_epoch参数,修正如下:

import json
# 读取已完成的epoch数
try:
    with open('current_epoch.json', 'r') as f:
        epoch_data = json.load(f)
        initial_epoch = epoch_data['current_epoch']
except FileNotFoundError:
    initial_epoch = 0  # 无记录时从0开始

loaded_model = tf.keras.models.load_model('best_model')
new_history = loaded_model.fit(train_generator, 
                               steps_per_epoch= train_generator.n // 16, 
                               validation_data=valid_generator,
                               validation_steps= valid_generator.n // 16,  # 修正拼写错误:va.lidation_steps → validation_steps
                               callbacks=[cp_callback],
                               shuffle=False,
                               initial_epoch=initial_epoch)  # 指定从该epoch继续训练

4. 额外注意事项

  • 当save_weights_only=False时,模型会保存optimizer的状态(如Adam的动量、学习率衰减),恢复后无需重新编译,你的代码这部分是正确的。
  • 保持shuffle=False可以保证生成器的数据顺序一致,避免重复或遗漏数据,适合断点续训场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 22:01:36