如何在PyTorch中按类别单独保存生成图像(非网格格式)
按类别单独保存ACGAN生成图像的修改方案
嘿,这个需求很常见,要把ACGAN生成的图像按类别单独保存而不是网格图,只需要调整一下样本生成保存的逻辑就行。下面是两种实用的修改方案,你可以根据自己的需求选择:
方案1:单目录下按「类别_序号」命名保存
这个方案会把所有类别生成的图像放在同一个目录里,文件名包含类别ID和样本序号,比如class_0_0.png(第0类的第0张)。
修改后的样本生成函数如下:
import os from torchvision.utils import save_image import torch from torch.autograd import Variable import numpy as np # 替换原来的sample_image函数,或者新增这个函数 def sample_images_by_class(num_samples_per_class, batches_done): """按类别单独保存生成图像,每个类别生成num_samples_per_class张""" # 创建当前批次的保存目录,避免文件覆盖 save_dir = f"images/class_samples_{batches_done}" os.makedirs(save_dir, exist_ok=True) # 遍历每个类别 for class_idx in range(opt.n_classes): # 为当前类别生成指定数量的噪声向量 z = Variable(torch.FloatTensor(np.random.normal(0, 1, (num_samples_per_class, opt.latent_dim)))) # 生成对应类别的标签(全部为当前class_idx) labels = Variable(torch.LongTensor([class_idx] * num_samples_per_class)) # 生成图像 gen_imgs = generator(z, labels) # 逐个保存每张图像 for img_idx in range(num_samples_per_class): img_path = f"{save_dir}/class_{class_idx}_{img_idx}.png" # 取出单张图像并保存,normalize=True确保图像像素值在[0,1]之间 save_image(gen_imgs[img_idx].data, img_path, normalize=True)
方案2:按类别分文件夹保存
如果希望每个类别的图像单独放在一个文件夹里(比如images/class_samples_1000/class_0/0.png),可以用这个版本:
def sample_images_by_class(num_samples_per_class, batches_done): base_save_dir = f"images/class_samples_{batches_done}" os.makedirs(base_save_dir, exist_ok=True) for class_idx in range(opt.n_classes): # 为每个类别创建单独的文件夹 class_dir = f"{base_save_dir}/class_{class_idx}" os.makedirs(class_dir, exist_ok=True) z = Variable(torch.FloatTensor(np.random.normal(0, 1, (num_samples_per_class, opt.latent_dim)))) labels = Variable(torch.LongTensor([class_idx] * num_samples_per_class)) gen_imgs = generator(z, labels) for img_idx in range(num_samples_per_class): img_path = f"{class_dir}/{img_idx}.png" save_image(gen_imgs[img_idx].data, img_path, normalize=True)
关键细节说明
- Generator的Forward方法:确保你的Generator类已经实现了接受噪声
z和标签labels的forward方法,比如:
class Generator(nn.Module): def __init__(self): # 你的初始化代码...(和原来一致) def forward(self, z, labels): # 将标签嵌入向量与噪声向量拼接 label_embedding = self.label_emb(labels) gen_input = torch.cat((label_embedding, z), dim=-1) # 后续的线性层和卷积层处理 out = self.l1(gen_input) out = out.view(out.shape[0], 128, self.init_size, self.init_size) img = self.conv_blocks(out) return img
(你的代码里已经调用了generator(z, labels),所以这个forward方法应该已经存在,只是确认一下逻辑没问题)
参数调整:你可以根据需求修改
num_samples_per_class的值,比如每个类别生成10张;文件名格式也可以自定义,比如直接用{class_idx}.png(如果每个类别只保存1张)。依赖导入:确保已经导入
os模块,因为用到了创建目录的函数os.makedirs。
最后,训练过程中你只需要调用这个新的sample_images_by_class函数,替代原来的sample_image就可以了。
内容的提问来源于stack exchange,提问作者dayday
相关产品推荐
相关产品推荐

