从零实现GAN:手把手教你用PyTorch训练图像生成模型
1. 从零开始:为什么GAN是图像生成的“炼金术”?
如果你对AI生成图片感兴趣,刷到过那些真假难辨的人脸、风格奇特的画作,那你大概率已经接触过GAN(生成对抗网络)的成果了。但你可能也听过,GAN训练起来“玄学”得很,动不动就崩,效果时好时坏。今天,我们不谈那些高深的理论,就从最实际的角度出发,手把手带你从零开始,用代码“炼”出一张属于自己的生成图片。我会把整个过程掰开揉碎,告诉你每一步在干什么、为什么这么干,以及我踩过的那些坑。我们的目标不是复现一篇顶会论文,而是让你能真正跑通一个GAN模型,看到它从生成一片噪声到逐渐“学会”画图的全过程,并理解这背后的“道”与“术”。
首先,得说清楚GAN到底是什么。你可以把它想象成一个造假币的罪犯(生成器G)和一个经验老道的警察(判别器D)之间的猫鼠游戏。生成器的目标是造出以假乱真的假币,让警察认不出来;判别器的目标是练就火眼金睛,准确分辨真币和假币。两者在不断的对抗中共同进化:造假者技术越来越高超,警察的鉴伪能力也越来越强。最终,当警察再也无法区分真假时,我们就得到了一个强大的造假者——也就是我们的图像生成模型。这个核心思想,就是GAN一切魅力和难点的根源。
那么,为什么我们要“从零开始”呢?现在有很多现成的工具,比如一些在线平台或封装好的库,输入关键词就能出图。但依赖这些工具,就像只会开车却不懂发动机原理,一旦抛锚就束手无策。从零实现,意味着你能掌控数据预处理、模型架构设计、损失函数计算、训练循环调试每一个环节。当生成图片一片模糊或者直接崩掉时,你才知道该从哪里入手调整。这对于想深入AI生成领域,甚至未来想自己设计新模型的人来说,是必不可少的一课。
接下来,我们会用最经典的DCGAN(深度卷积生成对抗网络)结构,在MNIST手写数字数据集上实战。选择它们是因为结构清晰、数据简单,非常适合入门理解。整个旅程包括:搭建开发环境、理解并准备数据、亲手编写生成器和判别器网络、实现那个精妙的对抗训练过程、最后启动训练并观察和分析结果。我会穿插大量代码和注释,并解释每一个超参数选择的理由。准备好了吗?我们开始“炼丹”。
2. 环境搭建与数据准备:给“炼金术”准备坩埚和原料
工欲善其事,必先利其器。第一步,我们要搭建一个稳定、可复现的Python开发环境。我强烈推荐使用Anaconda来管理环境,它能很好地解决不同项目间包版本冲突的问题。
2.1 创建并配置专属的Python环境
打开你的终端(或Anaconda Prompt),执行以下命令来创建一个新的环境,我将其命名为 gan_study ,并指定Python版本为3.8(这是一个在深度学习领域兼容性非常好的版本):
conda create -n gan_study python=3.8
conda activate gan_study
环境激活后,我们来安装核心的深度学习框架。这里我选择PyTorch,因为它动态图机制对研究和实验非常友好,调试起来直观。你需要根据你的电脑是否有NVIDIA显卡(以及CUDA版本)去PyTorch官网选择对应的安装命令。如果没有显卡,就安装CPU版本。以下是一个安装PyTorch 1.12.0(CPU版本)及一些必要工具库的命令示例:
pip install torch==1.12.0 torchvision==0.13.0 --index-url https://download.pytorch.org/whl/cpu
pip install matplotlib numpy tqdm jupyter
注意:如果你有GPU且安装了CUDA,请务必使用官网命令安装对应的CUDA版本(如
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118)。GPU训练速度通常是CPU的几十倍以上,能极大节省实验时间。安装后,可以在Python中运行import torch; print(torch.cuda.is_available())来验证GPU是否可用。
2.2 理解并加载MNIST数据集
我们选用MNIST数据集作为“原料”。它包含6万张28x28像素的灰度手写数字图片(0-9),结构简单,训练速度快,非常适合GAN的入门实验。使用 torchvision 库可以非常方便地下载和加载它。
import torch
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
# 定义数据预处理变换
# 1. ToTensor(): 将PIL图像或NumPy数组转换为PyTorch张量,并自动缩放到[0,1]范围。
# 2. Normalize(mean, std): 进行标准化,这里均值和标准差设为0.5,是为了将数据范围从[0,1]映射到[-1,1]。
# 为什么是[-1,1]?因为GAN的生成器通常使用tanh作为输出层激活函数,其值域就是[-1,1],这样数据分布和模型输出能对齐。
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)) # 对于单通道灰度图,均值和标准差都是一个值
])
# 下载并加载训练数据集
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
# 创建数据加载器,batch_size是关键参数。
# batch_size太小,梯度更新噪声大,训练不稳定;太大,可能内存不够,且降低了梯度更新的频率。
# 对于GAN,通常可以从64或128开始尝试。这里我们设为128。
train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=2)
# 我们可以可视化看一下数据
import matplotlib.pyplot as plt
def show_images(images):
fig, axes = plt.subplots(1, 5, figsize=(10, 2))
for i, ax in enumerate(axes):
# 反标准化:将[-1,1]映射回[0,1]以便显示
img = images[i].squeeze().numpy() * 0.5 + 0.5
ax.imshow(img, cmap='gray')
ax.axis('off')
plt.show()
# 获取一个批次的数据
data_iter = iter(train_loader)
real_images, _ = next(data_iter) # _ 是标签,我们生成任务暂时用不到
show_images(real_images[:5])
运行这段代码,你应该能看到5个清晰的手写数字。这一步确保了我们的“原料”是正确可用的。数据加载器 train_loader 会在训练时,每次给我们提供一批(128张)处理好的图像张量,形状为 [128, 1, 28, 28] (批大小,通道数,高度,宽度)。
3. 构建对抗双方:生成器与判别器的网络设计
现在,我们来打造“造假者”(生成器G)和“警察”(判别器D)。DCGAN提出了一系列设计准则,使得训练更加稳定,我们遵循这些准则来构建网络。
3.1 生成器(Generator)的设计:从噪声到图像
生成器的任务是接收一个随机噪声向量(通常从标准正态分布中采样),然后通过一系列上采样(反卷积)层,将其“翻译”成一张图片。你可以把它想象成一个“解码器”。
import torch.nn as nn
class Generator(nn.Module):
def __init__(self, latent_dim=100, img_channels=1):
super(Generator, self).__init__()
self.latent_dim = latent_dim
self.init_size = 7 # 初始特征图大小。我们从100维噪声开始,先映射到一个7x7的特征图。
self.fc = nn.Linear(latent_dim, 128 * self.init_size ** 2) # 全连接层,将噪声展开
self.model = nn.Sequential(
nn.BatchNorm2d(128), # 批归一化,稳定训练,加速收敛
nn.Upsample(scale_factor=2), # 上采样,7x7 -> 14x14
nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1), # 卷积层,提炼特征
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.Upsample(scale_factor=2), # 上采样,14x14 -> 28x28
nn.Conv2d(128, 64, kernel_size=3, stride=1, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.Conv2d(64, img_channels, kernel_size=3, stride=1, padding=1),
nn.Tanh() # 输出层,将值约束在[-1,1],与我们的数据标准化范围对应
)
def forward(self, z):
# z的形状: [batch_size, latent_dim]
out = self.fc(z) # 全连接层输出: [batch_size, 128*7*7]
out = out.view(out.shape[0], 128, self.init_size, self.init_size) # 重塑为 [batch_size, 128, 7, 7]
img = self.model(out) # 通过卷积上采样网络
return img # 输出形状: [batch_size, img_channels, 28, 28]
关键点解析:
-
latent_dim(潜在空间维度) :这是输入噪声的维度,通常设为100。它代表了生成图像的“创意空间”大小。维度太低,模型表达能力不足;太高,可能增加训练难度和不确定性。 - 上采样 vs 转置卷积 :早期常用转置卷积(
nn.ConvTranspose2d),但它容易产生棋盘伪影。这里我使用了Upsample(最近邻或双线性插值)后接普通卷积的组合,这是现在更推荐的做法,能生成质量更高的图像。 - 批归一化(BatchNorm) :在生成器的每一层(除了输出层)之后使用,它有助于缓解训练初期因输入数据分布差异大导致的梯度问题,是稳定GAN训练的关键技术之一。
- 激活函数 :隐藏层使用ReLU,输出层使用Tanh,这与我们数据标准化到[-1,1]是匹配的。
3.2 判别器(Discriminator)的设计:火眼金睛的鉴伪专家
判别器是一个标准的二分类卷积神经网络。输入一张图片(真实或生成的),输出一个标量,代表该图片为“真”的概率。
class Discriminator(nn.Module):
def __init__(self, img_channels=1):
super(Discriminator, self).__init__()
self.model = nn.Sequential(
# 输入: [batch_size, img_channels, 28, 28]
nn.Conv2d(img_channels, 64, kernel_size=4, stride=2, padding=1), # 28x28 -> 14x14
nn.LeakyReLU(0.2, inplace=True),
nn.Conv2d(64, 128, kernel_size=4, stride=2, padding=1), # 14x14 -> 7x7
nn.BatchNorm2d(128),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv2d(128, 256, kernel_size=3, stride=1, padding=1), # 保持7x7
nn.BatchNorm2d(256),
nn.LeakyReLU(0.2, inplace=True),
nn.Flatten(), # 将特征图展平
nn.Linear(256 * 7 * 7, 1), # 全连接层,输出一个值
nn.Sigmoid() # 使用Sigmoid将输出映射到[0,1],表示概率
)
def forward(self, img):
# img的形状: [batch_size, img_channels, 28, 28]
validity = self.model(img) # 输出形状: [batch_size, 1]
return validity
关键点解析:
- LeakyReLU :判别器中使用LeakyReLU而不是普通的ReLU,主要是为了防止梯度稀疏。当输入为负数时,LeakyReLU有一个很小的斜率(如0.2),允许梯度反向传播,避免了“神经元死亡”问题,这在判别器中尤为重要。
- 下采样 :通过卷积的步长(stride=2)实现,逐步降低空间维度,增加通道数,以提取更高层次的特征。
- 输出层 :一个全连接层接Sigmoid,输出一个介于0到1之间的概率值。1代表判别器认为图片是真实的,0代表认为是生成的。
- 判别器中没有在输入层使用BatchNorm :这是一个常见的技巧。有些研究指出,在判别器的第一层不使用BatchNorm可能有助于学习到更正确的数据分布统计量。
现在,我们可以初始化这两个网络,并看一下它们的结构概要。
# 初始化模型
latent_dim = 100
generator = Generator(latent_dim=latent_dim)
discriminator = Discriminator()
# 打印模型结构,查看参数数量
print(generator)
print(f"\nGenerator参数数量: {sum(p.numel() for p in generator.parameters())}")
print("\n" + "="*50 + "\n")
print(discriminator)
print(f"\nDiscriminator参数数量: {sum(p.numel() for p in discriminator.parameters())}")
# 简单测试前向传播
z = torch.randn(4, latent_dim) # 4个随机噪声
fake_imgs = generator(z)
print(f"\n生成图片形状: {fake_imgs.shape}") # 应为 [4, 1, 28, 28]
d_real = discriminator(real_images[:4])
d_fake = discriminator(fake_imgs.detach()) # 使用.detach()切断与生成器的计算图
print(f"判别器对真实图片的输出(概率): {d_real}")
print(f"判别器对生成图片的输出(概率): {d_fake}")
如果一切正常,你会看到模型结构打印出来,并且生成器成功输出了4张“图片”(虽然初期只是随机噪声),判别器也对它们给出了概率判断(初期判别器未经训练,输出应该接近0.5,即随机猜测)。
4. 定义损失函数与优化器:制定游戏规则与学习策略
模型搭好了,接下来要定义它们如何“学习”,也就是损失函数和优化器。这是GAN训练中最核心、也最微妙的部分。
4.1 对抗损失:二元交叉熵损失
GAN的训练目标可以形式化为一个极小极大博弈。对应到我们的代码,判别器D试图最大化它判断真实图片和识别假图片的能力,而生成器G试图最小化判别器识别出假图片的能力(或者说,最大化判别器将假图片误判为真的概率)。这个目标通常用二元交叉熵损失(BCELoss)来实现。
# 定义损失函数
adversarial_loss = nn.BCELoss()
# 定义优化器
# 使用Adam优化器,它是目前训练GAN最常用的选择,结合了动量和自适应学习率。
# 关键参数:学习率(lr)和动量项(betas)。
lr = 0.0002 # 学习率,一个比较通用的起始值。GAN对学习率很敏感,通常设置较小。
beta1 = 0.5 # Adam优化器的第一个动量衰减率,这是GAN论文中常用的一个经验值。
beta2 = 0.999 # 第二个动量衰减率,通常保持默认。
optimizer_G = torch.optim.Adam(generator.parameters(), lr=lr, betas=(beta1, beta2))
optimizer_D = torch.optim.Adam(discriminator.parameters(), lr=lr, betas=(beta1, beta2))
为什么是Adam且beta1=0.5? 在原始的DCGAN论文中,作者发现使用Adam优化器,并将 beta1 设为0.5(而不是默认的0.9)能够帮助稳定训练。 beta1 控制了一阶动量(梯度均值)的衰减速度,设为0.5意味着更少的动量,更新更依赖于当前梯度,这可能有助于在对抗的动态过程中更快地调整方向。
4.2 训练标签的“小把戏”
在计算损失时,我们需要为真实图片和生成图片分配标签。这里有一个常见的技巧:
# 假设batch_size = 128
batch_size = real_images.size(0)
# 为真实图片创建标签“1”,为生成图片创建标签“0”
valid = torch.ones(batch_size, 1, requires_grad=False) # 形状 [128, 1],值全为1
fake = torch.zeros(batch_size, 1, requires_grad=False) # 形状 [128, 1],值全为0
# 有时会使用“软标签”或“单侧标签平滑”来稳定训练,例如将真实的标签设为0.9,假的标签设为0.1。
# valid = torch.full((batch_size, 1), 0.9, requires_grad=False)
# fake = torch.full((batch_size, 1), 0.1, requires_grad=False)
标签平滑(Label Smoothing) :这是一个正则化技巧。如果一直用硬标签(1和0),判别器可能会对自己的判断过于自信,导致梯度消失(梯度变得非常小),生成器学不到东西。将真实标签设为略低于1(如0.9),假标签设为略高于0(如0.1),可以缓解这个问题,让判别器始终保持一定的“不确定性”,从而为生成器提供更有用的梯度信号。在训练初期,我们可以先使用硬标签,如果发现训练不稳定(如判别器loss迅速降到0),再尝试启用标签平滑。
5. 核心训练循环:上演警察与造假者的动态博弈
这是整个项目最激动人心的部分。我们将在一个循环中,交替训练判别器和生成器。理解这个循环的每一步至关重要。
import time
from torchvision.utils import save_image
# 训练参数
num_epochs = 50 # 训练轮数。对于MNIST,50轮通常能看到不错的效果。
sample_interval = 200 # 每隔多少批次保存一次生成样本
latent_dim = 100
# 记录损失,用于后续绘制曲线
G_losses = []
D_losses = []
for epoch in range(num_epochs):
start_time = time.time()
for i, (real_imgs, _) in enumerate(train_loader): # 遍历数据加载器
batch_size = real_imgs.size(0)
# 配置真实和假的标签
valid = torch.ones(batch_size, 1, requires_grad=False)
fake = torch.zeros(batch_size, 1, requires_grad=False)
# ---------------------
# 训练判别器 (D)
# ---------------------
optimizer_D.zero_grad() # 清空判别器梯度
# 计算真实图片的损失
real_pred = discriminator(real_imgs) # 判别器对真实图片的判断
d_real_loss = adversarial_loss(real_pred, valid) # 希望判别器输出接近1
# 计算生成图片的损失
z = torch.randn(batch_size, latent_dim) # 生成随机噪声
gen_imgs = generator(z) # 生成器生成假图片
fake_pred = discriminator(gen_imgs.detach()) # 关键:detach(),切断假图片与生成器的连接
d_fake_loss = adversarial_loss(fake_pred, fake) # 希望判别器对假图片输出接近0
# 判别器总损失
d_loss = (d_real_loss + d_fake_loss) / 2
d_loss.backward() # 反向传播,计算梯度
optimizer_D.step() # 更新判别器参数
# ---------------------
# 训练生成器 (G)
# ---------------------
optimizer_G.zero_grad() # 清空生成器梯度
# 生成新的噪声(也可以复用之前的,但重新生成更清晰)
z = torch.randn(batch_size, latent_dim)
gen_imgs = generator(z)
# 生成器的目标:让判别器对生成的图片输出接近1(即判别器认为它是真的)
g_pred = discriminator(gen_imgs) # 注意这里没有detach!
g_loss = adversarial_loss(g_pred, valid) # 希望判别器的输出接近“真”标签
g_loss.backward() # 反向传播,梯度会通过判别器一直传递到生成器
optimizer_G.step() # 更新生成器参数
# 记录损失
D_losses.append(d_loss.item())
G_losses.append(g_loss.item())
# 定期打印和保存样本
batches_done = epoch * len(train_loader) + i
if batches_done % sample_interval == 0:
print(f"[Epoch {epoch}/{num_epochs}] [Batch {i}/{len(train_loader)}] "
f"[D loss: {d_loss.item():.4f}] [G loss: {g_loss.item():.4f}]")
# 保存一组生成图片
save_image(gen_imgs.data[:25], f"images/{batches_done}.png", nrow=5, normalize=True)
# normalize=True会自动将图片从[-1,1]映射回[0,1]保存
epoch_time = time.time() - start_time
print(f"Epoch {epoch} 完成,耗时 {epoch_time:.2f}秒")
print("训练完成!")
训练循环的逐行解读与避坑指南:
-
optimizer.zero_grad(): 这是必须的! 在每次计算新梯度前,必须将上一轮迭代中累积的梯度清零。否则梯度会不断累加,导致训练失控。 -
判别器训练中的
detach():gen_imgs.detach()是 关键操作 。在训练判别器时,我们只希望更新判别器自身的参数。detach()方法将gen_imgs从当前的计算图中分离出来,创建一个新的、不需要梯度的张量。这样,当计算d_fake_loss并反向传播时,梯度只会更新判别器D,而不会影响到生成器G。如果不加detach(),梯度会沿着gen_imgs回溯到生成器G,这违背了“固定G,训练D”的步骤。 -
生成器训练中不使用
detach():在训练生成器时,我们的目标是让判别器D对生成图片gen_imgs的输出概率g_pred接近1。因此,我们需要梯度从判别器的损失g_loss一路回溯,经过判别器D,最终到达生成器G,从而更新G的参数。所以这里绝对不能detach。 -
损失计算顺序 :先训练判别器一步,再训练生成器一步,这是最经典的交替训练方式。也有研究尝试一次训练判别器多次再训练一次生成器(例如D训练5次,G训练1次),以保持判别器的能力略强于生成器,这在某些复杂数据集上可能更稳定。对于MNIST,1:1的交替通常就够了。
-
损失值解读 :
d_real_loss和d_fake_loss都下降,且d_loss稳定在某个值(比如0.5左右),说明判别器在学习。g_loss下降,说明生成器在进步,它生成的图片越来越能骗过判别器。- 一个健康的训练过程 :D_loss和G_loss都在波动,但整体呈下降趋势,最终达到一个动态平衡。如果D_loss迅速降到接近0,而G_loss居高不下,说明判别器太强,生成器学不到东西(“梯度消失”)。如果G_loss迅速降到0,说明可能发生了“模式崩溃”(生成器只学会生成少数几种图片来糊弄判别器)。
-
保存样本 :定期保存生成图片至关重要。通过观察这些图片的演变,你能最直观地判断模型是否在正常学习。如果图片一直是一片模糊的噪声,或者很快收敛到几个固定的奇怪图案,那就需要调整了。
6. 监控、调试与结果分析:如何判断你的GAN在“健康”学习?
训练启动后,你不能只是干等着。你需要像医生一样,持续监控模型的“生命体征”——损失曲线和生成样本。
6.1 可视化损失曲线
在训练循环中,我们已经记录了每一批次的 D_losses 和 G_losses 。训练结束后(或中途),我们可以绘制损失曲线。
import matplotlib.pyplot as plt
plt.figure(figsize=(10,5))
plt.title("Generator and Discriminator Loss During Training")
plt.plot(G_losses, label="G", alpha=0.7)
plt.plot(D_losses, label="D", alpha=0.7)
plt.xlabel("Iterations")
plt.ylabel("Loss")
plt.legend()
plt.grid(True, alpha=0.3)
plt.show()
如何解读损失曲线?
- 理想情况 :两条曲线都有波动,但整体趋势是逐渐下降并最终在一个非零值附近震荡。这表示对抗达到了一个纳什均衡。
- 判别器过强(D Loss → 0) :如果D_loss迅速降到接近0且不再回升,而G_loss很高,说明判别器轻易识破了所有假货,生成器获得的梯度非常小(消失),无法继续学习。 对策 :可以尝试降低判别器的学习率、减少判别器的更新频率(比如D更新5次,G更新1次)、或者使用标签平滑。
- 生成器过强/模式崩溃(G Loss → 0) :如果G_loss迅速降到0,而D_loss很高,这可能意味着生成器找到了判别器的某个致命弱点,只生成极少数能骗过判别器的图片,失去了多样性。生成的样本看起来几乎一模一样。 对策 :可以尝试增加判别器的能力(如加深网络)、在判别器中使用Dropout、或者使用更先进的GAN变体(如WGAN-GP)。
- 损失剧烈震荡 :曲线上下跳动非常厉害。这通常是学习率设置过高导致的。 对策 :降低学习率(
lr),比如从0.0002降到0.0001或0.00005。
6.2 观察生成样本的演变
打开你保存的 images 文件夹,按顺序查看生成的图片。一个成功的训练过程应该是这样的:
- 初期(前几个epoch) :图片完全是随机噪声,看不出任何结构。
- 中期(10-20个epoch) :开始出现模糊的数字轮廓,可能像墨团,但已经能隐约看出是数字。
- 后期(30-50个epoch) :数字变得清晰,笔画分明,并且具有多样性(生成不同的数字,且同一数字也有不同写法)。
如果发现:
- 图片始终是噪声 :可能是模型架构有问题、学习率太低、或者梯度根本没有回传(检查
detach()使用是否正确)。 - 图片模糊不清 :这是GAN的常见问题。可以尝试使用更深的网络、不同的上采样方式(如PixelShuffle)、或在损失函数中加入感知损失等。
- 所有图片都长得一样(模式崩溃) :如前所述,需要增强判别器或改用改进的GAN结构。
6.3 体验“潜在空间”的插值
一个训练好的生成器,其潜在空间(latent space)应该是连续且有意义的。我们可以通过插值来验证这一点。
# 加载训练好的生成器模型(假设已保存为 generator.pth)
# generator.load_state_dict(torch.load('generator.pth'))
# generator.eval()
with torch.no_grad(): # 不计算梯度,节省内存
# 随机选取两个噪声向量
z1 = torch.randn(1, latent_dim)
z2 = torch.randn(1, latent_dim)
# 在两个向量之间进行线性插值,生成10个点
n_interpolation = 10
interpolated_imgs = []
for alpha in torch.linspace(0, 1, n_interpolation):
z = alpha * z1 + (1 - alpha) * z2 # 线性插值
img = generator(z)
interpolated_imgs.append(img)
# 可视化插值结果
fig, axes = plt.subplots(1, n_interpolation, figsize=(15, 2))
for i, ax in enumerate(axes):
img = interpolated_imgs[i].squeeze().numpy() * 0.5 + 0.5
ax.imshow(img, cmap='gray')
ax.axis('off')
ax.set_title(f'{i}')
plt.show()
如果潜在空间学习得好,你会看到生成的数字从 z1 对应的数字,平滑地、连续地过渡到 z2 对应的数字,而不是突兀地跳跃。这证明了生成器不是简单记忆,而是学习到了数字的抽象特征和流形结构。
7. 从MNIST走向更复杂的世界:进阶思路与常见问题深挖
在MNIST上跑通DCGAN只是一个开始。如果你想生成更复杂的图像(如人脸、风景、动漫头像),会遇到更多挑战。这里分享一些进阶思路和我踩过的坑。
7.1 面对更复杂数据集的调整策略
-
数据预处理 :
- 尺寸 :MNIST是28x28,对于人脸(如CelebA 178x218)或自然场景,需要统一缩放到更大的尺寸,如64x64, 128x128。网络结构中的初始特征图大小和上采样次数需要相应调整。
- 通道 :彩色图像有3个通道(RGB),记得将模型中的
img_channels从1改为3。 - 增强 :对于数据量有限的情况,可以使用随机水平翻转、小角度旋转等数据增强,但需谨慎,因为GAN本身对数据分布很敏感。
-
模型容量 :
- 更复杂的数据需要更深的网络和更多的特征通道。可以增加生成器和判别器中卷积层的通道数(如从128/64增加到512/256)。
- 考虑使用残差块(ResBlock)。对于生成高分辨率图像(如256x256以上),像ProGAN或StyleGAN那样,采用渐进式增长或风格迁移结构是更主流的选择。
-
训练技巧与稳定化 :
- Wasserstein GAN with Gradient Penalty (WGAN-GP) :这是目前最流行、最稳定的GAN变体之一。它用Wasserstein距离代替JS散度作为损失,并添加梯度惩罚项,从根本上缓解了模式崩溃和训练不稳定的问题。 如果你在复杂数据集上训练原始GAN失败,强烈建议转向WGAN-GP。 它的损失函数不再是对抗性的,判别器(在WGAN中称为Critic)的输出是一个分数,而不是概率。
- 谱归一化(Spectral Normalization) :通过对判别器每一层的权重矩阵进行谱范数归一化,来限制判别器的Lipschitz常数,也能有效稳定训练,常与WGAN-GP结合使用。
- 学习率调度 :使用学习率衰减,在训练后期降低学习率,有助于模型收敛到更优的点。
7.2 实战中高频问题排查清单
- 问题:生成图片全是黑色/白色/单一颜色块。
- 检查 :输出层激活函数是否为
Tanh,数据标准化范围是否为[-1,1]。如果用了Sigmoid输出[0,1]而数据是[-1,1],会导致输出饱和。 - 检查 :损失函数是否爆炸(变成NaN)。可能是学习率太高。
- 检查 :输出层激活函数是否为
- 问题:损失值长时间不下降。
- 检查 :模型参数是否在更新?可以在训练循环中打印某一层的权重,看其值是否变化。
- 检查 :梯度是否存在?在
backward()之后,检查generator.parameters()的.grad属性是否非空。 - 尝试 :使用更简单的架构或数据集(如MNIST)验证代码流程是否正确。
- 问题:训练速度慢。
- 确保 :使用了GPU(
torch.cuda.is_available()为True)。 - 检查 :
DataLoader的num_workers是否大于0(如4或8),以并行加载数据。 - 考虑 :使用混合精度训练(
torch.cuda.amp),可以显著减少显存占用并加速训练。
- 确保 :使用了GPU(
- 问题:生成图片有规律的棋盘伪影。
- 原因 :这通常是由转置卷积(
ConvTranspose2d)的步长和核大小不匹配造成的。 - 解决 :用
Upsample(最近邻或双线性)+ 普通Conv2d的组合替代转置卷积,或者确保转置卷积的核大小能被步长整除。
- 原因 :这通常是由转置卷积(
从零实现并训练一个GAN,就像完成一次精密的手工实验。你会对数据流、梯度传播、模型博弈有刻骨铭心的理解。这份理解,是调用任何高级API都无法获得的。当你看到自己编写的代码从一片混沌的噪声中,逐渐描绘出清晰的图像时,那种成就感是无与伦比的。希望这篇超详细的指南,能成为你探索生成式AI世界的一块坚实垫脚石。记住,遇到问题是常态,耐心观察损失曲线和生成样本,系统地排查,你总能找到让模型“炼”出好图的方法。
更多推荐



所有评论(0)