共计 3677 个字符,预计需要花费 10 分钟才能阅读完成。
从零搭建 GAN 生成对抗网络:环境配置、数据集加载与模型训练实战
背景介绍
GAN(Generative Adversarial Network,生成对抗网络)是一种强大的生成模型,由生成器(Generator)和判别器(Discriminator)两部分组成。生成器的任务是生成逼真的数据,而判别器的任务是区分生成的数据和真实数据。两者通过对抗训练不断提升性能。

GAN 的应用场景非常广泛,包括图像生成、风格迁移、超分辨率重建等。然而,初学者在搭建 GAN 时常常会遇到以下问题:
- 环境配置复杂:依赖库版本冲突、CUDA 与 GPU 驱动不匹配等。
- 训练不稳定:模式崩溃(Mode Collapse)、梯度消失或爆炸。
- 数据集处理困难:数据标准化、分批加载等预处理步骤容易出错。
本教程将一步步带你解决这些问题,从环境配置到模型训练,让你快速上手 GAN 实战。
环境配置
首先,我们需要创建一个独立的 Python 虚拟环境,避免依赖冲突。推荐使用 conda 管理环境:
conda create -n gan_env python=3.8
conda activate gan_env
接下来,安装必要的依赖库。这里以 PyTorch 为例:
conda install pytorch torchvision torchaudio cudatoolkit=11.3 -c pytorch
pip install matplotlib numpy tqdm
- PyTorch:深度学习框架,版本 1.10.0。
- torchvision:提供常用数据集和图像变换工具。
- CUDA 11.3:确保 GPU 加速可用(需显卡支持)。
- matplotlib/numpy:用于数据可视化和数值计算。
- tqdm:显示训练进度条。
数据集处理
我们以 MNIST 手写数字数据集为例,演示如何加载和预处理数据。
import torch
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
# 数据预处理:标准化到 [-1, 1] 范围
transform = transforms.Compose([transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)) # 均值 0.5,标准差 0.5
])
# 加载 MNIST 数据集
train_dataset = datasets.MNIST(
root='./data',
train=True,
download=True,
transform=transform
)
# 分批加载数据
train_loader = DataLoader(
train_dataset,
batch_size=64,
shuffle=True
)
关键点说明:
- 标准化 :将像素值从[0, 1] 映射到[-1, 1],有利于 GAN 训练的稳定性。
- 分批加载 :使用
DataLoader实现数据的按批加载和打乱顺序。
模型搭建
生成器(Generator)
生成器的任务是将随机噪声(潜在空间向量)转换为逼真的图像。以下是一个简单的全连接网络结构:
import torch.nn as nn
class Generator(nn.Module):
def __init__(self, latent_dim=100, img_dim=784):
super(Generator, self).__init__()
self.model = nn.Sequential(nn.Linear(latent_dim, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, img_dim),
nn.Tanh() # 输出范围[-1, 1]
)
def forward(self, z):
return self.model(z)
关键层说明:
LeakyReLU:解决梯度消失问题,负值斜率设为 0.2。Tanh:将输出限制在[-1, 1],与标准化后的数据范围一致。
判别器(Discriminator)
判别器的任务是判断输入图像是真实的还是生成的。以下是判别器的代码:
class Discriminator(nn.Module):
def __init__(self, img_dim=784):
super(Discriminator, self).__init__()
self.model = nn.Sequential(nn.Linear(img_dim, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 1),
nn.Sigmoid() # 输出 0 到 1 的概率)
def forward(self, img):
return self.model(img)
关键层说明:
Sigmoid:输出 0 到 1 的概率,表示图像为真的可能性。
训练过程
损失函数与优化器
GAN 的训练需要定义生成器和判别器的损失函数,并分别配置优化器:
import torch.optim as optim
# 初始化模型
generator = Generator()
discriminator = Discriminator()
# 二元交叉熵损失
criterion = nn.BCELoss()
# 优化器
lr = 0.0002
optimizer_G = optim.Adam(generator.parameters(), lr=lr)
optimizer_D = optim.Adam(discriminator.parameters(), lr=lr)
训练循环
以下是训练的核心代码,包含生成器和判别器的交替训练:
num_epochs = 50
for epoch in range(num_epochs):
for i, (real_imgs, _) in enumerate(train_loader):
batch_size = real_imgs.size(0)
real_imgs = real_imgs.view(batch_size, -1)
# 真实标签和生成标签
real_labels = torch.ones(batch_size, 1)
fake_labels = torch.zeros(batch_size, 1)
# 训练判别器
optimizer_D.zero_grad()
# 真实图像的损失
outputs = discriminator(real_imgs)
d_loss_real = criterion(outputs, real_labels)
# 生成图像的损失
z = torch.randn(batch_size, 100) # 随机噪声
fake_imgs = generator(z)
outputs = discriminator(fake_imgs.detach())
d_loss_fake = criterion(outputs, fake_labels)
# 总判别器损失
d_loss = d_loss_real + d_loss_fake
d_loss.backward()
optimizer_D.step()
# 训练生成器
optimizer_G.zero_grad()
outputs = discriminator(fake_imgs)
g_loss = criterion(outputs, real_labels) # 生成器希望判别器认为生成图像是真的
g_loss.backward()
optimizer_G.step()
# 打印训练状态
print(f'Epoch [{epoch+1}/{num_epochs}], d_loss: {d_loss.item():.4f}, g_loss: {g_loss.item():.4f}')
训练监控:
- 损失曲线 :观察
d_loss和g_loss是否收敛。 - 生成样本可视化:定期保存生成器输出的图像,检查生成质量。
避坑指南
- 模式崩溃(Mode Collapse):生成器只生成单一类型的样本。
-
解决方案:使用 Mini-batch Discrimination 或 Wasserstein GAN(WGAN)。
-
梯度消失或爆炸:判别器或生成器的梯度异常。
-
解决方案:使用梯度裁剪(Gradient Clipping)或调整学习率。
-
训练不稳定:判别器或生成器一方过于强势。
-
解决方案:平衡两者的训练频率(如判别器训练 5 次,生成器训练 1 次)。
-
生成图像模糊:损失函数不适合。
-
解决方案:改用 L1/L2 损失或感知损失(Perceptual Loss)。
-
数据预处理不当:输入范围与激活函数不匹配。
- 解决方案:确保数据标准化与输出激活函数(如 Tanh)的范围一致。
延伸思考
- 改进网络结构:尝试 DCGAN(深度卷积 GAN),使用卷积层提升生成质量。
- 添加正则化:在判别器中加入梯度惩罚(Gradient Penalty),提升训练稳定性。
- 探索其他 GAN 变体:如 CycleGAN(图像风格迁移)、StyleGAN(高分辨率生成)。
参考资源
希望这篇教程能帮助你顺利搭建第一个 GAN 模型!如果有任何问题,欢迎在评论区交流。
