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
相关产品推荐
相关产品推荐

