拓冰建站拓冰建站
首页 / 资讯中心 / 正文

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
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门