PyTorch计算机视觉——WGAN-GP在图像生成中的应用
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_epochs=25image_size=64img_channels=3batch_size=64z_dim=128lr=1e-4n_critic=1lamda_gp=10fixed_latent=torch.randn(48,z_dim,device='cuda')data_path='./data/AnimeFaces'# data_path = './data/CelebA/img_align_celeba'train_dataset=ImageFolder(data_path,transform=T.Compose([T.Resize(image_size),T.CenterCrop(image_size),T.ToTensor(),T.Normalize([0.5]*3,[0.5,0.5,0.5])]))n_samples=len(train_dataset)图像被调整为64x64像素,并进行归一化处理,将像素值范围从[0,1]映射到[-1,1],这有助于模型更好地学习数据分布。
接下来,创建数据加载器并观察数据集示例:
train_dataloader=DataLoader(train_dataset,batch_size=batch_size,shuffle=True,num_workers=3,pin_memory=True)n_batch=len(train_dataloader)#n_batch=994forimgs,_intrain_dataloader:print("imgs_batch.shape=",imgs.shape)breakdefdenorm(img_tensors):returnimg_tensors*0.5+0.5defshow_imgs(images):fig,ax=plt.subplots(figsize=(16,12))input=make_grid(denorm(images[:48]),nrow=16)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,f=4,s=2,p=1):returnnn.Sequential(nn.ConvTranspose2d(in_channles,out_channels,kernel_size=f,stride=s,padding=p,bias=False),nn.BatchNorm2d(out_channels),nn.ReLU(True))classGenerator(nn.Module):def__init__(self):super().__init__()self.net=nn.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):input=z.view(-1,z_dim,1,1)images=self.net(input)returnimages G=Generator().cuda()G.apply(weights_init)3.2 判别器
定义判别器 (Critic),使用实例归一化和LeakyReLU激活函数:
defbasic_D(in_channles,out_channels,f=4,s=2,p=1):returnnn.Sequential(nn.Conv2d(in_channles,out_channels,kernel_size=f,stride=s,padding=p,bias=False),nn.InstanceNorm2d(out_channels,affine=True),nn.LeakyReLU(0.2,inplace=True))classDiscriminator(nn.Module):def__init__(self):super().__init__()self.net=nn.Sequential(nn.Conv2d(img_channels,64,4,2,1),nn.LeakyReLU(0.2,inplace=True),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):scalars=self.net(images)returnscalars D=Discriminator().cuda()D.apply(weights_init)3.3 梯度惩罚实现
梯度惩罚是WGAN-GP的核心组件,确保判别器满足Lipschitz约束:
defgradient_penalty(D,real_data,fake_data):batch_size=real_data.size(0)eps=torch.rand(batch_size,1,1,1).cuda()# uniform distributioneps=eps.expand_as(real_data)# eps.shape=batch_size x 3 x 64^2# Interpolation between real data and fake data.interpolation=eps*real_data+(1-eps)*fake_data logits=D(interpolation)#logits for interpolated imagesgradients=torch.autograd.grad(outputs=logits,inputs=interpolation,grad_outputs=torch.ones_like(logits),create_graph=True,retain_graph=True)[0]gradients=gradients.view(batch_size,-1)grad_norm=gradients.norm(2,1)gradient_penalty=torch.mean((grad_norm-1)**2)returngradient_penalty4. 训练模型
定义模型优化器:
optimizer_D=torch.optim.RMSprop(D.parameters(),lr=lr)optimizer_G=torch.optim.RMSprop(G.parameters(),lr=lr)#optimizer_G = torch.optim.Adam(G.parameters(), lr=lr, betas=(0.0, 0.9))#optimizer_D = torch.optim.Adam(D.parameters(), lr=lr, 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_size=inputs.shape[0]real_preds=D(inputs)real_score=torch.mean(real_preds)# create fake images with random numberslatent=torch.randn(batch_size,z_dim).cuda()fake_images=G(latent)fake_preds=D(fake_images.detach())fake_score=torch.mean(fake_preds)# Update discriminator weightsgp=gradient_penalty(D,inputs,fake_images)loss=fake_score-real_score+lamda_gp*gp optimizer_D.zero_grad()loss.backward()optimizer_D.step()returnloss.item(),real_score.item(),fake_score.item()deftrain_G(optimizer_G):latent=torch.randn(batch_size,z_dim).cuda()fake_images=G(latent)# Create fake images from latentpreds=D(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 historydf=pd.DataFrame(np.empty([epochs,4]),index=np.arange(epochs),columns=['Loss_G','Loss_D','D(X)','D(G(Z))'])foriintrange(epochs):loss_G=0.0;loss_D=0.0;real_sc=0.0;fake_sc=0.0forreal_images,labelsintrain_dataloader:inputs=real_images.cuda()labels=labels.cuda()loss_d,real_score,fake_score=train_D(inputs,optimizer_D)loss_D+=loss_d;real_sc+=real_score;fake_sc+=fake_score loss_g=train_G(optimizer_G)loss_G+=loss_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_batchifi==0or(i+1)%5==0:print("Epoch={:2}, Ls_G={:.2f}, Ls_D={:.2f}, D(X)={:.2f}, D(G(Z))={:.2f}".format(i+1,df.iloc[i,0],df.iloc[i,1],df.iloc[i,2],df.iloc[i,3]))fake_images=G(fixed_latent)show_imgs(fake_images.detach().cpu())returndf history=fit(n_epochs)关键参数设置:
n_critic = 1:每更新一次生成器,更新一次判别器lambda_gp = 10:梯度惩罚系数- 使用
RMSprop优化器,学习率lr = 1e-4
5. 实验结果与分析
WGAN-GP的显著优势在于训练过程的稳定性。传统GAN需要精心调整超参数以避免模式坍塌,而WGAN-GP通过Wasserstein距离和梯度惩罚机制,大大降低了对超参数的敏感性。从训练过程曲线可以看出:
- 生成器损失和判别器损失保持相对稳定的变化趋势
- 梯度惩罚项在整个训练过程中维持在合理范围内
- 没有出现传统
GAN常见的梯度消失或爆炸问题
df=history fig,ax=plt.subplots(1,2,figsize=(9,4),sharex=True)df.plot(ax=ax[0],y=[0,1],style=['r-','b-+'])gp=df.iloc[:,1]-df.iloc[:,3]+df.iloc[:,2]ax[0].plot(gp,label='Gradient Penalty',color='k',linestyle=':')ax[0].set(ylabel='loss')ax[0].legend()df.plot(ax=ax[1],y=[2,3],style=['r-+','b-'])foriinrange(2):ax[i].grid(which='major',axis='both',color='g',linestyle=':')ax[i].set(xlabel='epoch')plt.show()训练完成后,可以通过以下代码生成图像:
n_images=1z=torch.randn(n_images,z_dim).cuda()img=G(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 Network,GAN)
PyTorch计算机视觉(6)——深度卷积对抗神经网络(DCGAN)
PyTorch计算机视觉(7)——条件生成对抗网络(cGAN)
PyTorch计算机视觉(8)——WGAN及其变体WGAN-GP
