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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 07:54:06