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

Colab中U-Net模型中断后无法续训的问题求助

U-Net续训失败:从断点重启后从头训练的解决方法

问题现象

在Google Colab训练单目标检测U-Net时,训练中断后尝试从检查点续训,但每次都会从第1个epoch重新开始,而非从之前完成的轮次(比如完成30轮后从第31轮)继续。

核心原因

  1. 未记录训练进度:仅加载了模型权重,但未保存/读取之前已完成的epoch数,model.fit()默认从epoch 0开始计数
  2. 检查点策略限制:原代码使用save_best_only=True仅保存验证集表现最优的权重,未保留每个训练轮次的进度记录,无法确定中断时的训练轮次
  3. 优化器状态未恢复:仅加载模型权重时,优化器的动量、学习率调度等状态未被恢复,会导致续训的训练状态不连贯

解决方案

1. 修改检查点保存策略,保留训练进度

调整ModelCheckpoint配置,保存每个epoch的权重,同时保留最优权重的备份:

checkpoint_dir = '/content/drive/MyDrive/Node21/YOLOv8/tentativa_1/checkpoints/'
os.makedirs(checkpoint_dir, exist_ok=True)

# 保存每个epoch的检查点(带epoch编号)
epoch_checkpoint = ModelCheckpoint(
    filepath=os.path.join(checkpoint_dir, 'checkpoint_epoch_{epoch:03d}.ckpt'),
    save_weights_only=True,
    save_freq='epoch',  # 每个epoch都保存
    verbose=1
)

# 同时保存最优权重(保留原有需求)
best_checkpoint = ModelCheckpoint(
    filepath=os.path.join(checkpoint_dir, 'best_checkpoint.ckpt'),
    save_weights_only=True,
    save_best_only=True,
    monitor='val_loss',
    verbose=1
)

2. 续训时加载最后一轮权重并指定起始epoch

在续训代码中,自动识别最新的检查点,解析出已完成的epoch数,通过initial_epoch参数指定续训起始点:

import os
from tensorflow.keras.callbacks import ModelCheckpoint, EarlyStopping

checkpoint_dir = '/content/drive/MyDrive/Node21/YOLOv8/tentativa_1/checkpoints/'

# 加载模型架构并编译(与原训练参数一致)
model = unet()
optimizer = SGD(learning_rate=0.001)
model.compile(optimizer=optimizer, loss='binary_crossentropy', metrics=['accuracy'])

# 查找并加载最新的epoch检查点
initial_epoch = 0
checkpoint_files = [f for f in os.listdir(checkpoint_dir) if f.startswith('checkpoint_epoch_')]
if checkpoint_files:
    # 按epoch编号排序,取最后一个
    checkpoint_files.sort()
    latest_checkpoint = checkpoint_files[-1]
    # 解析已完成的epoch数
    initial_epoch = int(latest_checkpoint.split('_')[-1].split('.')[0])
    # 加载对应权重
    model.load_weights(os.path.join(checkpoint_dir, latest_checkpoint))

# 定义回调(包含双检查点策略)
early_stopping = EarlyStopping(
    patience=10,
    restore_best_weights=True,
    verbose=1
)

# 启动续训:指定initial_epoch参数
history = model.fit(
    np.array(train_images),
    np.array(train_masks),
    epochs=500,
    batch_size=6,
    initial_epoch=initial_epoch,  # 关键:从已完成的下一轮开始
    validation_data=(np.array(test_images), np.array(test_masks)),
    callbacks=[epoch_checkpoint, best_checkpoint, early_stopping]
)

3. 额外注意事项

  • 数据一致性:续训时必须保证训练集/验证集的划分、图像加载方式与原训练完全一致,避免因数据变化导致训练逻辑混乱
  • 完整状态恢复:若需要完全还原训练状态(如学习率衰减、优化器动量),建议使用tf.train.Checkpoint保存模型+优化器的完整状态,而非仅保存权重

内容的提问来源于stack exchange,提问作者Mauricésar Barbosa

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 03:52:02