【问题标题】:How to load GAN checkpoint properly in PyTorch?如何在 PyTorch 中正确加载 GAN 检查点?
【发布时间】:2022-11-07 03:03:25
【问题描述】:

我在 256x256 图像上训练了 GAN,基本上扩展了 PyTorch 自己的 DCGAN tutorial 中的代码以适应更大分辨率的图像。模型和优化器初始化如下所示:

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

gen = Generator(...).to(device)
disc = Discriminator(...).to(device)

opt_gen = optim.Adam(gen.parameters(), ...)
opt_disc = optim.Adam(disc.parameters(), ...)

gen.train()
disc.train()

GAN 产生了高质量的样本。在每个 epoch 中,我使用与生成器相同的输入向量 fixed_noise 生成了一些图像(并使用 SummaryWriter 在 Tensorboard 上查看它们):

with torch.no_grad():
    fake = gen(fixed_noise)

    img_grid_real = torchvision.utils.make_grid(
        real[:NUM_VISUALIZATION_SAMPLES], normalize=True
    )
    img_grid_fake = torchvision.utils.make_grid(
        fake[:NUM_VISUALIZATION_SAMPLES], normalize=True
    )

    writer_real.add_image("Real", img_grid_real, global_step=step)
    writer_fake.add_image("Fake", img_grid_fake, global_step=step)

我在每个训练周期后保存了 GAN,如下所示:

checkpoint = {
    "gen_state": gen.state_dict(),
    "gen_optimizer": opt_gen.state_dict(),
    "disc_state": disc.state_dict(),
    "disc_optimizer": opt_disc.state_dict()
}
torch.save(checkpoint, f"checkpoints/checkpoint_{epoch_number}.pth.tar")

到目前为止,我已经在 CentOS7.9 机器上用 NVIDIA T4 GPU 和 PyTorch 1.11.0 训练了 GAN。然后我在我的个人计算机(Windows 10、NVIDIA GTX1050Ti、PyTorch 1.10.1)上rsync'd 了几个检查点(已按上述方式保存)。对 GAN 使用完全相同的类定义,并以相同的方式对其进行初始化(参见第一个代码 sn-p,除了将它们设置为训练模式),我加载了一个检查点,如下所示:

checkpoint = torch.load(f"checkpoints/checkpoint_10.pth.tar")
gen.load_state_dict(checkpoint["gen_state"])
opt_gen.load_state_dict(checkpoint["gen_optimizer"])
disc.load_state_dict(checkpoint["disc_state"])
opt_disc.load_state_dict(checkpoint["disc_optimizer"])

然后,我使用与第二个代码 sn-p 相同的代码,使用经过训练的 GAN 生成一些图像,现在在我的机器中加载了检查点。这产生了垃圾输出:

我尝试使用我拥有的所有检查点,以及所有输出的废话。我在 PyTorch 论坛中查找问题(123),但似乎没有任何帮助。

我保存/加载模型是否错误?

【问题讨论】:

    标签: python-3.x image pytorch generative-adversarial-network dcgan


    【解决方案1】:

    您找到任何解决方案了吗?

    【讨论】:

    猜你喜欢
    • 2019-07-07
    • 2021-11-29
    • 1970-01-01
    • 2021-10-05
    • 2023-01-26
    • 1970-01-01
    • 1970-01-01
    • 2021-01-15
    • 2013-11-13
    相关资源
    最近更新 更多