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

如何让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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 17:18:03