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
第四步:修改训练循环与生成样本可视化
- 将训练循环的遍历语句
for n, (real_samples, mnist_labels) in enumerate(train_loader):修改为for n, (real_samples, _) in enumerate(train_loader):,丢弃无用的哑标签,其余训练逻辑完全保持不变。 - 将原代码末尾生成样本可视化的整段替换为以下代码,适配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
相关产品推荐
相关产品推荐

