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

Tensorflow2/Keras下含多模块的GAN如何保存训练状态实现断点续训

ACGAN断点续训完整解决方案

你担心单独保存三个子模型会失去关联是误解:GAN的三个模型(生成器、判别器、复合GAN)的依赖关系是训练逻辑层面的调用规则,并非模型结构内部的绑定,只要三个模型各自的权重、优化器状态与中断时完全一致,按原有训练逻辑交替调用,恢复后的训练效果就能和未中断的连续训练完全一致。

具体实现步骤

1. 修改断点保存逻辑

原有仅保存复合GAN权重的逻辑需要替换为保存三个模型的完整对象(包含结构、权重、优化器状态),同时记录当前训练步数:

import pickle
import random
import numpy as np

# 每10个epoch触发的保存逻辑
def summarize_performance(step, g_model, d_model, gan_model):
    # 原有生成样本、评估的逻辑保持不变
    # ...

    # 保存三个完整模型(推荐用Keras原生.keras格式,无需额外依赖)
    step_str = f"{step+1:08d}"
    g_model.save(f'g_model_{step_str}.keras')
    d_model.save(f'd_model_{step_str}.keras')
    gan_model.save(f'gan_model_{step_str}.keras')
    
    # 记录当前步数,避免后续解析文件名出错
    with open('latest_step.txt', 'w') as f:
        f.write(str(step+1))
    
    # 可选:保存随机数状态,追求完全一致的训练轨迹时需要
    np.save(f'np_random_state_{step_str}.npy', np.random.get_state())
    with open(f'py_random_state_{step_str}.pkl', 'wb') as f:
        pickle.dump(random.getstate(), f)
    
    print(f'>已保存第{step+1}步的所有训练状态')

2. 修改断点加载逻辑

启动训练时先判断是否存在断点,存在则直接加载三个完整模型、恢复训练步数和随机状态,无需重新手动构建模型架构:

import os
import pickle
import random
import numpy as np
from keras.models import load_model

def load_training_state():
    # 无断点时返回新构建的模型和起始步数0,build_*为你原有的模型构建函数
    if not os.path.exists('latest_step.txt'):
        return build_generator(), build_discriminator(), build_gan(), 0
    
    with open('latest_step.txt', 'r') as f:
        latest_step = int(f.read().strip())
    step_str = f"{latest_step:08d}"
    
    # 加载三个完整模型,优化器状态会自动恢复
    g_model = load_model(f'g_model_{step_str}.keras')
    d_model = load_model(f'd_model_{step_str}.keras')
    gan_model = load_model(f'gan_model_{step_str}.keras')
    
    # 可选:恢复随机数状态
    np.random.set_state(np.load(f'np_random_state_{step_str}.npy', allow_pickle=True))
    with open(f'py_random_state_{step_str}.pkl', 'rb') as f:
        random.setstate(pickle.load(f))
    
    print(f'>断点加载完成,从第{latest_step}步开始训练')
    return g_model, d_model, gan_model, latest_step

3. 后续训练逻辑保持不变

拿到加载好的三个模型和起始步数后,按你原有的训练逻辑交替训练判别器和生成器即可,无需额外调整。


内容的提问来源于stack exchange,提问作者Moritz Grünbauer

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 14:24:03