如何在Keras中加载DCGAN训练检查点并生成对应像素数组
Keras DCGAN加载指定训练检查点生成数组实现方案
前置要求
- 先在当前代码环境中定义和训练时结构、参数完全一致的生成器、判别器、两个优化器对象,不需要手动初始化权重,检查点加载时会自动覆盖
- 确认要加载的轮次对应检查点文件存在,比如第N轮训练对应文件为
./training_checkpoints/ckpt-N.index和./training_checkpoints/ckpt-N.data-00000-of-00001
加载指定轮次检查点
import tensorflow as tf import os # 这里替换为你训练时用的生成器、判别器、优化器定义代码,结构参数必须和训练时完全匹配 generator = 【你的生成器定义代码】 discriminator = 【你的判别器定义代码】 generator_optimizer = tf.keras.optimizers.【你使用的优化器,参数和训练时一致】 discriminator_optimizer = tf.keras.optimizers.【你使用的优化器,参数和训练时一致】 # Checkpoint定义和你保存时的代码完全一致 checkpoint_dir = "./training_checkpoints" checkpoint = tf.train.Checkpoint( generator_optimizer=generator_optimizer, discriminator_optimizer=discriminator_optimizer, generator=generator, discriminator=discriminator ) # 加载指定轮次检查点,ckpt后的数字对应训练epoch轮次,比如加载第20轮就填ckpt-20 # 加expect_partial()避免未用到的权重(判别器、优化器)加载警告,不影响生成结果 checkpoint.restore(os.path.join(checkpoint_dir, "ckpt-20")).expect_partial()
生成目标数组
# 噪声维度必须和你训练时输入生成器的噪声维度完全一致,原版MNIST DCGAN默认用100维 noise_dim = 100 # 自定义需要生成的样本数量 num_generate_samples = 100 # 生成符合维度要求的随机噪声 test_noise = tf.random.normal([num_generate_samples, noise_dim]) # 推理模式生成数组,training设为False关闭Dropout、BatchNormalization等训练专属逻辑 generated_arrays = generator(test_noise, training=False) # 后处理:和训练时的预处理逻辑对齐,比如训练时你把MNIST像素从[0,1]缩到了[-1,1],这里反向转换 generated_arrays = (generated_arrays + 1) / 2.0 # 转为numpy数组格式方便后续计算观测值 generated_arrays_np = generated_arrays.numpy()
注意:如果生成数组不符合预期,优先排查三个点:生成器结构是否和训练时一致、噪声维度是否匹配、后处理归一化逻辑是否和训练时的预处理对应。
内容的提问来源于stack exchange,提问作者Barry
相关产品推荐
相关产品推荐

