如何让Deepfake训练使用的自编码器从中断的epoch处恢复训练?
Deepfake自编码器断点续训实现方法
完全可以实现,你现有代码已经具备每100轮自动存权重的基础,只需要补充3部分逻辑即可:
调整方案
1. 优化权重存储逻辑
修改现有的save_model_weights()函数,存储权重的同时,同步记录当前训练到的epoch数到本地文件:
def save_model_weights(current_epoch): # 原有保存权重的逻辑保留,示例如下,和你原有逻辑对齐即可 aeA.save_weights(f"./aeA_weights_{current_epoch}.h5") aeB.save_weights(f"./aeB_weights_{current_epoch}.h5") # 新增记录当前epoch的逻辑 with open("./last_epoch.txt", "w") as f: f.write(str(current_epoch))
调用这个函数的地方,把当前epoch传入即可,比如原来的save_model_weights()改成save_model_weights(epoch)。
2. 新增启动时断点加载逻辑
在训练循环开始前,先检查是否有已保存的训练记录,如果有就加载对应权重和起始epoch:
import os # 原有数据集加载逻辑保留 train_setA = video.loading_images(setA_path)/255.0 train_setB = video.loading_images(setB_path)/255.0 train_setA += train_setB.mean( axis=(0,1,2) ) - train_setA.mean( axis=(0,1,2) ) batch_size = int(len(os.listdir(setA_path))/20) # 新增断点加载逻辑 start_epoch = 0 last_epoch_path = "./last_epoch.txt" if os.path.exists(last_epoch_path): with open(last_epoch_path, "r") as f: start_epoch = int(f.read().strip()) # 加载对应轮数的权重,如果你之前是存固定文件名,直接填固定路径即可 aeA.load_weights(f"./aeA_weights_{start_epoch}.h5") aeB.load_weights(f"./aeB_weights_{start_epoch}.h5") print(f"加载断点成功,从第{start_epoch+1}轮开始训练") # 初始化测试样本,避免断点启动时变量未定义 test_A = None test_B = None
3. 修改训练循环的起始位置
把原来的for epoch in range(1000000):改成从start_epoch开始循环:
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 % 100 == 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()
注意事项
- 如果你原来的
save_model_weights()没有给权重文件加epoch后缀,直接存的固定文件名,那加载的时候直接填固定路径即可,不需要拼接epoch参数 - 权重文件和last_epoch.txt要放在同一目录下,避免读取失败
- 如果训练到10000轮中断,启动后会自动加载10000轮的权重,从10001轮开始继续训练,不会覆盖之前的训练结果
内容的提问来源于stack exchange,提问作者pasho_6798
相关产品推荐
相关产品推荐

