PyTorch计算机视觉——WGAN-GP在图像生成中的应用0. 前言1. WGAN-GP 技术原理简述2. 数据集分析2.1 数据集简介2.2. 数据加载与预处理3. 模型构建3.1 生成器3.2 判别器3.3 梯度惩罚实现4. 训练模型5. 实验结果与分析小结相关链接0. 前言生成对抗网络 (Generative Adversarial Network, GAN) 自2014年由Ian Goodfellow提出以来已成为深度学习领域最具创新性的技术之一。然而原始GAN面临着训练不稳定、模式坍塌等挑战这些问题限制了其在实际应用中的效果。Wasserstein GAN with Gradient Penalty(WGAN-GP) 作为一种改进方案通过引入Wasserstein距离和梯度惩罚项有效解决了这些问题。本节将使用WGAN-GP在CelebA人脸数据集和动漫面孔数据集上实现图像生成包括代码实现、训练过程分析以及结果评估。1. WGAN-GP 技术原理简述WGAN-GP 的核心创新在于Wasserstein距离替代传统GAN使用的JS散度提供更平滑的梯度使训练过程更加稳定梯度惩罚 (Gradient Penalty)强制判别器 (Critic) 的梯度范数接近1满足Lipschitz约束条件弃用批归一化在判别器中使用实例归一化 (Instance Normalization) 替代批归一化避免批次内样本间的相互影响这些改进使得WGAN-GP对超参数的选择不那么敏感减少了模式坍塌的风险。2. 数据集分析2.1 数据集简介CelebA数据集包含202599张名人面部图像广泛用于人脸识别和生成任务动漫面孔数据集包含63566张动漫风格的面部图像来自AnimeFaces项目2.2. 数据加载与预处理我们使用ImageFolder和torchvision.transforms进行数据预处理importtorchimporttorch.nnasnnfromtorch.utils.dataimportDataLoaderfromtorchvision.utilsimportmake_gridimporttorchvision.transformsasTfromtorchvision.datasetsimportImageFolderimportmatplotlib.pyplotaspltimportpandasaspdimportnumpyasnpfromtqdmimporttrange n_epochs25image_size64img_channels3batch_size64z_dim128lr1e-4n_critic1lamda_gp10fixed_latenttorch.randn(48,z_dim,devicecuda)data_path./data/AnimeFaces# data_path ./data/CelebA/img_align_celebatrain_datasetImageFolder(data_path,transformT.Compose([T.Resize(image_size),T.CenterCrop(image_size),T.ToTensor(),T.Normalize([0.5]*3,[0.5,0.5,0.5])]))n_sampleslen(train_dataset)图像被调整为64x64像素并进行归一化处理将像素值范围从[0,1]映射到[-1,1]这有助于模型更好地学习数据分布。接下来创建数据加载器并观察数据集示例train_dataloaderDataLoader(train_dataset,batch_sizebatch_size,shuffleTrue,num_workers3,pin_memoryTrue)n_batchlen(train_dataloader)#n_batch994forimgs,_intrain_dataloader:print(imgs_batch.shape,imgs.shape)breakdefdenorm(img_tensors):returnimg_tensors*0.50.5defshow_imgs(images):fig,axplt.subplots(figsize(16,12))inputmake_grid(denorm(images[:48]),nrow16)ax.imshow(input.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)创建生成器 (Generator)采用转置卷积逐步上采样# Generator classdefbasic_G(in_channles,out_channels,f4,s2,p1):returnnn.Sequential(nn.ConvTranspose2d(in_channles,out_channels,kernel_sizef,strides,paddingp,biasFalse),nn.BatchNorm2d(out_channels),nn.ReLU(True))classGenerator(nn.Module):def__init__(self):super().__init__()self.netnn.Sequential(basic_G(z_dim,512,4,1,0),basic_G(512,256,4,2,1),basic_G(256,128,4,2,1),basic_G(128,64,4,2,1),nn.ConvTranspose2d(64,3,4,2,1),nn.Tanh())defforward(self,z):inputz.view(-1,z_dim,1,1)imagesself.net(input)returnimages GGenerator().cuda()G.apply(weights_init)3.2 判别器定义判别器 (Critic)使用实例归一化和LeakyReLU激活函数defbasic_D(in_channles,out_channels,f4,s2,p1):returnnn.Sequential(nn.Conv2d(in_channles,out_channels,kernel_sizef,strides,paddingp,biasFalse),nn.InstanceNorm2d(out_channels,affineTrue),nn.LeakyReLU(0.2,inplaceTrue))classDiscriminator(nn.Module):def__init__(self):super().__init__()self.netnn.Sequential(nn.Conv2d(img_channels,64,4,2,1),nn.LeakyReLU(0.2,inplaceTrue),basic_D(64,128,4,2,1),basic_D(128,256,4,2,1),basic_D(256,512,4,2,1),nn.Conv2d(512,1,4,1,0),nn.Flatten())defforward(self,images):scalarsself.net(images)returnscalars DDiscriminator().cuda()D.apply(weights_init)3.3 梯度惩罚实现梯度惩罚是WGAN-GP的核心组件确保判别器满足Lipschitz约束defgradient_penalty(D,real_data,fake_data):batch_sizereal_data.size(0)epstorch.rand(batch_size,1,1,1).cuda()# uniform distributionepseps.expand_as(real_data)# eps.shapebatch_size x 3 x 64^2# Interpolation between real data and fake data.interpolationeps*real_data(1-eps)*fake_data logitsD(interpolation)#logits for interpolated imagesgradientstorch.autograd.grad(outputslogits,inputsinterpolation,grad_outputstorch.ones_like(logits),create_graphTrue,retain_graphTrue)[0]gradientsgradients.view(batch_size,-1)grad_normgradients.norm(2,1)gradient_penaltytorch.mean((grad_norm-1)**2)returngradient_penalty4. 训练模型定义模型优化器optimizer_Dtorch.optim.RMSprop(D.parameters(),lrlr)optimizer_Gtorch.optim.RMSprop(G.parameters(),lrlr)#optimizer_G torch.optim.Adam(G.parameters(), lrlr, betas(0.0, 0.9))#optimizer_D torch.optim.Adam(D.parameters(), lrlr, betas(0.0, 0.9))定义生成器和判别器训练函数deftrain_D(inputs,optimizer_D):for_inrange(n_critic):# The inputs are real images from a batch of DataLoader loaded in cudabatch_sizeinputs.shape[0]real_predsD(inputs)real_scoretorch.mean(real_preds)# create fake images with random numberslatenttorch.randn(batch_size,z_dim).cuda()fake_imagesG(latent)fake_predsD(fake_images.detach())fake_scoretorch.mean(fake_preds)# Update discriminator weightsgpgradient_penalty(D,inputs,fake_images)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):latenttorch.randn(batch_size,z_dim).cuda()fake_imagesG(latent)# Create fake images from latentpredsD(fake_images)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,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)show_imgs(fake_images.detach().cpu())returndf historyfit(n_epochs)关键参数设置n_critic 1每更新一次生成器更新一次判别器lambda_gp 10梯度惩罚系数使用RMSprop优化器学习率lr 1e-45. 实验结果与分析WGAN-GP的显著优势在于训练过程的稳定性。传统GAN需要精心调整超参数以避免模式坍塌而WGAN-GP通过Wasserstein距离和梯度惩罚机制大大降低了对超参数的敏感性。从训练过程曲线可以看出生成器损失和判别器损失保持相对稳定的变化趋势梯度惩罚项在整个训练过程中维持在合理范围内没有出现传统GAN常见的梯度消失或爆炸问题dfhistory 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()训练完成后可以通过以下代码生成图像n_images1ztorch.randn(n_images,z_dim).cuda()imgG(z).data.cpu()show_imgs(img)小结本节详细介绍了WGAN-GP在CelebA和动漫面孔数据集上的应用实践。实验结果表明WGAN-GP有效解决了传统GAN训练不稳定和模式坍塌的问题生成的图像质量优于DCGAN面部特征更加清晰自然训练过程稳定超参数调试工作量大大减少相关链接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