共计 3342 个字符,预计需要花费 9 分钟才能阅读完成。
背景介绍
GAN(Generative Adversarial Network,生成对抗网络)是一种强大的生成模型,由生成器(Generator)和判别器(Discriminator)两部分组成。生成器负责生成假数据,判别器负责判断数据是真实的还是生成的。两者通过对抗训练不断提升性能,最终生成高质量的数据。GAN 广泛应用于图像生成、风格迁移、数据增强等领域。

环境配置
搭建 GAN 模型需要以下 Python 库和工具:
- Python 3.7+
- PyTorch 1.8+
- torchvision
- matplotlib
- numpy
安装命令如下:
pip install torch torchvision matplotlib numpy
数据集准备
我们以 MNIST 数据集为例,演示如何加载和处理数据。MNIST 是一个手写数字数据集,包含 60000 张训练图片和 10000 张测试图片。
import torch
from torchvision import datasets, transforms
# 定义数据预处理
transform = transforms.Compose([transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
# 加载数据集
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=64, shuffle=True)
模型架构
生成器
生成器的作用是将随机噪声转换为与真实数据相似的样本。我们使用全连接层实现一个简单的生成器。
import torch.nn as nn
class Generator(nn.Module):
def __init__(self, input_dim=100, output_dim=784):
super(Generator, self).__init__()
self.fc = nn.Sequential(nn.Linear(input_dim, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, output_dim),
nn.Tanh())
def forward(self, x):
return self.fc(x)
判别器
判别器的作用是判断输入数据是真实的还是生成的。
class Discriminator(nn.Module):
def __init__(self, input_dim=784):
super(Discriminator, self).__init__()
self.fc = nn.Sequential(nn.Linear(input_dim, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 1),
nn.Sigmoid())
def forward(self, x):
return self.fc(x)
训练过程
初始化模型和优化器
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# 初始化模型
generator = Generator().to(device)
discriminator = Discriminator().to(device)
# 定义损失函数和优化器
criterion = nn.BCELoss()
g_optimizer = torch.optim.Adam(generator.parameters(), lr=0.0002)
d_optimizer = torch.optim.Adam(discriminator.parameters(), lr=0.0002)
训练循环
num_epochs = 50
for epoch in range(num_epochs):
for i, (real_images, _) in enumerate(train_loader):
batch_size = real_images.size(0)
real_images = real_images.view(batch_size, -1).to(device)
# 训练判别器
d_optimizer.zero_grad()
# 真实数据
real_labels = torch.ones(batch_size, 1).to(device)
real_outputs = discriminator(real_images)
d_loss_real = criterion(real_outputs, real_labels)
# 生成数据
noise = torch.randn(batch_size, 100).to(device)
fake_images = generator(noise)
fake_labels = torch.zeros(batch_size, 1).to(device)
fake_outputs = discriminator(fake_images.detach())
d_loss_fake = criterion(fake_outputs, fake_labels)
# 总损失
d_loss = d_loss_real + d_loss_fake
d_loss.backward()
d_optimizer.step()
# 训练生成器
g_optimizer.zero_grad()
# 让生成的数据尽可能被判别为真
fake_outputs = discriminator(fake_images)
g_loss = criterion(fake_outputs, real_labels)
g_loss.backward()
g_optimizer.step()
print(f'Epoch [{epoch+1}/{num_epochs}], d_loss: {d_loss.item():.4f}, g_loss: {g_loss.item():.4f}')
性能优化
-
学习率调整 :学习率太大可能导致模型不稳定,太小则收敛慢。可以使用学习率调度器动态调整。
-
批次大小 :较大的批次可以提高训练稳定性,但会消耗更多内存。
-
网络深度 :增加网络深度可以提高模型能力,但也会增加训练难度。
-
激活函数 :生成器最后一层使用 Tanh,中间层使用 LeakyReLU;判别器使用 Sigmoid 输出。
避坑指南
-
模式崩溃 :生成器只生成少数几种样本。解决方法:增加判别器能力,使用不同的损失函数。
-
训练不稳定 :损失值剧烈波动。解决方法:调整学习率,使用梯度裁剪。
-
生成质量差 :样本模糊或不符合预期。解决方法:增加训练轮次,调整网络结构。
结果验证
训练完成后,我们可以生成一些样本并可视化:
import matplotlib.pyplot as plt
# 生成样本
noise = torch.randn(16, 100).to(device)
generated_images = generator(noise).cpu().detach().numpy()
# 可视化
fig, axes = plt.subplots(4, 4, figsize=(8, 8))
for i, ax in enumerate(axes.flatten()):
ax.imshow(generated_images[i].reshape(28, 28), cmap='gray')
ax.axis('off')
plt.show()
延伸阅读
- 原始 GAN 论文:Generative Adversarial Networks by Ian Goodfellow et al.
- DCGAN:Deep Convolutional GAN
- WGAN:Wasserstein GAN
练习题
- 尝试使用 CIFAR-10 数据集训练 GAN。
- 修改网络结构,使用卷积层代替全连接层。
- 实现学习率调度器,动态调整学习率。
希望这篇教程能帮助你快速上手 GAN 模型的搭建和训练。GAN 是一个强大但需要耐心调试的工具,多尝试不同的参数和结构,你会得到更好的结果。
