GAN模型仅保存首个样本生成图,求实现全样本Batch/Epoch生成图保存
GAN生成图像仅保存第一张问题排查与解决
问题核心分析
你遇到的问题分两种情况,对应不同的解决方向:
- 生成逻辑误解:当前代码是无条件GAN,生成图像完全由随机噪声生成,和每个batch的真实图像没有对应关系——并非“针对每张真实图像生成匹配结果”,而是生成符合数据集分布的随机图像。
- 数据集加载异常:如果确实只生成了第一张的结果,优先排查数据集是否正确加载:
- 确认图片后缀是小写
.jpg(代码中匹配规则为*.jpg,大写.JPG会被忽略) - 打印数据集长度验证加载数量:
print(f"实际加载图片数: {len(dataset)}") print(f"图片路径: {dataset.img_paths}")
- 确认图片后缀是小写
代码修正方案
1. 确保数据集全量加载
如果图片后缀有大写,修改MyDataset中的路径匹配规则:
self.img_paths = sorted(glob.glob(os.path.join(path, '*.[jJ][pP][gG]')))
2. 若需生成与真实图像对应的结果(条件GAN改造)
如果你确实想针对每张真实图像生成相似的生成图,需要将无条件GAN改为条件GAN,给模型传入真实图像特征作为生成条件:
修改Generator类:
class Generator(nn.Module): def __init__(self): super(Generator, self).__init__() # 新增真实图像特征编码器 self.img_encoder = nn.Sequential( nn.Conv2d(1, 64, 4, 2, 1), nn.LeakyReLU(0.2), nn.Conv2d(64, 128, 4, 2, 1), nn.LeakyReLU(0.2), nn.Conv2d(128, 256, 4, 2, 1), nn.LeakyReLU(0.2), nn.Conv2d(256, 512, 4, 2, 1), nn.LeakyReLU(0.2), ) self.feature_dim = 512 * (image_size//16) * (image_size//16) # 合并噪声与图像特征的全连接层 self.fc1 = nn.Linear(noise_dim + self.feature_dim, 128) self.fc2 = nn.Linear(128, 256) self.fc3 = nn.Linear(256, 512) self.fc4 = nn.Linear(512, image_size*image_size) def forward(self, noise, real_img): # 提取真实图像特征 img_feat = self.img_encoder(real_img) img_feat = img_feat.view(img_feat.size(0), -1) # 合并噪声与图像特征 x = torch.cat([noise.view(noise.size(0), -1), img_feat], dim=1) x = nn.functional.leaky_relu(self.fc1(x), 0.2) x = nn.functional.leaky_relu(self.fc2(x), 0.2) x = nn.functional.leaky_relu(self.fc3(x), 0.2) x = nn.functional.tanh(self.fc4(x)) x = x.view(x.size(0), 1, image_size, image_size) return x
修改训练循环中的生成与保存逻辑:
# 训练生成器时传入真实图像作为条件 G.zero_grad() noise = torch.randn(images.size(0), noise_dim, 1, 1).cuda() fake_images = G(noise, real_images) fake_labels = torch.ones((images.size(0), 1)).cuda() fake_outputs = D(fake_images) g_loss = criterion(fake_outputs, fake_labels.view(-1, 1)) g_loss.backward() optimizer_G.step() # 保存生成图像时,传入当前batch的真实图像 generated_images = G(fixed_noise, real_images) for j, img in enumerate(generated_images): save_image(img, f'generated_images/image_epoch{epoch}_batch{i}_idx{j}.png')
3. 冗余代码清理
删除重复的import torch语句,整理代码结构。
内容的提问来源于stack exchange,提问作者Candy
相关产品推荐
相关产品推荐

