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

如何在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)

关键细节说明

  1. 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方法应该已经存在,只是确认一下逻辑没问题)

  1. 参数调整:你可以根据需求修改num_samples_per_class的值,比如每个类别生成10张;文件名格式也可以自定义,比如直接用{class_idx}.png(如果每个类别只保存1张)。

  2. 依赖导入:确保已经导入os模块,因为用到了创建目录的函数os.makedirs。

最后,训练过程中你只需要调用这个新的sample_images_by_class函数,替代原来的sample_image就可以了。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 09:16:46