生成对抗网络GAN基础:原理与PyTorch实战实现
发布时间:2026/9/29 2:25:17来源:尧图网络
近几年生成模型发展很快你可能常听到“AI 绘画”“人脸生成”“图像修复”这些词。背后有一个绕不开的基础模型——GAN中文通常叫生成对抗网络Generative Adversarial Network。如果你刚进入深度学习的学习路线看到“9 GAN 9.1 GAN 基础 9.1.0 什么是 GAN”这一节第一反应可能是一堆疑问它和普通的卷积神经网络有什么区别为什么叫“对抗”两个网络怎么训练这篇文章就把 GAN 基础部分展开讲清楚。我们会从“什么是生成模型”开始拆解生成器和判别器的关系推导最核心的目标函数然后基于 PyTorch 从零实现一个能在 MNIST 手写数字数据集上训练的 GAN。代码可以直接复制运行你也可以把结果打印出来观察每一轮生成效果的变化。无论你是初学者还是正在补深度学习基础的开发者这篇文章都适合。1. 背景与核心概念1.1 什么是生成模型深度学习模型大体可以分成两类判别模型和生成模型。判别模型解决的是分类、回归问题。给一张图判断里面是猫还是狗给一段文本判断情感是正向还是负向。这类模型学习的是条件概率分布 (p(y|x))也就是“输入 (x) 时输出 (y) 的概率”。生成模型解决的是“生成数据”的问题。给模型一批真实图片让它学会这些图片的分布规律然后从分布中采样得到新的、之前没见过的样本。比如让它学习一万张猫的图片之后它能自己画出新的猫。这类模型学习的是数据本身的分布 (p(x))。听起来很简单难点在于真实图片的分布极其复杂。一张 (28\times28) 的 MNIST 灰度图每个像素取值 0 到 255如果直接建模就是一个几万维空间中的分布人工写出概率密度函数根本不可能。传统方法里我们可以用最大似然估计去拟合一个假设的分布比如高斯混合模型但表达能力有限。直到 GAN 出现才真正提供了一种“用神经网络去逼近复杂分布”的可行思路。1.2 GAN 的核心思想让两个网络互相对抗GAN 是 Ian Goodfellow 等人在 2014 年提出的模型。它的核心思路不是直接计算分布而是构造两个角色生成器Generator简称 G负责从随机噪声中生成伪造样本目标是骗过判别器。判别器Discriminator简称 D负责区分输入样本是来自真实数据还是来自生成器伪造的数据。你可以把生成器理解为一个“仿造者”把判别器理解为一个“鉴定师”。仿造者不断改进造假技术鉴定师不断升级识别能力。两者互相竞争、互相促进。最终状态是仿造者造出来的东西已经足以以假乱真鉴定师无法区分真假。这个过程在数学上就是一个极小极大博弈minimax game。GAN 的训练目标可以写成min_G max_D V(D, G)其中目标函数 (V(D,G)) 的含义是最大化判别器的判别能力同时最小化生成器的生成误差。1.3 GAN 的典型应用场景GAN 虽然最初是为了生成图像而生但已经扩展到非常多领域。常见应用包括图像生成生成人脸、风景、动漫角色等。图像修复对破损图片进行补全、去噪、去雨。图像超分辨率把低分辨率图片放大成高分辨率并补充细节。风格迁移把照片转换成油画风格、把白天的图片转换成夜晚风格。数据增强为训练集生成额外样本缓解样本不足的问题。异常检测利用生成器的重建误差定位图像中的异常区域。这些应用大多是在 GAN 基础理论上发展出来的。所以先把 9.1 的“什么是 GAN”学扎实后面的进阶版本才会更容易理解。2. GAN 核心原理拆解2.1 生成器 G从噪声到样本生成器的输入是一个随机噪声向量 (z)常见维度是 100 或 128服从标准正态分布。经过若干层神经网络的映射输出一个和真实样本同尺寸的向量。对于 MNIST输出就是 784 维的向量再 reshape 成 (1 \times 28 \times 28) 的图像。生成器的本质是[ G(z; \theta_g) ]其中 (\theta_g) 是生成器的参数。训练初期生成器输出的是纯噪声随着训练进行参数不断调整输出越来越接近真实图片。它没有见过任何真实图片的“像素级标签”只是通过判别器的反馈来改进。这里有一个初学者容易困惑的点生成器并没有直接计算“这张图像不像 MNIST”它只是根据判别器给它的梯度信号来调整自己的参数。这个信号从哪来来自判别器对它的评价。2.2 判别器 D真假鉴定器判别器是一个二分类网络。输入一张图片 (x)输出一个标量 (D(x))表示这张图片是真实图片的概率。如果输入来自真实数据集(D(x)) 的期望目标是 1。如果输入来自生成器 (G(z))(D(G(z))) 的期望目标是 0。判别器的本质是[ D(x; \theta_d) ]它的结构和普通分类网络类似可以是一个卷积网络也可以是一个全连接网络。在简单实验中全连接网络往往就够用。训练过程中生成器和判别器是交替更新的。不能同时更新也不能只更新其中一个否则博弈会失衡。2.3 目标函数与训练流程原始 GAN 的目标函数是[ \min_G \max_D V(D,G) \mathbb{E}{x \sim p{data}(x)}[\log D(x)] \mathbb{E}_{z \sim p_z(z)}[\log(1 - D(G(z)))] ]拆开看第一个期望项真实样本 (x) 输入判别器我们希望 (\log D(x)) 尽量大也就是判别器对真实样本输出接近 1。第二个期望项固定真实样本对随机噪声 (z)我们希望 (\log(1 - D(G(z)))) 尽量大也就是判别器对伪造样本输出接近 0。这是判别器的视角。但从生成器的角度看它希望 (D(G(z))) 输出接近 1也就是让判别器认不出伪造样本所以它希望第二项中 (\log(1 - D(G(z)))) 尽量小。那么训练流程就是从真实数据集中采样一批真实样本。从随机噪声分布中采样一批噪声用生成器生成一批伪造样本。把真实样本和伪造样本一起输入判别器计算判别损失更新判别器参数。固定当前判别器参数把新的噪声输入生成器计算生成器损失更新生成器参数。重复上述过程直到两者达到纳什均衡。在实现时初学者最容易犯的错误是把训练判别器和训练生成器合并成一次反向传播。正确的做法是交替更新并且训练生成器时要避免让判别器这一路的梯度干扰生成器参数更新。另外还有一个常见的“负梯度消失”问题。原始损失函数 (\log(1 - D(G(z)))) 在判别器过于强大时梯度会非常小导致生成器几乎学不动。所以在实际实现中通常会把生成器的目标改为最大化 (\log(D(G(z))))也就是让生成器尝试“让判别器认为伪造样本是真的”。这个改进在代码里很常见但理解了原始公式后你会知道这本质上只是在优化目标上做了一个等价变换目的是提供更好的梯度信号。3. 环境准备与实验约定3.1 软件环境本文代码基于 PyTorch 编写。版本需要根据你的项目实际情况调整本文示例以常见环境为例重点演示配置思路。操作系统Windows / Linux / macOS 均可。Python建议 3.8 及以上版本。PyTorch建议 1.13 或 2.x 系列CPU 版本即可运行有 GPU 会更快。torchvision与 PyTorch 版本匹配用于加载 MNIST 数据集。matplotlib用于可视化生成效果。对于初学者尤其注意不要只安装 torch 而忘了 torchvision。MNIST 数据集下载和预处理依赖 torchvision 的接口。3.2 项目结构为了便于管理建议按下面的结构组织代码gan-tutorial/ ├── train.py # 主训练脚本 ├── generator.py # 生成器定义 ├── discriminator.py # 判别器定义 └── result/ # 保存生成的图片当然为了减少文件依赖本文会给出一个可以直接保存为train.py的完整脚本。你也可以在 Jupyter Notebook 中按小节逐步运行。4. 从零实现一个 GAN下面我们基于 PyTorch 实现一个最简单的 GAN用于生成 MNIST 手写数字。这个版本重点突出基础流程网络结构保持简单方便你观察每一步在做什么。4.1 数据加载与预处理MNIST 是 28×28 的灰度图数据集。常见做法是把像素值从 [0,1] 归一化到 [-1,1]。这样做的原因是生成器最后使用 Tanh 激活函数输出范围也是 [-1,1]两者匹配训练更加稳定。import torch import torch.nn as nn import torch.optim as optim import torchvision import torchvision.transforms as transforms from torch.utils.data import DataLoader import os # 超参数 latent_dim 100 # 噪声向量维度 batch_size 128 epochs 50 lr 0.0002 # 学习率 device torch.device(cuda if torch.cuda.is_available() else cpu) # 数据预处理转换为 Tensor并归一化到 [-1, 1] transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset torchvision.datasets.MNIST( root./data, trainTrue, transformtransform, downloadTrue ) train_loader DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue)代码说明transforms.Normalize((0.5,), (0.5,))中的均值 0.5、标准差 0.5对应 MNIST 灰度图归一化到 [-1,1] 的常见参数。downloadTrue表示本地没有数据时会自动下载。如果你所在环境无法访问外网可以提前把 MNIST 数据文件放到./data目录下。4.2 定义生成器生成器接收 100 维随机噪声经过三层全连接层最终输出 784 维向量再 reshape 成 1×28×28 图像。class Generator(nn.Module): def __init__(self, latent_dim100): super(Generator, self).__init__() self.model nn.Sequential( nn.Linear(latent_dim, 128), nn.ReLU(inplaceTrue), nn.Linear(128, 256), nn.ReLU(inplaceTrue), nn.Linear(256, 512), nn.ReLU(inplaceTrue), nn.Linear(512, 784), nn.Tanh() ) def forward(self, z): # z 形状: (batch_size, latent_dim) img self.model(z) # 输出形状: (batch_size, 1, 28, 28) img img.view(img.size(0), 1, 28, 28) return img几个要点隐藏层激活函数用 ReLU。ReLU 计算简单且梯度不容易饱和。输出层用 Tanh输出范围 [-1,1]与数据预处理一致。view操作把一维向量恢复成图像形状方便后续输入判别器和保存图片。4.3 定义判别器判别器接收 1×28×28 图像先展平成 784 维向量再经过全连接层最终输出一个概率值。class Discriminator(nn.Module): def __init__(self): super(Discriminator, self).__init__() self.model nn.Sequential( nn.Linear(784, 512), nn.LeakyReLU(0.2, inplaceTrue), nn.Linear(512, 256), nn.LeakyReLU(0.2, inplaceTrue), nn.Linear(256, 1), nn.Sigmoid() ) def forward(self, img): # img 形状: (batch_size, 1, 28, 28) x img.view(img.size(0), -1) return self.model(x)这里使用 LeakyReLU 而不是 ReLU原因是判别器面对的是真假样本分布差异很大LeakyReLU 在负半轴保留了一个很小的梯度不容易让神经元“死掉”。0.2是负半轴斜率是常见配置。最终的 Sigmoid 输出是一个 0 到 1 的概率值表示输入为真实样本的置信度。4.4 定义损失函数与优化器二分类任务使用 BCEWithLogitsLoss 或 BCELoss。由于判别器最后已经有 Sigmoid这里使用 BCELoss 比较直观。# 损失函数 criterion nn.BCELoss() # 优化器 optimizer_G optim.Adam(generator.parameters(), lrlr, betas(0.5, 0.999)) optimizer_D optim.Adam(discriminator.parameters(), lrlr, betas(0.5, 0.999)) # 初始化网络 generator Generator(latent_dim).to(device) discriminator Discriminator().to(device)Adam 是 GAN 训练中最常见的优化器。值得注意的是betas(0.5, 0.999)这和默认值(0.9, 0.999)略有区别。设置第一动量系数为 0.5可以减小训练过程中的惯性让模型对参数变化更敏感这在 GAN 场景中被广泛采用。如果你用的是 PyTorch 2.xLazy模块会触发形状推理但这里没有使用避免兼容性问题。4.5 训练循环核心代码训练过程中每一轮迭代包含两部分先更新判别器再更新生成器。二者交替进行。generator.train() discriminator.train() for epoch in range(epochs): for batch_idx, (real_imgs, _) in enumerate(train_loader): current_batch real_imgs.size(0) real_imgs real_imgs.to(device) # 真实样本标签为 1伪造样本标签为 0 real_labels torch.ones(current_batch, 1, devicedevice) fake_labels torch.zeros(current_batch, 1, devicedevice) # ---------- 训练判别器 ---------- # 从正态分布采样噪声 z torch.randn(current_batch, latent_dim, devicedevice) fake_imgs generator(z) # 判别器对真实样本的损失 real_loss criterion(discriminator(real_imgs), real_labels) # 判别器对伪造样本的损失 fake_loss criterion(discriminator(fake_imgs.detach()), fake_labels) # 总损失是两者平均 d_loss (real_loss fake_loss) / 2 optimizer_D.zero_grad() d_loss.backward() optimizer_D.step() # ---------- 训练生成器 ---------- # 重新采样噪声生成新伪造样本 z torch.randn(current_batch, latent_dim, devicedevice) fake_imgs generator(z) # 生成器希望判别器把伪造样本识别为 1 g_loss criterion(discriminator(fake_imgs), real_labels) optimizer_G.zero_grad() g_loss.backward() optimizer_G.step()解释一下几个容易出错的地方fake_imgs.detach()训练判别器时我们不需要更新生成器。如果不 detach梯度会同时回传到生成器导致一次迭代更新了两个网络逻辑上混成一锅粥。训练生成器时重新采样了z也可以直接复用上一份fake_imgs。但确保不再对判别器做 detach因为生成器需要从判别器输出的梯度中学习。real_loss和fake_loss求平均是原始论文中的做法可以避免 batch 内真假样本数量不一致带来的偏差。4.6 保存与可视化生成效果每个 epoch 结束后把当前生成器产出的图片保存到result/目录方便观察训练效果。def save_samples(generator, epoch, save_dir./result, num_samples16): os.makedirs(save_dir, exist_okTrue) generator.eval() with torch.no_grad(): z torch.randn(num_samples, latent_dim, devicedevice) samples generator(z).cpu() samples (samples 1) / 2 # 从 [-1,1] 映射回 [0,1] grid torchvision.utils.make_grid(samples, nrow4) torchvision.utils.save_image(grid, os.path.join(save_dir, fepoch_{epoch:03d}.png)) generator.train()说明推理时使用torch.no_grad()关闭梯度计算节省内存并避免意外把噪声数据的梯度传入网络。make_grid将多张图片拼接成一张网格图便于直接查看。从 [-1,1] 映射回 [0,1] 是为了保存成正常可见的灰度图片。在训练循环内部每个 epoch 结束后调用一次print(fEpoch [{epoch1}/{epochs}] D loss: {d_loss.item():.4f}, G loss: {g_loss.item():.4f}) save_samples(generator, epoch 1)4.7 运行与预期结果将以上代码完整保存为train.py然后运行python train.py如果没有 GPU训练速度会比较慢。你可以把 epochs 调小例如 10 或 20先验证整个流程能跑通。预期你会看到类似输出Epoch [1/50] D loss: 0.6912, G loss: 0.6823 Epoch [2/50] D loss: 0.5801, G loss: 1.2301 ...第一次运行时如果目录中还没有 MNIST 数据PyTorch 会自动下载。输出 loss 时前几轮 D loss 和 G loss 都会在 0.6 到 1.0 附近徘徊。随着训练推进你再打开result/epoch_010.png会发现图片从一团噪声慢慢变成模糊的数字轮廓。这里需要特别说明loss 数值本身并不能直接说明生成质量。判别器 loss 降到很低可能是判别器“太强”了生成器跟不上生成器 loss 很低也不代表它生成的图片清晰、多样。GS 网络的评估最终要看视觉效果这也是 GAN 训练中比较复杂的地方。5. 训练过程中的常见问题与排查思路初学者第一次训练 GAN大概率会踩到下面几个坑。我们把常见问题整理成一张排查表。问题现象常见原因解决思路训练刚开始 loss 很快降到接近 0判别器过于强大生成器完全被骗不过去调小判别器学习率、减少判别器更新次数或增加生成器网络容量生成器 loss 持续很高图片仍然是噪声梯度消失生成器训练信号太弱将生成器损失从最小化log(1-D(G(z)))改为最大化log(D(G(z)))并检查是否存在梯度为 0 的情况生成图片只有少数几种数字模式坍塌Mode Collapse尝试更小的学习率、使用标签平滑、增大噪声维度、换用更复杂的架构生成图片有棋盘格纹路上采样操作不合理在卷积网络中使用ConvTranspose2d后跟上卷积层消除重叠或改用Upsample配合卷积D loss 和 G loss 剧烈震荡模型不稳定降低学习率、增大 batch size、加入批归一化层图片整体偏亮或偏暗数据归一化与生成器激活函数不匹配确认真实图片是否归一化到 [-1,1]生成器是否使用 Tanh 输出其中模式坍塌是 GAN 训练中最经典也最棘手的问题。它表现为生成器只找到了一小部分能被判别器接受的输出于是反复生成同一类图片。比如 MNIST 只生成“1”或“7”缺少多样性。解决模式坍塌没有万能药通常需要综合调节网络结构、损失函数和训练策略。另外初学者还容易遇到一个现象前几轮图片噪声明显之后突然变清晰然后又变模糊。这通常是训练过程不稳定的表现。建议你记录每个 epoch 的图片不要只看最后一个 epoch 的结果。保存历史输出这一点在你排查问题时能起到很大作用。6. 最佳实践与工程建议6.1 超参数选择与训练稳定性根据大量 GAN 工程的实践经验以下超参数是一个比较稳妥的起点学习率0.0002。Adam 动量betas (0.5, 0.999)。batch size64 或 128。噪声维度100 到 128。隐藏层激活生成器用 ReLU判别器用 LeakyReLU。输出层激活生成器用 Tanh判别器用 Sigmoid。在训练时可以适当“让判别器慢一点”。比如每更新一次生成器只更新一次判别器如果判别器 loss 下降太快就调整为每两次生成器更新后再更新一次判别器。6.2 标签平滑技巧把真实样本的目标标签从 1 改成 0.9这是一个非常简单但有效的技巧称为标签平滑Label Smoothing。它能防止判别器对真实样本过于自信从而给生成器留出更多改进空间。代码改动只需一行real_labels torch.ones(current_batch, 1, devicedevice) * 0.96.3 从全连接网络升级到卷积网络上面的示例使用全连接层便于理解但效果有限。真正应用时推荐使用 DCGANDeep Convolutional GAN架构生成器用ConvTranspose2d逐步放大特征图。判别器用Conv2d逐步缩小特征图。大量使用BatchNorm2d和 LeakyReLU。不在判别器中使用池化层而是用带步长的卷积替代。DCGAN 训练更稳定生成图片也更清晰。理解基础 GAN 之后可以把网络结构替换成 DCGAN超参数基本不用大改效果立刻会有明显提升。6.4 从原始 GAN 到 WGAN原始 GAN 的损失函数基于 JS 散度当真实分布和生成分布重叠很少时JS 散度会失去梯度方向。WGAN 改用 Wasserstein 距离从理论上缓解了训练不稳定和模式坍塌问题。随后出现的 WGAN-GP 又加了梯度惩罚项是目前应用中最常被使用的 GAN 变体之一。所以当你在基础 GAN 上遇到难调的稳定性问题时不用死磕原始损失函数可以直接跳到 WGAN 系列学习更先进的思路。6.5 项目落地时的建议如果你要把 GAN 用在真实项目中下面几条建议值得记住数据隐私训练数据可能包含人脸、身份信息使用前需要获得合规授权。生成内容审核GAN 生成图片可能被滥用务必在应用层加入内容过滤与来源标注。模型评估不能只看 loss要建立评价指标比如 FIDFréchet Inception Distance用于衡量真实图片和生成图片的分布距离。环境一致性使用torch.save保存模型时同时保存生成器和判别器的state_dict并记录当时的超参数方便复盘。如果你用自己的数据集训练注意先做数据清洗。GAN 对脏数据非常敏感样本中混入错误标签或大量重复图片都会让生成结果异常。7. 进阶学习路线如果你已经能跑通上面的代码GAN 基础算是入门了。接下来建议按这条路线深入先复现 DCGAN观察卷积层如何提升生成质量。再学习条件 GANConditional GANcGAN在生成时加入标签信息控制生成的数字类别。学习 WGAN 和 WGAN-GP理解损失函数改进背后的数学直觉。尝试 CycleGAN实现未配对图像风格迁移。阿里感兴趣的方向可以继续看 StyleGAN这是人脸生成领域非常稳定的架构。每一步都可以配合小实验来验证。比如在 MNIST 上先加一个条件标签让生成器输出指定数字再换到真实照片数据集观察模型对复杂分布的适应能力。GAN 的入门曲线其实不算特别高真正的门槛在于训练稳定性和调参经验。把上面的示例代码完整跑一遍再手动改几次超参数你很快就能体会到“博弈”这两个字在代码里到底是怎么发生的。等你能稳定训练出一个像样的 GAN深度学习生成模型这条支线就算真正走通了。
网站建设高端定制企业官网