加载TensorFlow模型.h5权重后能否断点续训?现有训练代码需修改吗?
问题解答
断点续训的核心逻辑
model.load_weights 加载.h5权重文件时,仅会恢复模型的可训练参数(权重、偏置),不会自动恢复训练进度相关的状态:
- 你举例的场景里,重启后模型权重确实是第10000个epoch保存的结果,参数更新会基于这个权重继续进行
- 但你的现有代码会从epoch=0重新开始计数,且优化器状态、学习率进度等都会重置,不属于完整的断点续训
现有代码需要调整的点
- 新增训练状态的持久化存储
每次保存模型权重时,同步保存当前的epoch数、优化器状态(如果使用Adam、RMSprop等带历史状态的优化器,重置状态会导致训练出现异常波动)、学习率等训练参数。可以简单用json文件存储数值类状态,优化器状态可以和模型权重一起存储。
- 新增训练状态的持久化存储
- 调整epoch循环的起始值
训练启动时先读取本地存储的上次训练epoch数,循环从该值+1开始,而不是固定从0开始。
- 调整epoch循环的起始值
- 修复变量未定义的bug
现有代码中test_A、test_B仅在epoch % 50 == 0时赋值,重启后第一个epoch如果不是50的倍数,会直接报错变量未定义,需要在循环外提前初始化这两个变量,或者首次循环时强制采样一次测试数据。
- 修复变量未定义的bug
- 按需存储优化器状态
如果使用自适应优化器,建议改为用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
相关产品推荐
相关产品推荐

