发布时间:2026/9/1 15:07:36
PyTorch计算机视觉——cWGAN-GP实现可控图像生成 PyTorch计算机视觉——cWGAN-GP实现可控图像生成0. 前言1. cWGAN-GP 的核心思想2. 数据集2.1 数据集介绍2.2 数据集加载与预处理3. 模型架构3.1 条件生成器3.2 条件判别器3.3 梯度惩罚函数4. 模型训练小结相关链接0. 前言我们已经学习了 WGAN-GP (Wasserstein GAN with Gradient Penalty) 在CelebA和动漫面孔数据集上的应用展示了其生成高质量随机图像的能力。然而在实际应用中我们往往需要更有针对性的图像生成——例如我们希望生成特定类别的图像而不是完全随机的样本。这就引出了条件生成对抗网络 (Conditional GAN, cGAN) 的概念。本节将详细介绍条件WGAN-GP(cWGAN-GP) 的实现与应用以石头剪刀布彩色图像数据集为例展示如何通过引入类别标签信息实现对生成图像类别的精确控制。1. cWGAN-GP 的核心思想传统的GAN只能从随机噪声中生成图像无法控制生成图像的具体类别。cWGAN-GP通过在生成器和判别器中同时输入类别标签信息实现了条件生成的能力。cWGAN-GP的核心创新点包括条件嵌入将类别标签转换为向量表示与噪声或图像特征进行拼接可控生成通过指定不同的标签生成对应类别的图像保持优势继承了WGAN-GP的训练稳定性和高质量生成能力2. 数据集2.1 数据集介绍本节选择使用石头剪刀布数据集主要考虑以下因素小型数据集仅2520张图像适合快速实验和演示三类明确石头 (rock)、剪刀 (scissors)、布 (paper)类别清晰彩色图像128x128分辨率包含丰富的视觉特征2.2 数据集加载与预处理将图像调整为128x128像素并进行归一化确保输入数据在[-1,1]范围内这有利于模型的训练收敛importtorch;importtorch.nnasnnfromtorch.utils.dataimportDataLoaderimporttorchvision.transformsasTfromtorchvision.utilsimportmake_gridfromtorchvision.datasetsimportImageFolderfromtqdmimporttrangeimportnumpyasnpimportpandasaspdimportmatplotlib.pyplotasplt n_epochs100batch_size32z_dim100lr4e-4n_critic1lamda_gp10img_size128img_channels3n_class3fixed_latenttorch.randn(48,z_dim).cuda()# noises.shape 48 x z_dimfixed_labelstorch.LongTensor([iforiinrange(3)forjinrange(16)]).cuda()train_datasetImageFolder(./data/RockPaperScissors/train,transformT.Compose([T.Resize(img_size),T.ToTensor(),T.Normalize([0.5,0.5,0.5],[0.5,0.5,0.5])]))n_sampleslen(train_dataset)train_dataloaderDataLoader(train_dataset,batch_sizebatch_size,shuffleTrue,num_workers4,pin_memoryTrue)n_batchlen(train_dataloader)#n_batch79forimgs,labelsintrain_dataloader:print(imgs.shape,imgs.shape)print(labels,\n,labels.view(-1,16))breakdefdenorm(img_tensors):# Shift image pixel values to [0,1]returnimg_tensors*0.50.5defshow_imgs(images):fig,axplt.subplots(figsize(16,10))inputsmake_grid(denorm(images),nrow16)#inputs make_grid(images, nrow16)ax.imshow(inputs.permute(1,2,0))ax.set(xticks[],yticks[])plt.show()show_imgs(imgs)3. 模型架构3.1 条件生成器定义函数weights_init()用于模型参数初始化defweights_init(m):if(type(m)nn.ConvTranspose2dortype(m)nn.Conv2d):nn.init.normal_(m.weight.data,0.0,0.02)elif(type(m)nn.BatchNorm2d):nn.init.normal_(m.weight.data,0.0,0.02)nn.init.constant_(m.bias.data,0)生成器需要接收两个输入随机噪声向量z和类别标签labels# Generator Classdefbasic_G(in_channels):returnnn.Sequential(nn.ConvTranspose2d(in_channels,int(in_channels/2),4,2,1,biasFalse),nn.BatchNorm2d(int(in_channels/2)),nn.ReLU(True))classGenerator(nn.Module):def__init__(self):super().__init__()self.netnn.Sequential(basic_G(512),basic_G(256),basic_G(128),basic_G(64),nn.ConvTranspose2d(32,3,kernel_size4,stride2,padding1),nn.Tanh())self.label_embnn.Embedding(n_class,4*4)self.latentnn.Linear(z_dim,511*4*4)defforward(self,z,labels):yself.latent(z)cself.label_emb(labels)y_ctorch.cat([y,c],dim1)inputy_c.view(-1,512,4,4)outputself.net(input)returnoutput GGenerator().cuda()G.apply(weights_init)在以上代码中使用nn.Embedding将类别标签(0,1,2)映射为16维向量 (4x4)通过nn.Linear将100维噪声扩展为8176维 (511x4x4)将标签嵌入与噪声特征拼接形成512x4x4的初始特征图通过5层转置卷积逐步上采样至128x128x3。3.2 条件判别器判别器同样接收两个输入图像和对应的标签# Discriminator classdefbasic_D(in_channels):returnnn.Sequential(nn.Conv2d(in_channels,2*in_channels,4,2,1,biasFalse),nn.InstanceNorm2d(2*in_channels),nn.LeakyReLU(0.2,inplaceTrue))classDiscriminator(nn.Module):def__init__(self):super().__init__()self.netnn.Sequential(nn.Conv2d(img_channels1,64,4,2,1),nn.LeakyReLU(0.2,inplaceTrue),basic_D(64),basic_D(128),basic_D(256),basic_D(512),nn.Conv2d(1024,1,kernel_size4,stride1,padding0),nn.Flatten())self.label_codenn.Embedding(n_class,1*img_size*img_size)defforward(self,images,labels):ximages.view(-1,img_channels*img_size*img_size)cself.label_code(labels)x_ctorch.cat([x,c],dim1)inputx_c.view(-1,img_channels1,img_size,img_size)outself.net(input)returnout DDiscriminator().cuda()D.apply(weights_init)在判别器中使用InstanceNorm2d替代BatchNorm2d提高了训练稳定性通过nn.Embedding将标签编码为16384维 (128x128) 的向量将编码后的标签与展平的图像特征拼接形成4通道输入 (RGB 标签信息)。3.3 梯度惩罚函数在cWGAN-GP中梯度惩罚函数需要特别注意标签的使用# CGradient-Penalty functiondefgradient_penalty(D,real_data,fake_data,fake_labels):batch_sizereal_data.size(0)#real_data.shape batch_size x 3 x128^2# Sample Epsilon from uniform distributionepstorch.rand(batch_size,1,1,1).cuda()epseps.expand_as(real_data)#eps.shapebatch_size x 3 x 128^2# Interpolation between real data and fake data.interpolationeps*real_data(1-eps)*fake_data# get logits for interpolated imageslogitsD(interpolation,fake_labels)#shape batch_size x 1gradientstorch.autograd.grad(outputslogits,inputsinterpolation,grad_outputstorch.ones_like(logits),create_graphTrue,retain_graphTrue)[0]# Gradientsgradientsgradients.view(batch_size,-1)grad_normgradients.norm(2,1)gradient_penaltytorch.mean((grad_norm-1)**2)returngradient_penalty# Compute and return the gradient norm在计算梯度惩罚时必须使用伪造标签而非真实标签。这是因为梯度惩罚旨在约束判别器在真实数据分布和生成数据分布之间的行为使用伪造标签更符合实际生成场景实验表明使用真实标签会导致模型崩溃4. 模型训练定义模型优化器optimizer_Dtorch.optim.Adam(D.parameters(),lrlr,betas(0.0,0.9))optimizer_Gtorch.optim.Adam(G.parameters(),lrlr,betas(0.0,0.9))定义生成器和判别器训练函数deftrain_D(inputs,labels,optimizer_D):for_inrange(n_critic):real_predsD(inputs,labels)real_scoretorch.mean(real_preds)# create fake images and labels with random numberslatenttorch.randn(inputs.shape[0],z_dim).cuda()fake_labelstorch.LongTensor(torch.randint(0,n_class,(inputs.shape[0],))).cuda()fake_imagesG(latent,fake_labels)fake_predsD(fake_images.detach(),fake_labels.detach())fake_scoretorch.mean(fake_preds)# Train the optimizer_D with real_loss and fake_lossgpgradient_penalty(D,inputs,fake_images,fake_labels)lossfake_score-real_scorelamda_gp*gp optimizer_D.zero_grad()loss.backward()optimizer_D.step()returnloss.item(),real_score.item(),fake_score.item()deftrain_G(optimizer_G):# Create fake images and labelslatenttorch.randn(batch_size,z_dim).cuda()fake_labelstorch.LongTensor(torch.randint(0,n_class,(batch_size,))).cuda()fake_imagesG(latent,fake_labels)# Try to fool the discriminatorpredsD(fake_images,fake_labels)loss-torch.mean(preds)optimizer_G.zero_grad()loss.backward()optimizer_G.step()returnloss.item()训练过程的核心是交替更新判别器和生成器deffit(epochs):torch.cuda.empty_cache()# The DataFrame df is a recorder of the training historydfpd.DataFrame(np.empty([epochs,4]),indexnp.arange(epochs),columns[Loss_G,Loss_D,D(X),D(G(Z))])foriintrange(epochs):loss_G0.0;loss_D0.0;real_sc0.0;fake_sc0.0forreal_images,labelsintrain_dataloader:inputsreal_images.cuda()labelslabels.cuda()loss_d,real_score,fake_scoretrain_D(inputs,labels,optimizer_D)loss_Dloss_d;real_screal_score;fake_scfake_score loss_gtrain_G(optimizer_G)loss_Gloss_g# Record losses scoresdf.iloc[i,0]loss_G/n_batch df.iloc[i,1]loss_D/n_batch df.iloc[i,2]real_sc/n_batch df.iloc[i,3]fake_sc/n_batchifi0or(i1)%50:print(Epoch{:2}, Ls_G{:.2f}, Ls_D{:.2f}, D(X){:.2f}, D(G(Z)){:.2f}.format(i1,df.iloc[i,0],df.iloc[i,1],df.iloc[i,2],df.iloc[i,3]))fake_imagesG(fixed_latent,fixed_labels)show_imgs(fake_images.detach().cpu())returndf historyfit(n_epochs)查看模型训练过程性能变化曲线# Show the training historydfhistory fig,axplt.subplots(1,2,figsize(9,4),sharexTrue)df.plot(axax[0],y[0,1],style[r,b:])gpdf.iloc[:,1]-df.iloc[:,3]df.iloc[:,2]ax[0].plot(gp,labelGradient Penalty,colork,linestyle-)ax[0].set(ylabelloss)ax[0].legend()df.plot(axax[1],y[2,3],style[r-,b:])foriinrange(2):ax[i].grid(whichmajor,axisboth,colorg,linestyle:)ax[i].set(xlabelepoch)plt.show()通过指定不同的标签我们可以精确控制生成图像的类别defgenerate_image(G,digital):ztorch.randn(1,100).cuda()Nlen(train_dataset.classes)-1ifdigitalN:labeltorch.LongTensor([digital]).cuda()imgG(z,label).data.cpu()show_imgs(img)else:print(Your label is bigger than ,N)generate_image(G,0)小结本节详细介绍了cWGAN-GP在石头剪刀布数据集上的实现与应用展示了条件GAN在可控图像生成方面的强大能力。通过引入类别标签信息我们能够精确控制生成图像的类别同时保持了WGAN-GP的训练稳定性和高质量生成能力。相关链接PyTorch计算机视觉1——计算机视觉的数学工具PyTorch计算机视觉2——神经网络模型训练与PyTorch基础PyTorch计算机视觉3——卷积神经网络CNN详解与实现PyTorch计算机视觉4——迁移学习Transfer Learning详解与实现PyTorch计算机视觉5——生成对抗网络Generative Adversarial NetworkGANPyTorch计算机视觉6——深度卷积对抗神经网络DCGANPyTorch计算机视觉7——条件生成对抗网络cGANPyTorch计算机视觉8——WGAN及其变体WGAN-GP

相关新闻

2026/9/1 15:07:36

基于STM32的GPS智能公交报站系统:从原理到实车调试

简介:本资源是一套基于STM32F103C8T6最小系统板实现的智能公交报站系统完整嵌入式源码,面向嵌入式初学者、STM32课程设计学生及物联网应用开发爱好者,解决公交场景下自动定位、语音播报与站名显示等核心功能的软硬件协同实现问题。压缩包共10…

2026/9/1 15:07:36

MKVToolNix:无损混流工具,高效处理视频封装与合并

你是不是也遇到过这样的问题:下载了一部电影,结果视频和字幕是分开的;或者从不同来源收集了多段视频素材,想要合并成一个文件;又或者想给视频换个音轨,却发现常规的视频编辑软件要么操作复杂,要…

2026/9/1 15:07:36

C语言结构体快速通关:从定义到内存对齐的实战指南

这次我们来看一个C语言结构体的快速通关教程。对于很多初学者来说,结构体是C语言从基础语法迈向复杂数据结构的关键一步,它直接关系到后续链表、文件操作乃至数据结构的学习。这个教程的核心目标不是讲复杂的理论,而是让你能快速上手、理解结…

2026/9/1 15:22:40

2024小红书iOS笔试深度复盘:底层原理与实战避坑全解析

2024年春招,小红书iOS开发岗第三批笔试我完整走了一轮。说实话,动笔之前我以为是常规的八股题堆砌,真坐到电脑前才发现,这份卷子几乎是把iOS开发者在日常工作中会踩到的坑铺开当考题用。整场笔试涉及OC和Swift的底层原理、内存管理…

2026/9/1 15:22:40

数字艺术资源技术解析:从获取到应用的全流程实践指南

这次我们来看一个名为“【furryDiives】《Binggan&Xingyun - Ox New Year 2021》”的项目。从标题看,这很可能是一个与Furry(兽人)艺术创作相关的数字作品,具体指向Diives创作的角色Binggan和Xingyun在2021牛年新年的主题内容…

2026/9/1 15:22:40

单词举一反三技巧的核心要点与实践经验

很多人背单词一直陷入“单词单点记忆”:背一个、会一个,不延伸、不拓展。看似每天打卡、词汇量稳步上涨,实则词汇碎片化、不成体系、转化率极低。 考场遇到同源变形、同义替换、衍生搭配立刻卡顿;写作永远只会用基础简单词&#x…

2026/9/1 15:17:38

Go结构体与面向对象编程实战:嵌入组合取代继承

Go结构体与面向对象编程实战:嵌入组合取代继承 文章导语 Go没有class,没有extends,没有构造函数——这在OOP语言开发者看来简直是"反面向对象"。但Go用结构体嵌入(Embedding)和接口组合实现了更优雅的代码复…

2026/8/31 1:05:20

vSound小提琴数字处理器实操指南:从接线到演出的完整配置

电小提琴或者原声小提琴插电演出,第一个绕不开的坎就是声音难听。原声琴的共鸣和空气感一旦进了拾音器,出来的往往是一坨干瘪、发尖、带着奇怪塑料味的信号。我当初第一次把琴接上乐队调音台,直接被主唱吐槽"你这声音像在锯钢丝"。…

2026/9/1 8:27:47

传感器接口IC如何攻克生物化学传感的微弱信号难题?

1. 从电极到比特流:为什么生物化学传感必须依赖专用接口IC 做生物化学传感的人都有过类似的经历:明明传感器本身性能很好,信号输出却一塌糊涂——噪声大、漂移明显、重复性差,怎么调都达不到预期。很多时候问题并不在传感器&#…

2026/9/1 7:04:43

STM32F411CEU6多通道ADC采集:扫描模式+DMA实现详解

1. 多通道 ADC 的用武之地把“Multichannel ADC”和“STM32F411CEU6”这两个关键字放在一起,其实就是嵌入式开发里最常遇到的一类需求:用一块不算贵的 MCU,同时采集多路模拟信号。STM32F411CEU6 是 48 引脚的 Cortex-M4F 主控,主频…

2026/9/1 0:00:42

USB Type-C PCB布局分区设计:电源、高速信号与PD协议全攻略

做硬件这行,Type-C接口算是典型的“看着简单,做起来全坑”的东西。光引脚就24个,高低速信号、电源、控制线全部塞在一个小小的连接器里,如果PCB布局不做规划,打样回来基本就是“插上没反应”、“高速掉线”、“静电一打…

2026/9/1 0:00:42

系统编程学习原型如何补齐稳定性边界

系统编程学习原型如何补齐稳定性边界预算有限时&#xff0c;我先优化明显多余的复制&#xff0c;而不是猜测性地换容器。用借用传递只读数据通常就能减少分配&#xff1a; fn parse(line: &str) -> Result<Item, Error> { /* ... */ }用基准确认热点确实在分配&am…

2026/9/1 0:00:42

雨花区哪家财务公司代理记账比较好?

在雨花区&#xff0c;企业处理财税事务常常面临诸多挑战&#xff0c;选择一家靠谱的财务公司至关重要。湖南巨勤财务管理咨询有限公司就是本地正规实体财税服务机构&#xff0c;深耕本地工商财税行业多年&#xff0c;熟悉当地工商局、税务局最新政策与申报流程。主营公司注册、…

2026/9/1 0:00:42

USB Type-C PCB布局分区设计:电源、高速信号与PD协议全攻略

做硬件这行&#xff0c;Type-C接口算是典型的“看着简单&#xff0c;做起来全坑”的东西。光引脚就24个&#xff0c;高低速信号、电源、控制线全部塞在一个小小的连接器里&#xff0c;如果PCB布局不做规划&#xff0c;打样回来基本就是“插上没反应”、“高速掉线”、“静电一打…

2026/9/1 0:00:42

系统编程学习原型如何补齐稳定性边界

系统编程学习原型如何补齐稳定性边界预算有限时&#xff0c;我先优化明显多余的复制&#xff0c;而不是猜测性地换容器。用借用传递只读数据通常就能减少分配&#xff1a; fn parse(line: &str) -> Result<Item, Error> { /* ... */ }用基准确认热点确实在分配&am…

2026/9/1 0:00:42

雨花区哪家财务公司代理记账比较好?

在雨花区&#xff0c;企业处理财税事务常常面临诸多挑战&#xff0c;选择一家靠谱的财务公司至关重要。湖南巨勤财务管理咨询有限公司就是本地正规实体财税服务机构&#xff0c;深耕本地工商财税行业多年&#xff0c;熟悉当地工商局、税务局最新政策与申报流程。主营公司注册、…