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
相关产品推荐
相关产品推荐

