0. 前言

生成对抗网络 (Generative Adversarial Network, GAN) 2014 年由 Ian Goodfellow 提出以来,已成为深度学习领域最具创新性的技术之一。然而,原始 GAN 面临着训练不稳定、模式坍塌等挑战,这些问题限制了其在实际应用中的效果。Wasserstein GAN with Gradient Penalty (WGAN-GP) 作为一种改进方案,通过引入 Wasserstein 距离和梯度惩罚项,有效解决了这些问题。本节将使用 WGAN-GPCelebA 人脸数据集和动漫面孔数据集上实现图像生成,包括代码实现、训练过程分析以及结果评估。

1. WGAN-GP 技术原理简述

WGAN-GP 的核心创新在于:

  1. Wasserstein 距离:替代传统 GAN 使用的 JS 散度,提供更平滑的梯度,使训练过程更加稳定
  2. 梯度惩罚 (Gradient Penalty):强制判别器 (Critic) 的梯度范数接近 1,满足 Lipschitz 约束条件
  3. 弃用批归一化:在判别器中使用实例归一化 (Instance Normalization) 替代批归一化,避免批次内样本间的相互影响

这些改进使得 WGAN-GP 对超参数的选择不那么敏感,减少了模式坍塌的风险。

2. 数据集分析

2.1 数据集简介

  • CelebA 数据集:包含 202599 张名人面部图像,广泛用于人脸识别和生成任务
  • 动漫面孔数据集:包含 63566 张动漫风格的面部图像,来自 AnimeFaces 项目

2.2. 数据加载与预处理

我们使用 ImageFoldertorchvision.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-GPCelebA 和动漫面孔数据集上的应用实践。实验结果表明:

  1. WGAN-GP 有效解决了传统 GAN 训练不稳定和模式坍塌的问题
  2. 生成的图像质量优于 DCGAN,面部特征更加清晰自然
  3. 训练过程稳定,超参数调试工作量大大减少

相关链接

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

更多推荐