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

加载TensorFlow模型.h5权重后能否断点续训?现有训练代码需修改吗?

问题解答

断点续训的核心逻辑

model.load_weights 加载.h5权重文件时,仅会恢复模型的可训练参数(权重、偏置),不会自动恢复训练进度相关的状态:

  • 你举例的场景里,重启后模型权重确实是第10000个epoch保存的结果,参数更新会基于这个权重继续进行
  • 但你的现有代码会从epoch=0重新开始计数,且优化器状态、学习率进度等都会重置,不属于完整的断点续训

现有代码需要调整的点

    1. 新增训练状态的持久化存储
      每次保存模型权重时,同步保存当前的epoch数、优化器状态(如果使用Adam、RMSprop等带历史状态的优化器,重置状态会导致训练出现异常波动)、学习率等训练参数。可以简单用json文件存储数值类状态,优化器状态可以和模型权重一起存储。
    1. 调整epoch循环的起始值
      训练启动时先读取本地存储的上次训练epoch数,循环从该值+1开始,而不是固定从0开始。
    1. 修复变量未定义的bug
      现有代码中test_A、test_B仅在epoch % 50 == 0时赋值,重启后第一个epoch如果不是50的倍数,会直接报错变量未定义,需要在循环外提前初始化这两个变量,或者首次循环时强制采样一次测试数据。
    1. 按需存储优化器状态
      如果使用自适应优化器,建议改为用model.save()存储完整模型(包含结构、权重、优化器状态),或者单独保存优化器的权重,加载时同步恢复。

调整后参考代码

import json
import os
import numpy as np
import cv2

# 初始化训练起始epoch
start_epoch = 0
state_path = 'models/train_state.json'

# 加载历史训练状态
if os.path.exists(state_path):
    with open(state_path, 'r') as f:
        train_state = json.load(f)
        start_epoch = train_state['current_epoch']

# 加载模型权重
enc.load_weights('models/encoder.h5')
decoderA.load_weights('models/decoder_A.h5')
decoderB.load_weights('models/decoder_B.h5')
# 如果用完整模型存储的话可以改为:
# aeA = tf.keras.models.load_model('models/aeA.h5')
# aeB = tf.keras.models.load_model('models/aeB.h5')

def save_model_weights(current_epoch):
    enc.save_weights('models/encoder.h5')
    decoderA.save_weights('models/decoder_A.h5')
    decoderB.save_weights('models/decoder_B.h5')
    # 保存训练状态
    with open(state_path, 'w') as f:
        json.dump({'current_epoch': current_epoch}, f)

# 提前初始化测试数据,避免变量未定义报错
test_A, test_B = None, None

for epoch in range(start_epoch, 1000000):
    batch_size = 64
    warped_A, target_A = train_util.training_data(train_setA, batch_size)
    warped_B, target_B = train_util.training_data(train_setB, batch_size)

    loss_A = aeA.train_on_batch(warped_A, target_A)
    loss_B = aeB.train_on_batch(warped_B, target_B)
    print(loss_A, loss_B)
    print('Current epoch no... ' + str(epoch))

    if epoch % 50 == 0:
        save_model_weights(epoch)
        print('Model weights saved')
        test_A = target_A[0:14]
        test_B = target_B[0:14]
    # 首次运行没拿到测试数据时临时采样
    if test_A is None or test_B is None:
        test_A = target_A[0:14]
        test_B = target_B[0:14]

    figure_A = np.stack([
        test_A,
        aeA.predict(test_A),
        aeB.predict(test_A),
        ], axis=1 )
    figure_B = np.stack([
        test_B,
        aeB.predict(test_B),
        aeA.predict(test_B),
        ], axis=1 )

    figure = np.concatenate( [ figure_A, figure_B ], axis=0 )
    figure = figure.reshape( (4,7) + figure.shape[1:] )
    figure = train_util.stack_images( figure )

    figure = np.clip( figure * 255, 0, 255 ).astype('uint8')

    cv2.imshow( "", figure )
    key = cv2.waitKey(1)
    if key == ord('q'):
        save_model_weights(epoch)
        exit()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 02:36:05