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

PyTorch简易GAN如何用无标签自定义PNG图替换MNIST/CIFAR数据集

PyTorch GAN 自定义PNG数据集改造方案

所有改造完全保留原有代码逻辑、API与参数设计,仅针对336*336 RGB格式无标签PNG图做适配修改。


第一步:替换数据加载段代码

将原代码中从# PREPARING TRAINING DATA到# PLOTTING SAMPLES的整段内容替换为以下代码:

# PREPARING TRAINING DATA
# 适配RGB三通道336*336图像的预处理逻辑
transform = transforms.Compose(
    [
        transforms.Resize((336, 336)), # 强制统一所有图像尺寸,避免个别图尺寸异常报错
        transforms.ToTensor(),
        transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) # 三通道归一化到[-1,1]区间,匹配生成器Tanh输出范围
    ]
)

# LOADING DATA
train_set = torchvision.datasets.ImageFolder(
    root="./custom_dataset",
    transform=transform
)

# CREATE DATALOADER
batch_size = 32
train_loader = torch.utils.data.DataLoader(
    train_set, batch_size=batch_size, shuffle=True, drop_last=True # 丢弃最后一个不完整batch,避免尺寸不匹配报错
)

# PLOTTING SAMPLES
real_samples, _ = next(iter(train_loader)) # 无标签场景用匿名变量接收返回的哑标签
plt.figure(figsize=(8,8))
for i in range(16):
    ax = plt.subplot(4, 4, i + 1)
    # 通道顺序转换+反归一化,适配matplotlib RGB图像显示要求
    img = real_samples[i].permute(1,2,0) * 0.5 + 0.5
    plt.imshow(img)
    plt.xticks([])
    plt.yticks([])
plt.show()

第二步:修改判别器适配输入维度

原判别器针对2828单通道MNIST设计,需修改展平维度适配336336三通道输入,替换原Discriminator类为以下代码:

class Discriminator(nn.Module):
    def __init__(self):
        super().__init__()
        img_flatten_dim = 3 * 336 * 336 # 三通道336*336图像展平后的总维度
        self.model = nn.Sequential(
            nn.Linear(img_flatten_dim, 1024),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(1024, 512),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(512, 256),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(256, 1),
            nn.Sigmoid(),
        )

    def forward(self, x):
        img_flatten_dim = 3 * 336 * 336
        x = x.view(x.size(0), img_flatten_dim)
        output = self.model(x)
        return output

第三步:修改生成器适配输出维度

原生成器输出为2828单通道尺寸,需修改输出维度匹配336336三通道要求,替换原Generator类为以下代码:

class Generator(nn.Module):
    def __init__(self):
        super().__init__()
        img_flatten_dim = 3 * 336 * 336
        self.model = nn.Sequential(
            nn.Linear(100, 256),
            nn.ReLU(),
            nn.Linear(256, 512),
            nn.ReLU(),
            nn.Linear(512, 1024),
            nn.ReLU(),
            nn.Linear(1024, img_flatten_dim),
            nn.Tanh(),
        )

    def forward(self, x):
        output = self.model(x)
        output = output.view(x.size(0), 3, 336, 336) # 输出reshape为三通道336*336图像格式
        return output

第四步:修改训练循环与生成样本可视化

  1. 将训练循环的遍历语句for n, (real_samples, mnist_labels) in enumerate(train_loader):修改为for n, (real_samples, _) in enumerate(train_loader):,丢弃无用的哑标签,其余训练逻辑完全保持不变。
  2. 将原代码末尾生成样本可视化的整段替换为以下代码,适配RGB图像显示:
# SAMPLES
latent_space_samples = torch.randn(batch_size, 100).to(device=device)
generated_samples = generator(latent_space_samples)

generated_samples = generated_samples.cpu().detach()
plt.figure(figsize=(8,8))
for i in range(16):
    ax = plt.subplot(4, 4, i + 1)
    img = generated_samples[i].permute(1,2,0) * 0.5 + 0.5
    plt.imshow(img)
    plt.xticks([])
    plt.yticks([])
plt.show()

注意事项

  • 提前在VS Code工作目录下新建名为custom_dataset的文件夹,在该文件夹内新建一个任意名称的子文件夹(例如命名为imgs),将所有PNG训练图全部存入这个子文件夹即可,无需额外标注。ImageFolder会自动扫描目录下的图片,完全适配无标签训练场景。
  • 全连接网络处理336336分辨率图像时参数量远大于MNIST数据集,训练显存占用会明显升高。如果出现显存不足报错,可适当调小batch_size数值,或在预处理的Resize步骤中将图像统一缩放到更小尺寸(如6464)再训练。

内容的提问来源于stack exchange,提问作者Sterik

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 12:06:51