
这次我们来看 AI Study 深度学习系列里一个绕不开的模块——GAN生成对抗网络。如果你正在学深度学习一定看过“AI 凭空画出一张人脸”这类演示大多数早期实现背后的主力模型就是 GAN。它由 Ian Goodfellow 等人在 2014 年提出思路非常反直觉不教模型直接背诵数据而是让两个网络互相博弈一个负责造假一个负责打假打到最后假货越来越接近真货生成器就练成了。这篇博文会把 GAN 基础这件事讲透什么是 GAN、它的数学目标怎么理解、如何用 PyTorch 从零写一个能生成手写数字的最小 GAN、训练时最容易踩哪些坑。和只看理论笔记不同这里会给出可直接复现的代码和完整的验证思路你可以在此基础上改参数、换数据集、甚至往 DCGAN 和 WGAN 方向继续扩展。先给结论GAN 不需要昂贵的标注数据它从真实样本中自学数据分布入门实验用 CPU 就能跑有 NVIDIA 显卡训练速度更快训练过程对超参数非常敏感调参属于常态。学会 GAN 基础之后再去看图像修复、超分辨率、风格迁移这些应用理解成本会低很多。适合读者正在学深度学习原理的人、做图像生成或 AI 绘画方向的学生和开发者、准备跑通第一个生成模型但不知道从哪下手的人。下面直接进入正题。1. GAN 核心知识速览在做代码实验之前先把 GAN 的关键属性整理成一张表方便快速判断这个方向适不适合现在的你。项目说明项目定位深度学习生成模型属于无监督/自监督学习范式提出时间2014 年Ian Goodfellow 等人提出核心机制生成器Generator与判别器Discriminator对抗训练训练方式不需要逐样本人工标注从真实数据分布中学习硬件门槛CPU 可跑通 MNIST 级别入门实验训练更复杂图像建议使用 GPU启动方式PyTorch 脚本训练训练完成后保存生成器权重是否支持 API模型本身不含 API训练完成后可自行封装推理服务是否支持批量训练过程按 mini-batch 批量进行推理阶段也可批量生成主要应用图像生成、图像修复、超分辨率、风格迁移、数据增强典型风险训练不稳定、模式坍缩、深度伪造合规风险表格里值得注意的几点GAN 不是一个“开箱即用”的软件项目而是一类模型框架。所以后面的操作重点不是部署某个仓库而是理解结构、跑通训练脚本、学会观察训练日志和生成结果。这一点和部署 WebUI、一键包完全不同阅读时不要混淆。2. GAN 解决什么问题适用场景与使用边界2.1 适合做什么GAN 最擅长的事情是学习一个数据集的分布然后生成和它相似但不完全相同的新样本。在图像方向上典型场景包括数据增强当真实样本不足时用 GAN 合成补充样本。这在医学图像、工业缺陷检测等样本稀缺场景有应用价值但合成样本不能替代真实数据使用前需要验证分布偏差。图像修复老照片破损区域修复、图像去噪去模糊这是“GAN 图像修复”方向的核心思路生成器负责补全缺失内容判别器负责判断补全结果是否足够真实。图像超分辨率从低分辨率输入恢复高分辨率细节判别器判断重建结果是否清晰自然。风格迁移照片转油画、白天转夜景、人脸属性编辑本质都是域到域的生成转换。异常检测用生成器重建正常样本重建误差大的区域可以定位为异常。2.2 不适合做什么GAN 不是万能的。需要严格可控输出的任务比如表格识别、证件文字识别用 GAN 做很容易出现内容错乱数据量极小的情况下GAN 很容易过拟合到几个训练样本上需要完全可解释结果的医疗诊断、金融风控场景也要先评估生成结果的可验证性。2.3 合规与安全边界GAN 生成人脸、声音、视频内容时必须确认肖像权和声音授权不能用于伪造身份、绕过核验或制作误导性内容。训练数据如果来自网络图片需要确认是否允许用于机器学习训练遵守版权规定。生成内容用于商用发布前要做人工复核并保留数据来源和处理记录。这些不是空话而是实际项目上线前必须考虑的问题。3. GAN 本地实验环境准备3.1 软硬件检查清单项目建议操作系统Windows / Linux / macOS 均可Python3.8 或更高版本推荐 3.10 左右深度学习框架PyTorch 2.xCPU 版本即可跑入门实验显卡可选有 NVIDIA 显卡并安装 CUDA 可加速训练磁盘空间MNIST 数据集约 20 MB加上依赖约 2 GB 左右端口如果后续封装 API 服务预留 8000 端口没有材料依据时版本不写死以本机环境为准。如果安装了其他版本的 PyTorch代码接口基本一致。3.2 创建 Python 环境并安装依赖推荐用 conda 创建独立环境避免污染系统 Python。conda create -n gan_study python3.10 -y conda activate gan_study pip install torch torchvision如果显卡驱动和 CUDA 环境已经配好并且需要 GPU 加速可以按 PyTorch 官网给出的命令安装对应 CUDA 版本。第一次做 GAN 实验用 CPU 版把流程跑通最稳妥。3.3 准备数据集下面代码会通过torchvision.datasets.MNIST自动下载手写数字数据集。第一次运行需要联网数据集会保存到./data目录。如果下载速度慢可以手动下载后放到对应目录注意解压后的文件结构需要符合 torchvision 的读取要求。4. GAN 最小实现定义网络与启动训练4.1 生成器与判别器定义生成器输入一个随机噪声向量输出一张 28x28 的灰度图像判别器输入一张图像输出一个 0 到 1 之间的真实性分数。这里使用全连接结构MNIST 级别完全够用。import torch import torch.nn as nn class Generator(nn.Module): def __init__(self, latent_dim100, img_dim28 * 28): super(Generator, self).__init__() self.model nn.Sequential( nn.Linear(latent_dim, 256), nn.ReLU(), nn.Linear(256, 512), nn.ReLU(), nn.Linear(512, img_dim), nn.Tanh(), ) def forward(self, z): return self.model(z) class Discriminator(nn.Module): def __init__(self, img_dim28 * 28): 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(), ) def forward(self, img): return self.model(img)注意生成器最后一层用Tanh把输出压到 [-1, 1] 区间对应真实图像归一化到 [-1, 1] 之后的范围。判别器最后一层用Sigmoid输出真实性概率。这两个激活函数是 GAN 基础实现的标准选择。4.2 训练循环训练逻辑分两步先固定生成器训练判别器让判别器学会区分真实图像和伪造图像再固定判别器训练生成器让生成器学会骗过判别器。两者交替优化就构成了对抗训练。import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms latent_dim 100 batch_size 64 epochs 50 lr 0.0002 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(0.5, 0.999)) d_optimizer optim.Adam(discriminator.parameters(), lrlr, betas(0.5, 0.999)) transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)), ]) dataset datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) dataloader DataLoader(dataset, batch_sizebatch_size, shuffleTrue) for epoch in range(epochs): for i, (real_imgs, _) in enumerate(dataloader): real_imgs real_imgs.view(real_imgs.size(0), -1).to(device) cur_batch real_imgs.size(0) real_labels torch.ones(cur_batch, 1).to(device) fake_labels torch.zeros(cur_batch, 1).to(device) # 训练判别器真实图像判为 1伪造图像判为 0 z torch.randn(cur_batch, latent_dim).to(device) fake_imgs generator(z) d_loss_real criterion(discriminator(real_imgs), real_labels) d_loss_fake criterion(discriminator(fake_imgs.detach()), fake_labels) d_loss d_loss_real d_loss_fake d_optimizer.zero_grad() d_loss.backward() d_optimizer.step() # 训练生成器让伪造图像被判别器判为 1 z torch.randn(cur_batch, latent_dim).to(device) fake_imgs generator(z) g_loss criterion(discriminator(fake_imgs), real_labels) g_optimizer.zero_grad() g_loss.backward() g_optimizer.step() if i % 200 0: print(fEpoch [{epoch}/{epochs}] Batch [{i}/{len(dataloader)}] fD Loss: {d_loss.item():.4f} G Loss: {g_loss.item():.4f}) torch.save(generator.state_dict(), generator_mnist.pth)4.3 启动训练把上面代码保存为train_gan.py然后在终端运行python train_gan.py训练结束后当前目录会多出generator_mnist.pth文件和./data数据集目录。终端日志格式如下数值仅用于展示输出结构不同机器和超参数下差异很大Epoch [0/50] Batch [0/938] D Loss: 0.6914 G Loss: 0.7412 Epoch [0/50] Batch [200/938] D Loss: 0.5387 G Loss: 1.1204正常现象是 D Loss 和 G Loss 都在波动不要求某一方永远下降。如果生成器训练过快判别器会长期处于“打假失败”状态日志里 G Loss 会很低而 D Loss 会持续偏高这就是对抗失衡的信号。5. GAN 训练效果验证与可视化5.1 保存生成图片只观察 loss 数值不够直观最有效的验证方法是每隔几个 epoch 保存一批生成图像肉眼判断数字是否成型。import torchvision.utils as vutils import os save_dir ./generated os.makedirs(save_dir, exist_okTrue) def save_generated_images(epoch, num_images16, save_pathsave_dir): generator.eval() z torch.randn(num_images, latent_dim).to(device) with torch.no_grad(): fake_imgs generator(z).view(-1, 1, 28, 28) vutils.save_image(fake_imgs, f{save_path}/epoch_{epoch:03d}.png, nrow4, normalizeTrue) generator.train()把这个函数加到训练循环里每个 epoch 结束调用一次。生成目录下会出现epoch_000.png、epoch_001.png这类文件打开就能看到训练过程的可视化结果。5.2 判断训练是否成功判断标准可以分三层第一层生成图像是否像手写数字。训练到 20 到 50 个 epoch 时应该能看到清晰可辨的数字轮廓。第二层多样性是否足够。如果保存的 16 张图几乎一样说明出现模式坍缩生成器只学会了输出少量固定样本。第三层判别器 loss 是否处于合理区间。理想情况下 D Loss 和 G Loss 在 0.5 到 2 之间波动稳定到 0 或者涨到很大都说明训练失衡。5.3 常见失败表现如果图像整片模糊可能是网络太简单或训练轮数不够如果图像全黑检查真实图像的归一化范围是否和生成器输出范围匹配如果数字轮廓可见但细节粗糙可以尝试减小学习率、增加隐藏层宽度。这些都属于调参问题不是代码逻辑问题。6. 批量训练与推理接口封装6.1 mini-batch 批处理GAN 训练天然是批量任务。上面代码里的DataLoader按 batch_size 分批读取数据每批 64 张图像一个 epoch 共 938 个 batch。批量大小直接影响训练速度和稳定性batch 太小梯度噪声大batch 太大显存占用高训练速度不一定更快。实际使用中从 32 到 128 之间尝试即可。6.2 将生成器封装为推理服务训练完成后的生成器可以单独用于推理。GAN 基础版本不携带 API但你可以用 FastAPI 自己封装一个简单的生成接口。下面是一个通用示例实际使用时需要把第 4 节的Generator类定义复制进来并将模型路径替换成自己的 checkpoint。from fastapi import FastAPI from pydantic import BaseModel import torch app FastAPI() latent_dim 100 device torch.device(cpu) # 这里需要导入或重新定义 Generator 类然后加载权重 # generator Generator(latent_dim).to(device) # generator.load_state_dict(torch.load(generator_mnist.pth, map_locationdevice)) # generator.eval() class GenerateRequest(BaseModel): count: int 8 app.post(/generate) def generate(req: GenerateRequest): z torch.randn(req.count, latent_dim).to(device) with torch.no_grad(): imgs generator(z) return {count: req.count, shape: list(imgs.shape)}启动服务uvicorn app:app --host 127.0.0.1 --port 8000用 curl 测试curl -X POST http://127.0.0.1:8000/generate \ -H Content-Type: application/json \ -d {count: 8}返回的 JSON 会包含生成张量的形状。需要强调的是这个接口示例只是为了演示工程化思路GAN 基础实验阶段不必急着封装服务先把训练跑通更重要。7. 资源占用与性能观察7.1 观察显存和 CPU 占用训练时可以用以下命令实时观察资源占用watch -n 1 nvidia-smiPython 进程内部也可以查看当前显存分配print(torch.cuda.memory_allocated() / 1024 / 1024, MB)这里的数字会随 batch_size、图像分辨率和模型复杂度变化不存在固定的“推荐显存”数值。以本机实际测试为准。7.2 CPU 与 GPU 的差异CPU 训练 MNIST 级别 GAN 是可以接受的50 个 epoch 可能需要几十分钟到几小时取决于 CPU 核心数。GPU 训练时瓶颈主要在数据加载和 batch 之间的同步开销MNIST 这种小图像差异不明显换成 128x128 以上的图像后 GPU 优势才会完全体现出来。7.3 影响资源占用的因素batch_size显存占用近似线性增长调试时先调小。图像分辨率分辨率提高一倍特征图显存占用可能增长数倍。生成器和判别器宽度隐藏层节点越多参数和中间激活值越多。优化器状态Adam 会保存一阶和二阶动量参数量越大额外显存越大。降低显存占用的常用手段包括减小 batch_size、降低图像分辨率、使用混合精度训练、改用梯度累积模拟大 batch。GAN 训练中梯度累积要注意判别器和生成器的更新频率是否仍然均衡。8. GAN 训练常见问题与排查方法问题现象可能原因排查方式解决方案生成图像几乎一样模式坍缩生成器只学到少数样本查看不同 epoch 的生成图计算多样性和判别器 loss 走势减小学习率、使用 WGAN 等变体、增加噪声维度、调整生成器结构判别器 loss 快速降到接近 0判别器过强生成器完全无法对抗查看生成图像是否空白确认生成器梯度是否正常降低判别器学习率、增加生成器容量、减少判别器训练次数生成器 loss 不下降生成器训练不充分或被判别器压制单独查看生成器 loss 曲线检查梯度是否消失调整网络结构、更换激活函数、调整学习率训练时显存溢出batch_size 过大或分辨率过高用nvidia-smi查看显存使用减小 batch_size、降低分辨率、开启混合精度MNIST 下载失败网络问题或数据目录权限问题检查终端网络提示确认./data目录可写手动下载数据集放入对应目录或更换镜像源生成图片全黑或全白归一化范围不匹配对比生成器输出范围和实际图像范围确认真实图像归一化到 [-1, 1]生成器最后一层使用 Tanh训练速度越来越慢日志、图片保存或数据加载成为瓶颈检查磁盘占用和 dataloader 的 num_workers减少日志频率、独立保存图片、增加 num_workers排查时建议先确认模型能过拟合一个小数据集再扩大数据规模。如果小数据集上生成器和判别器都无法产生有意义输出问题通常出在网络结构或损失函数而不是数据量。9. 从 GAN 基础到进阶变体学会基础 GAN 之后接下来可以按这条路线继续深入DCGAN把全连接层替换成卷积层使用 BatchNorm 和 LeakyReLU训练稳定性明显提升适合从基础代码平滑过渡。CGAN在生成器和判别器输入中加入类别标签让生成结果受条件控制适合按指定类别生成图像。WGAN 与 WGAN-GP把损失函数换成 Wasserstein 距离缓解模式坍缩和训练不稳定问题是工程上更常用的方案。StyleGAN引入风格注入和渐进式训练生成人脸质量显著提高但对数据量和计算资源要求也更高。如果要落地“GAN 图像修复”应用常见做法是先用 GAN 学习正常图像的纹理先验再结合掩码区域重建输入破损图像和掩码生成器补全缺失区域判别器判断补全区域是否真实。理解基础 GAN 的对抗思想后这类应用的代码会好读很多。10. 最佳实践与合规使用建议做 GAN 实验时建议从最小配置开始先用 10 到 20 个 epoch 跑通流程确认代码无误后再增加训练轮数。固定随机种子能帮助复现实验结果尤其是判别器和生成器两个网络的初始化权重对训练结果影响很大。import random import numpy as np seed 42 random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)训练日志建议写入文件而不是只打印到终端方便事后对比不同超参数的训练曲线。模型权重、生成图片、数据集分别建立目录管理避免一个文件夹塞满各种中间文件。批量实验时每组实验要记录学习率、batch_size、网络结构、训练轮数和最终效果否则很难定位是哪一项改动导致了结果变化。合规方面需要再次强调使用真实人物照片、声音作为生成素材时必须提前获得授权训练数据来源要记录并确认版权允许生成内容如果涉及身份信息或敏感场景不能用于欺骗、伪造或规避平台规则。发布和商用之前必须对生成内容进行人工复核。在新一轮实验里建议把第 4 节的生成器换成 DCGAN 的卷积结构把损失函数换成 WGAN-GP 的 Wasserstein 距离感受一下训练稳定性的差别。GAN 基础这一章到这里就完整了接下来可以按本文的进阶路线继续往下学。