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]

关键点解析:

  1. latent_dim (潜在空间维度) :这是输入噪声的维度,通常设为100。它代表了生成图像的“创意空间”大小。维度太低,模型表达能力不足;太高,可能增加训练难度和不确定性。
  2. 上采样 vs 转置卷积 :早期常用转置卷积( nn.ConvTranspose2d ),但它容易产生棋盘伪影。这里我使用了 Upsample (最近邻或双线性插值)后接普通卷积的组合,这是现在更推荐的做法,能生成质量更高的图像。
  3. 批归一化(BatchNorm) :在生成器的每一层(除了输出层)之后使用,它有助于缓解训练初期因输入数据分布差异大导致的梯度问题,是稳定GAN训练的关键技术之一。
  4. 激活函数 :隐藏层使用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

关键点解析:

  1. LeakyReLU :判别器中使用LeakyReLU而不是普通的ReLU,主要是为了防止梯度稀疏。当输入为负数时,LeakyReLU有一个很小的斜率(如0.2),允许梯度反向传播,避免了“神经元死亡”问题,这在判别器中尤为重要。
  2. 下采样 :通过卷积的步长(stride=2)实现,逐步降低空间维度,增加通道数,以提取更高层次的特征。
  3. 输出层 :一个全连接层接Sigmoid,输出一个介于0到1之间的概率值。1代表判别器认为图片是真实的,0代表认为是生成的。
  4. 判别器中没有在输入层使用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("训练完成!")

训练循环的逐行解读与避坑指南:

  1. optimizer.zero_grad() 这是必须的! 在每次计算新梯度前,必须将上一轮迭代中累积的梯度清零。否则梯度会不断累加,导致训练失控。

  2. 判别器训练中的 detach() gen_imgs.detach() 关键操作 。在训练判别器时,我们只希望更新判别器自身的参数。 detach() 方法将 gen_imgs 从当前的计算图中分离出来,创建一个新的、不需要梯度的张量。这样,当计算 d_fake_loss 并反向传播时,梯度只会更新判别器D,而不会影响到生成器G。如果不加 detach() ,梯度会沿着 gen_imgs 回溯到生成器G,这违背了“固定G,训练D”的步骤。

  3. 生成器训练中不使用 detach() :在训练生成器时,我们的目标是让判别器D对生成图片 gen_imgs 的输出概率 g_pred 接近1。因此,我们需要梯度从判别器的损失 g_loss 一路回溯,经过判别器D,最终到达生成器G,从而更新G的参数。所以这里绝对不能 detach

  4. 损失计算顺序 :先训练判别器一步,再训练生成器一步,这是最经典的交替训练方式。也有研究尝试一次训练判别器多次再训练一次生成器(例如D训练5次,G训练1次),以保持判别器的能力略强于生成器,这在某些复杂数据集上可能更稳定。对于MNIST,1:1的交替通常就够了。

  5. 损失值解读

    • 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. 保存样本 :定期保存生成图片至关重要。通过观察这些图片的演变,你能最直观地判断模型是否在正常学习。如果图片一直是一片模糊的噪声,或者很快收敛到几个固定的奇怪图案,那就需要调整了。

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 文件夹,按顺序查看生成的图片。一个成功的训练过程应该是这样的:

  1. 初期(前几个epoch) :图片完全是随机噪声,看不出任何结构。
  2. 中期(10-20个epoch) :开始出现模糊的数字轮廓,可能像墨团,但已经能隐约看出是数字。
  3. 后期(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 面对更复杂数据集的调整策略

  1. 数据预处理

    • 尺寸 :MNIST是28x28,对于人脸(如CelebA 178x218)或自然场景,需要统一缩放到更大的尺寸,如64x64, 128x128。网络结构中的初始特征图大小和上采样次数需要相应调整。
    • 通道 :彩色图像有3个通道(RGB),记得将模型中的 img_channels 从1改为3。
    • 增强 :对于数据量有限的情况,可以使用随机水平翻转、小角度旋转等数据增强,但需谨慎,因为GAN本身对数据分布很敏感。
  2. 模型容量

    • 更复杂的数据需要更深的网络和更多的特征通道。可以增加生成器和判别器中卷积层的通道数(如从128/64增加到512/256)。
    • 考虑使用残差块(ResBlock)。对于生成高分辨率图像(如256x256以上),像ProGAN或StyleGAN那样,采用渐进式增长或风格迁移结构是更主流的选择。
  3. 训练技巧与稳定化

    • 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 ),可以显著减少显存占用并加速训练。
  • 问题:生成图片有规律的棋盘伪影。
    • 原因 :这通常是由转置卷积( ConvTranspose2d )的步长和核大小不匹配造成的。
    • 解决 :用 Upsample (最近邻或双线性)+ 普通 Conv2d 的组合替代转置卷积,或者确保转置卷积的核大小能被步长整除。

从零实现并训练一个GAN,就像完成一次精密的手工实验。你会对数据流、梯度传播、模型博弈有刻骨铭心的理解。这份理解,是调用任何高级API都无法获得的。当你看到自己编写的代码从一片混沌的噪声中,逐渐描绘出清晰的图像时,那种成就感是无与伦比的。希望这篇超详细的指南,能成为你探索生成式AI世界的一块坚实垫脚石。记住,遇到问题是常态,耐心观察损失曲线和生成样本,系统地排查,你总能找到让模型“炼”出好图的方法。

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐