ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

GAN原理与PyTorch实现:从训练逻辑到手写数字生成实战

GAN原理与PyTorch实现:从训练逻辑到手写数字生成实战 在深度学习领域里生成对抗网络Generative Adversarial Network简称 GAN可能是最让初学者既兴奋又困惑的方向之一。兴奋在于它能生成图片、音频、文本甚至完整人脸效果非常震撼困惑在于它的训练过程远比普通分类网络复杂两个网络相互博弈loss 总是不收敛生成图像也经常模糊或者模式单一。这篇文章会从 GAN 的基础概念讲起重点拆解其训练逻辑再用 PyTorch 手写一个可运行的简单 GAN 项目最后给出训练 GAN 时最常遇到的坑和排查思路。无论你是刚开始接触深度学习还是已经会训练 CNN、RNN都可以通过这篇文章把 GAN 的原理和代码真正串起来。1. GAN 是什么从生成任务说起1.1 生成模型要解决什么问题在传统深度学习任务里我们更多接触的是判别模型比如图像分类、目标检测、语义分割。这类模型的共同特点是给定输入 x预测标签 y。模型学到的是一张从输入到输出的映射关系本质上是在做“判断”。生成模型的目标完全不同。它希望学习训练数据的真实分布 p_data(x)然后从中学到的分布中采样出新的样本。换句话说判别模型在做“这张图是不是猫”生成模型在做“画一只不存在的猫”。常见的生成模型包括自回归模型、变分自编码器VAE、扩散模型Diffusion Model和生成对抗网络GAN。在 GAN 出现之前生成模型的难点在于如何度量生成样本和真实样本之间的差距。如果我们用像素级别的损失函数比如 MSE生成结果往往会模糊因为模型会把多种可能性平均起来。比如一张模糊的猫脸可能只是把各种猫的特征做了平均而不是真实清晰的猫脸。GAN 提出了一个非常巧妙的思路不直接定义生成样本的评价函数而是让另一个神经网络学会判断“这个样本是真实数据还是生成数据”。通过这种判别器的反馈生成器可以逐步把输出调整得越来越像真实数据。这个思路在 2014 年由 Ian Goodfellow 等人提出之后迅速成为生成领域的重要研究方向。1.2 GAN 的核心思想对抗训练GAN 的核心是引入两个网络生成器 G输入一个低维的随机噪声向量 z输出一个伪造样本 G(z)。判别器 D输入一个样本 x输出一个概率 D(x)表示它有多大的把握认为 x 来自真实数据。训练目标非常简单但执行起来非常有技巧生成器希望让判别器误以为伪造样本是真样本也就是让 D(G(z)) 尽量接近 1。判别器希望尽量区分真实样本和伪造样本。看到真实样本时输出接近 1看到伪造样本时输出接近 0。两个网络的优化目标相互对抗最终达到纳什均衡生成器生成的样本已经足够逼真判别器无法区分真假只能随机猜测。这里有个容易混淆的地方判别器不是我们最终需要的模型它只是一个“教练”角色。最终我们是希望生成器学会从随机噪声中生成逼真的数据。1.3 为什么叫“生成对抗网络”“生成”体现在生成器负责创造新数据“对抗”体现在两个网络的训练目标互为相反网络结构上两者都是深度神经网络所以合称生成对抗网络。我们可以把 GAN 想象成一个造假团队和鉴定团队的博弈。造假团队的目标是制造出无法被鉴定出的假画鉴定团队的目标是把所有假画都找出来。两者互相促进最终造假团队的水平越来越高鉴定团队的鉴别能力也越来越强。GAN 的训练就是反复交替进行这两个过程。这引出一个初学者很容易忽略的问题GAN 的训练不是一次前向传播加一次反向传播就能完成的而是一个交替优化的动态过程。下面重点拆解它的训练逻辑。2. GAN 的训练逻辑拆解2.1 两个网络的职责边界先约定符号真实数据分布p_data(x)先验噪声分布p_z(z)通常使用标准正态分布 N(0, 1)生成器参数θ_G判别器参数θ_D判别器的输入有两类从真实数据集中采样得到 x_real。从噪声分布采样 z经过生成器得到 x_fake G(z)。判别器的输出是输入为真实样本的概率。在二分类交叉熵视角下判别器的目标是对于真实样本尽量让 D(x_real) 接近 1。对于伪造样本尽量让 D(G(z)) 接近 0。生成器则不是直接最小化“生成样本与真实样本的像素差值”而是最小化判别器对伪造样本判断错误的程度也就是让 D(G(z)) 尽量接近 1。这个目标会让生成器更关注“能否骗过判别器”而不是简单套用某个像素距离。2.2 目标函数最小最大化博弈GAN 的原始目标函数如下min_G max_D V(D, G) E_x~p_data(x)[log D(x)] E_z~p_z(z)[log(1 - D(G(z)))]从判别器角度看它要让整个式子最大化。真实样本部分 log D(x) 越大越好说明判别器正确识别真实样本伪造样本部分 log(1 - D(G(z))) 越大越好说明判别器正确识别伪造样本。从生成器角度看它要让整个式子最小化。由于它无法控制真实样本部分只能让 log(1 - D(G(z))) 变小也就是让 D(G(z)) 变大让判别器把伪造样本误判为真实样本。这个目标函数理论上是完美的但在实际训练中会遇到一个问题生成器训练初期梯度非常弱。因为当生成器还很差时判别器可以轻易区分真伪此时 1 - D(G(z)) 接近 0log(1 - D(G(z))) 接近负无穷梯度饱和。所以实际代码中我们通常不直接最小化 log(1 - D(G(z)))而是反过来让生成器最小化 -log(D(G(z)))。这样生成器的目标是让 D(G(z)) 尽量接近 1梯度信号更强训练也更稳定。这个技巧在几乎所有 GAN 实现中都会出现。2.3 训练步骤的伪代码流程GAN 的标准训练流程可以拆成以下几个步骤从真实数据集中采样一个 batch 的真实样本 x。从先验噪声分布中采样一个 batch 的噪声 z。用生成器计算伪造样本 x_fake G(z)。用真实样本和伪造样本更新判别器参数判别器要尽可能识别出真实样本为真判别器要尽可能识别出伪造样本为假。再采样一个新的噪声 z_new注意这一步很关键。用伪造样本更新生成器参数目标是让判别器把伪造样本识别为真。这里的步骤 5 非常重要。更新生成器时不能沿用上一步判别器已经看过的同一个 z而是重新采样。这样可以让生成器在当前判别器下寻找新的欺骗方式避免判别器和生成器在同一个样本上来回纠缠。从代码实现角度来看更新判别器时会把真实样本和伪造样本拼在一起计算二分类交叉熵损失更新生成器时则把伪造样本标记为真样本让生成器尽量输出去骗过当前判别器。2.4 交替训练的含义两个网络不能同时更新参数原因是它们的目标函数相互依赖。如果同时更新一方面判别器能力飙升生成器完全跟不上另一方面生成器不稳定会给判别器提供混乱的梯度信号。标准做法是每轮迭代中先固定生成器更新判别器再固定判别器更新生成器交替进行。多数 GAN 实现中每一轮迭代会先更新判别器一次或多次再更新生成器一次。如果判别器太强可以调低判别器更新频率给生成器更多追赶机会。反过来如果生成器太强判别器完全无法区分真假训练也会失效。从数学角度看交替训练相当于在 G 和 D 的参数空间中用坐标下降法逼近极小极大问题的最优解。虽然实际训练中很难完美收敛到均衡点但交替训练通常是实践中效果最稳定的方式。3. 环境准备与依赖3.1 运行环境说明本文以 Python 和 PyTorch 为例演示一个最简单的 GAN 训练流程。示例环境如下操作系统Windows 10 / Ubuntu 20.04 均可Python 版本建议 3.8 及以上PyTorch 版本结合你本地的 CUDA 环境安装CPU 版本也能运行示例代码torchvision用于加载 MNIST 数据集matplotlib用于可视化生成结果如果你的机器没有 GPU本文给出的模型结构仍然可以在 CPU 上训练只是速度会慢一些。MNIST 图像只有 28x28 灰度图结构比较轻量CPU 训练几百个 epoch 也完全可以接受。3.2 依赖库安装建议使用 conda 或 venv 创建独立环境避免与现有项目冲突。安装命令如下pip install torch torchvision matplotlib如果没有 GPU可以在 PyTorch 官网选择 CPU 版本的安装命令。版本不需要完全一致只要保证 torch 和 torchvision 版本兼容即可。另外建议安装 tqdm用来显示训练进度pip install tqdm3.3 项目结构为了便于阅读我们采用一个简单的单文件结构gan_mnist/ ├── gan.py └── output/ ├── generated_images.png └── generator_final.pthgan.py 存放所有代码模型定义、训练循环、可视化。output 目录存放训练过程中生成的图片和最终模型权重。4. 从零实现一个简单 GANPyTorch 实战4.1 数据集准备MNISTMNIST 是手写数字数据集一共 10 类数字图像尺寸 28x28灰度图。我们用 torchvision 直接下载并加载不需要手动处理图片。在加载时需要注意三点图像范围是 0 到 255需要归一化到 -1 到 1 之间。这里选择 -1 到 1 而不是 0 到 1是因为生成器最后一层通常使用 tanh 激活函数输出范围正好是 [-1, 1]两者对齐后损失函数的语义更准确。数据管道使用 DataLoader方便按 batch 取数据。要把标签丢掉GAN 这里不需要使用数字标签。数据加载代码如下import torch import torchvision.datasets as datasets import torchvision.transforms as transforms from torch.utils.data import DataLoader transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) batch_size 128 train_dataset datasets.MNIST( root./data, trainTrue, transformtransform, downloadTrue ) train_loader DataLoader( datasettrain_dataset, batch_sizebatch_size, shuffleTrue )transforms.ToTensor()会把图像转为张量并把数值缩放到 [0, 1]Normalize((0.5,), (0.5,))会把数值线性变换到 [-1, 1]。很多初学者忽略这个细节后面对比生成图像时容易觉得生成结果一直发灰。4.2 定义生成器生成器的作用是接收一个随机噪声向量输出一张 28x28 的灰度图像。我们使用全连接网络加转置卷积的混合结构。这里为了代码简洁先用全连接网络实现便于初学者理解维度变化。输入维度是 latent_dim通常取 100输出维度是 784对应 28x28。import torch.nn as nn latent_dim 100 class Generator(nn.Module): def __init__(self, latent_dim): super(Generator, self).__init__() self.model nn.Sequential( nn.Linear(latent_dim, 256), nn.ReLU(True), nn.Linear(256, 512), nn.ReLU(True), nn.Linear(512, 1024), nn.ReLU(True), nn.Linear(1024, 784), nn.Tanh() ) def forward(self, z): img self.model(z) return img.view(-1, 1, 28, 28)这里的关键点是最后一层使用 Tanh 激活函数输出范围在 [-1, 1] 之间与 MNIST 预处理后的数据范围一致。中间层使用 ReLU 来提供非线性能力。forward方法最后把向量 reshape 成 1x28x28 的图像张量因为 MNIST 是单通道灰度图。如果后续要生成彩色头像或更复杂的图像需要把输出通道数和最后一层的输出维度做对应调整。但作为 GAN 基础这个结构已经足够。4.3 定义判别器判别器的作用是接收一张 28x28 图像输出一个标量表示图像来自真实数据的概率。输入维度是 784输出维度是 1。中间层可以与生成器对称也可以不对称。因为判别器是二分类任务最后一层使用 Sigmoid 激活函数把输出压缩到 [0, 1]。class Discriminator(nn.Module): def __init__(self): super(Discriminator, self).__init__() self.model nn.Sequential( nn.Linear(784, 1024), nn.LeakyReLU(0.2, inplaceTrue), nn.Dropout(0.3), nn.Linear(1024, 512), nn.LeakyReLU(0.2, inplaceTrue), nn.Dropout(0.3), nn.Linear(512, 256), nn.LeakyReLU(0.2, inplaceTrue), nn.Dropout(0.3), nn.Linear(256, 1), nn.Sigmoid() ) def forward(self, img): x img.view(-1, 784) return self.model(x)判别器使用 LeakyReLU 而不是 ReLU原因在于 ReLU 在负区间输出恒为 0可能导致神经元死亡而 LeakyReLU 在负区间保留一个小的斜率能让梯度持续流通。对于判别器这种需要区分真实和伪造样本的模型LeakyReLU 通常会有更好的梯度表现。Dropout 的作用是防止判别器过拟合。判别器的任务本质是一个二分类问题过拟合会让它在训练集上表现很好但在真实数据分布上泛化差进而给生成器传递错误的训练信号。4.4 定义训练循环这是整个 GAN 训练逻辑的核心部分。前面已经提到GAN 的训练需要交替更新判别器和生成器下面的代码展示了具体的实现细节。先定义优化器和损失函数import torch.optim as optim device torch.device(cuda if torch.cuda.is_available() else cpu) generator Generator(latent_dim).to(device) discriminator Discriminator().to(device) criterion nn.BCELoss() lr 0.0002 beta1 0.5 g_optimizer optim.Adam(generator.parameters(), lrlr, betas(beta1, 0.999)) d_optimizer optim.Adam(discriminator.parameters(), lrlr, betas(beta1, 0.999))这里使用 Adam 优化器学习率取 0.0002beta1 取 0.5这是很多 GAN 实现中比较稳定的默认配置。训练循环代码如下import torchvision.utils as vutils import matplotlib.pyplot as plt num_epochs 50 g_losses [] d_losses [] real_label 1.0 fake_label 0.0 for epoch in range(num_epochs): for i, (imgs, _) in enumerate(train_loader): current_batch_size imgs.size(0) imgs imgs.to(device) # 1. 训练判别器 discriminator.zero_grad() # 真实样本损失 real_output discriminator(imgs) d_real_loss criterion(real_output, torch.full((current_batch_size, 1), real_label, devicedevice)) # 生成伪造样本 z torch.randn(current_batch_size, latent_dim, devicedevice) fake_imgs generator(z) fake_output discriminator(fake_imgs.detach()) d_fake_loss criterion(fake_output, torch.full((current_batch_size, 1), fake_label, devicedevice)) d_loss d_real_loss d_fake_loss d_loss.backward() d_optimizer.step() # 2. 训练生成器 generator.zero_grad() z torch.randn(current_batch_size, latent_dim, devicedevice) fake_imgs generator(z) fake_output discriminator(fake_imgs) g_loss criterion(fake_output, torch.full((current_batch_size, 1), real_label, devicedevice)) g_loss.backward() g_optimizer.step() g_losses.append(g_loss.item()) d_losses.append(d_loss.item()) if i % 200 0: print(fEpoch [{epoch}/{num_epochs}] Batch [{i}/{len(train_loader)}] D Loss: {d_loss.item():.4f}, G Loss: {g_loss.item():.4f}) # 每个 epoch 结束后生成一组固定噪声的图片便于观察生成效果 if epoch % 5 0: with torch.no_grad(): fixed_z torch.randn(64, latent_dim, devicedevice) fixed_imgs generator(fixed_z).detach().cpu() grid vutils.make_grid(fixed_imgs, nrow8, normalizeTrue) plt.figure(figsize(8, 8)) plt.imshow(grid.permute(1, 2, 0)) plt.axis(off) plt.savefig(foutput/epoch_{epoch}.png) plt.close()在判别器训练步骤中我们调用fake_imgs.detach()目的是让伪造样本的梯度不会传到生成器因为这一步只更新判别器参数。如果不做 detach梯度会回溯到生成器导致本步骤同时修改两个模型参数训练会乱掉。在生成器训练步骤中我们重新采样了噪声 z并让 fake_imgs 直接输入判别器。此时不调用 detach因为需要梯度流回生成器让生成器根据判别器的反馈调整参数。从代码上可以看到生成器的目标不是最小化“图像与真实图片的像素差”而是让判别器输出接近 1也就是判别器把伪造图片当成真实图片。4.5 完整代码汇总下面是去掉注释后可以直接运行的完整代码。文件位置建议放在 gan_mnist/gan.py 中import torch import torch.nn as nn import torch.optim as optim import torchvision.datasets as datasets import torchvision.transforms as transforms import torchvision.utils as vutils from torch.utils.data import DataLoader import matplotlib.pyplot as plt latent_dim 100 batch_size 128 num_epochs 50 lr 0.0002 beta1 0.5 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset datasets.MNIST(root./data, trainTrue, transformtransform, downloadTrue) train_loader DataLoader(datasettrain_dataset, batch_sizebatch_size, shuffleTrue) class Generator(nn.Module): def __init__(self, latent_dim): super(Generator, self).__init__() self.model nn.Sequential( nn.Linear(latent_dim, 256), nn.ReLU(True), nn.Linear(256, 512), nn.ReLU(True), nn.Linear(512, 1024), nn.ReLU(True), nn.Linear(1024, 784), nn.Tanh() ) def forward(self, z): img self.model(z) return img.view(-1, 1, 28, 28) class Discriminator(nn.Module): def __init__(self): super(Discriminator, self).__init__() self.model nn.Sequential( nn.Linear(784, 1024), nn.LeakyReLU(0.2, inplaceTrue), nn.Dropout(0.3), nn.Linear(1024, 512), nn.LeakyReLU(0.2, inplaceTrue), nn.Dropout(0.3), nn.Linear(512, 256), nn.LeakyReLU(0.2, inplaceTrue), nn.Dropout(0.3), nn.Linear(256, 1), nn.Sigmoid() ) def forward(self, img): x img.view(-1, 784) return self.model(x) device torch.device(cuda if torch.cuda.is_available() else cpu) generator Generator(latent_dim).to(device) discriminator Discriminator().to(device) criterion nn.BCELoss() g_optimizer optim.Adam(generator.parameters(), lrlr, betas(beta1, 0.999)) d_optimizer optim.Adam(discriminator.parameters(), lrlr, betas(beta1, 0.999)) real_label 1.0 fake_label 0.0 for epoch in range(num_epochs): for i, (imgs, _) in enumerate(train_loader): current_batch_size imgs.size(0) imgs imgs.to(device) discriminator.zero_grad() real_output discriminator(imgs) d_real_loss criterion(real_output, torch.full((current_batch_size, 1), real_label, devicedevice)) z torch.randn(current_batch_size, latent_dim, devicedevice) fake_imgs generator(z) fake_output discriminator(fake_imgs.detach()) d_fake_loss criterion(fake_output, torch.full((current_batch_size, 1), fake_label, devicedevice)) d_loss d_real_loss d_fake_loss d_loss.backward() d_optimizer.step() generator.zero_grad() z torch.randn(current_batch_size, latent_dim, devicedevice) fake_imgs generator(z) fake_output discriminator(fake_imgs) g_loss criterion(fake_output, torch.full((current_batch_size, 1), real_label, devicedevice)) g_loss.backward() g_optimizer.step() if i % 200 0: print(fEpoch [{epoch}/{num_epochs}] Batch [{i}/{len(train_loader)}] D Loss: {d_loss.item():.4f}, G Loss: {g_loss.item():.4f}) if epoch % 5 0: with torch.no_grad(): fixed_z torch.randn(64, latent_dim, devicedevice) fixed_imgs generator(fixed_z).detach().cpu() grid vutils.make_grid(fixed_imgs, nrow8, normalizeTrue) plt.figure(figsize(8, 8)) plt.imshow(grid.permute(1, 2, 0)) plt.axis(off) plt.savefig(foutput/epoch_{epoch}.png) plt.close()4.6 运行与预期结果在项目目录下先创建 output 文件夹然后运行mkdir output python gan.py运行日志会持续打印每个 batch 的损失值。在 GAN 训练初期判别器损失会相对较大因为区分真实和伪造样本比较容易生成器损失也会经历波动。随着训练推进两者会进入一种动态平衡状态。每个 5 个 epoch程序会保存一张生成图片。你可以直观地看到图片从满屏噪声慢慢变成模糊的数字再到轮廓清晰的数字。需要说明的是这个示例模型结构比较简单生成效果可能不如论文中展示的那样完美。初学者不必追求极致的图片质量重点是通过代码理解 GAN 的训练流程。5. GAN 训练中的常见问题与排查思路GAN 训练不稳定是出了名的。很多初学者发现模型怎么调都出不来效果这里整理最常见的几类问题以及排查思路。问题现象常见原因解决思路生成器损失快速降为 0判别器损失也趋近于 0判别器能力太强生成器完全没有机会更新降低判别器学习率或者减少判别器更新频率判别器损失一直很大生成器一直生不出像样的图片生成器太弱或者模型结构不匹配增加生成器容量检查输入噪声维度是否合理生成图片模糊边缘不清晰使用 MSE 或 L1 损失或网络容量不足改用 BCE 对抗损失适当增加模型宽度生成图片非常单一只输出几种固定样式模式坍塌Mode Collapse生成器只学会了部分真实分布降低学习率增加训练样本多样性尝试 mini-batch discrimination 或引入多样性的正则约束训练过程不断震荡损失无法稳定下降生成器和判别器更新步调不协调调低学习率增加 batch size或者使用标签平滑判别器 loss 很低但生成图片仍然很差判别器过拟合没有真实区分能力增加 Dropout增加训练数据限制判别器更新次数生成图像出现大量棋盘格噪声转置卷积操作产生的重叠效应改用反卷积时调整 kernel_size 和 stride或使用上采样加普通卷积替代转置卷积从这些常见问题可以看出GAN 训练的本质是在平衡两个网络的学习速度。如果判别器学得太快生成器就无从更新如果生成器学得太快又可能让判别器快速失效。无数技巧本质上都是为了控制这个平衡。排查 GAN 不收敛时我建议按照下面的顺序执行先确认数据预处理是否正确特别是图像归一化范围是否与生成器输出激活函数一致。确认判别器是否真的能够区分真实数据和生成数据单独训练判别器试一下准确率。确认生成器的梯度是否正常可以在生成器训练步骤中打印梯度的均值或范数。检查损失值是否处于合理量级。BCE 损失在 0.1 到 0.7 之间波动是常见的如果出现极大或极小值要考虑训练是否崩溃。最后才去调整网络结构和超参数不要一开始就盲目换模型。6. GAN 训练的最佳实践与工程建议6.1 训练稳定性技巧在项目实践中下面几条经验对提高 GAN 训练稳定性非常有帮助。第一使用标签平滑。真实标签不要设为 1可以设为 0.9 或 0.8这样可以防止判别器过度自信从而减少生成器收到过强梯度信号的风险。第二优先使用 Adam 优化器并保持学习率适中。过高的学习率是 GAN 训练崩溃的主要原因之一。第三避免判别器和生成器模型差距过大。如果判别器参数量远超生成器两者能力不匹配训练就会长期停滞。第四可以考虑使用特征匹配损失。生成器不仅要骗过判别器的最终输出还可以尝试匹配判别器中间层对真实数据和伪造数据的不同特征。这种方式在某些任务中能显著提高稳定性。第五使用固定的噪声输入来周期性检查生成效果。在训练过程中保留一组固定 z每个 epoch 都生成同一组图片这样你可以直观地看到模型是否在逐步学习。如果固定 z 的输出没有变化说明生成器可能已经停滞。第六保存模型时不要只保存最终权重建议每个 epoch 或者每隔几个 epoch 保存一次检查点。GAN 训练过程中最优效果往往不是最后一个 epoch早期某个状态可能生成效果更好。保存检查点可以方便回溯。6.2 调参与工程化建议GAN 的工程化与普通深度学习任务不太一样。普通任务常常关注在验证集上的准确率而 GAN 没有稳定的验证指标。因此工程上通常采用人工检查生成图像质量、损失曲线稳定性、多样性评估等方法来综合判断。在调参时建议每次只改一个变量不要同时调整学习率、网络结构、损失函数等多个因素否则很难定位问题来源。另外写代码时应该把随机种子固定下来import random import numpy as np import torch def set_seed(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) set_seed(42)固定随机种子可以在调试时保证实验结果可复现对于定位问题和对比实验尤其重要。在生产环境或研究中应用 GAN 时还要注意隐私和安全问题。GAN 可以生成高仿真的人脸、声音等内容滥用可能带来隐私侵犯、虚假信息传播等风险。无论是做学术研究还是工程开发都应该遵守数据使用规范和相关法律法规在合法授权的前提下使用数据不要将生成模型用于制作虚假身份、伪造证据或其他恶意场景。6.3 性能优化建议GAN 训练有两个明显瓶颈计算资源消耗大、训练时间长。为了提升效率可以从以下几点入手优先使用 GPU 训练如果 GPU 显存有限可以适当降低 batch size。使用混合精度训练。PyTorch 的torch.cuda.amp可以显著减少显存占用并加速训练。提前终止不必要的 epoch。如果固定噪声生成的图片质量已经很久没有提升没有必要继续训练。分批保存中间结果。不要等训练完再统一生成图片否则如果训练中途崩溃前面全部作废。7. 总结与下一步学习方向通过这篇文章你应该已经掌握了 GAN 的核心概念、训练逻辑、目标函数背后的原理以及如何用 PyTorch 实现一个简单的 GAN 并完成 MNIST 手写数字生成。最关键的是理解了交替训练、判别器与生成器的对抗关系、以及两个网络如何通过对抗逐步提升能力。下一步可以从以下几个方向继续深入学习使用卷积网络替代全连接网络构建 DCGAN生成更清晰的图像。引入条件信息实现条件 GANcGAN让生成器按照指定类别生成图片。学习 WGAN、WGAN-GP解决原始 GAN 训练不稳定的问题。尝试 StyleGAN、CycleGAN 等更先进的生成模型了解不同任务的建模方式。把 GAN 与自注意力机制结合处理大尺寸图像生成任务。在实际项目中最需要优先关注的风险是模式坍塌和训练不收敛。不要指望一套参数能通吃所有数据集不同图像质量、不同数据分布下GAN 的超参数都需要反复尝试。建议你先从本文的基础代码跑通再用固定噪声观察输出变化最后逐步替换网络结构。这样能在掌握 GAN 训练逻辑的同时积累出属于自己的调参经验。
返回列表