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

GAN模型仅保存首个样本生成图,求实现全样本Batch/Epoch生成图保存

GAN生成图像仅保存第一张问题排查与解决

问题核心分析

你遇到的问题分两种情况,对应不同的解决方向:

  1. 生成逻辑误解:当前代码是无条件GAN,生成图像完全由随机噪声生成,和每个batch的真实图像没有对应关系——并非“针对每张真实图像生成匹配结果”,而是生成符合数据集分布的随机图像。
  2. 数据集加载异常:如果确实只生成了第一张的结果,优先排查数据集是否正确加载:
    • 确认图片后缀是小写.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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 00:02:04