PyTorch计算机视觉——WGAN-GP在图像生成中的应用
PyTorch计算机视觉——WGAN-GP在图像生成中的应用
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 进行数据预处理:
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torchvision.utils import make_grid
import torchvision.transforms as T
from torchvision.datasets import ImageFolder
import matplotlib.pyplot as plt
import pandas as pd
import numpy as np
from tqdm import trange
n_epochs = 25
image_size = 64
img_channels = 3
batch_size = 64
z_dim = 128
lr = 1e-4
n_critic = 1
lamda_gp = 10
fixed_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=994
for imgs, _ in train_dataloader:
print("imgs_batch.shape=", imgs.shape)
break
def denorm(img_tensors):
return img_tensors*0.5 + 0.5
def show_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(),用于模型参数初始化:
def weights_init(m):
if(type(m) == nn.ConvTranspose2d or type(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 class
def basic_G(in_channles, out_channels, f=4, s=2, p=1):
return nn.Sequential(
nn.ConvTranspose2d(in_channles, out_channels,
kernel_size=f, stride=s, padding=p, bias=False),
nn.BatchNorm2d(out_channels),
nn.ReLU(True)
)
class Generator(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()
)
def forward(self, z):
input = z.view(-1, z_dim, 1, 1)
images = self.net(input)
return images
G = Generator().cuda()
G.apply(weights_init)
3.2 判别器
定义判别器 (Critic),使用实例归一化和 LeakyReLU 激活函数:
def basic_D(in_channles, out_channels, f=4, s=2, p=1):
return nn.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)
)
class Discriminator(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()
)
def forward(self, images):
scalars = self.net(images)
return scalars
D = Discriminator().cuda()
D.apply(weights_init)
3.3 梯度惩罚实现
梯度惩罚是 WGAN-GP 的核心组件,确保判别器满足 Lipschitz 约束:
def gradient_penalty(D, real_data, fake_data):
batch_size = real_data.size(0)
eps = torch.rand(batch_size, 1, 1, 1).cuda()# uniform distribution
eps = 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 images
gradients = 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)
return gradient_penalty
4. 训练模型
定义模型优化器:
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))
定义生成器和判别器训练函数:
def train_D(inputs, optimizer_D):
for _ in range(n_critic):
# The inputs are real images from a batch of DataLoader loaded in cuda
batch_size = inputs.shape[0]
real_preds = D(inputs)
real_score= torch.mean(real_preds)
# create fake images with random numbers
latent = 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 weights
gp = gradient_penalty(D, inputs, fake_images)
loss = fake_score - real_score + lamda_gp*gp
optimizer_D.zero_grad()
loss.backward()
optimizer_D.step()
return loss.item(), real_score.item(), fake_score.item()
def train_G(optimizer_G):
latent = torch.randn(batch_size, z_dim).cuda()
fake_images = G(latent) # Create fake images from latent
preds = D(fake_images)
loss = -torch.mean(preds)
optimizer_G.zero_grad()
loss.backward()
optimizer_G.step()
return loss.item()
训练过程包括交替更新判别器和生成器:
def fit(epochs):
torch.cuda.empty_cache()
# The DataFrame df is a recorder of the training history
df = pd.DataFrame(np.empty([epochs, 4]),
index = np.arange(epochs),
columns=['Loss_G', 'Loss_D', 'D(X)', 'D(G(Z))'])
for i in trange(epochs):
loss_G = 0.0; loss_D = 0.0; real_sc = 0.0; fake_sc = 0.0
for real_images, labels in train_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 & scores
df.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_batch
if i==0 or (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())
return df
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-'])
for i in range(2):
ax[i].grid(which='major', axis='both', color='g', linestyle=':')
ax[i].set(xlabel='epoch')
plt.show()

训练完成后,可以通过以下代码生成图像:
n_images = 1
z = 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
更多推荐



所有评论(0)