做生成式图像的朋友应该都有过类似经历:在社区里刷到别人用 GAN 生成的逼真人像、抽象艺术或游戏原画,自己也兴致勃勃配了主机、装了 Ubuntu,结果训练跑到一半发现生成出来全是同一张脸的畸形变体,或者干脆全是噪点。我在这篇文章里围绕一个实际问题展开:如何在 Ubuntu 22.04 上搭建 GAN 训练环境、从零开始训练一个能出图的生成对抗网络,以及后续如何针对训练效果做稳定性优化。文章会结合我的实操经验,把能复现、能避坑的关键路径都拆开讲清楚,适合对生成式模型感兴趣的中级开发者、在 Ubuntu 上想快速开始实验的研究生或工程师,以及所有正被 GAN 训练不稳定折磨的人。
先说明一下,这篇文章不会直接给你一个“万能模型”,而是帮你建立一套从环境到数据、从架构到训练调优的完整决策链。GAN 不像普通分类模型,拿起来就能跑,它更像两个互相较劲的对手在持续博弈,任何一个环节失衡都会导致整个训练崩盘。下面直接从 GAN 的内核开始讲,这是理解后面所有调优技巧的前提。
1. 为什么GAN训练经常“不听话”:核心机制与主要坑点的地图
首先要搞清楚,GAN 不是一个“直接优化到目标”的模型,而是两个网络在互相博弈中共同进步的过程。生成器 G 的任务是从随机噪声 z 映射到图像 x;判别器 D 的任务则是把真实图像和 G 生成的假图像区分出来。原始 GAN 的核心目标函数是:
min_G max_D V(D, G) = E_x~p_data[log D(x)] + E_z~p_z[log(1 - D(G(z)))]
这个公式看上去很简洁,但工程上的所有麻烦都藏在这对 min 和 max 之间。整个训练是零和博弈:生成器每变强一点,判别器也随之变强;但如果一方领先太多,另一方就会陷入梯度消失或梯度爆炸的境地。要弄清它为什么改良起来如此费劲,得从下面几个机制说起。
1.1 生成器和判别器的博弈到底在优化什么
这里的 p_data 是真实数据分布,p_z 是噪声分布,通常是标准正态分布或均匀分布。生成器的目标是让生成分布 p_g 尽可能接近真实分布 p_data,实现方式有两种:一是直接最小化生成样本与真实样本之间的某种距离;二是通过骗过判别器来反向调整生成器,使其产生更逼真的图片。后者避免了直接计算复杂的生成分布密度,但也带来了训练不稳定的代价。
判别器 D 在经典 GAN 中是一个二分类器,对真实图像输出接近 1 的分数,对生成图像输出接近 0 的分数。生成器的损失并不直接用 D(G(z)) 接近 1,因为它无法直接控制 D 的输出,只能通过更新生成器 G 的参数,朝 D 更认可的方向移动。于是整个训练过程就像一场猫鼠游戏:判别器不断完善鉴别标准,生成器在同一条标准的压力下不断调整输出。
随着训练推进,一个关键指标是判别器的判别能力不能太强——如果它瞬间就把真假分得干干净净,生成器的梯度会因为 sigmoid 的饱和区而变得微乎其微,导致生成器长期得不到有效更新。这也是原版 GAN 最难伺候的地方,后面 WGAN-GP 等一系列改进,本质上都是在解决这个“判别器过强”的问题。
1.2 模式崩塌(Model Collapse)与梯度消失的本质
模式崩塌是指生成器最终只学会生成少数几类样本,这些样本在各批次之间几乎完全相同。这在实战中高频出现。一个典型的诱发因素是:当判别器被训练过度且对某些样本的判别极其自信时,生成器的损失梯度只会朝向极少数最容易骗过判别器的方向收敛,多样性彻底丢失。
更复杂的版本是循环性模式崩塌:生成器某一轮学会了某种偏方骗过判别器,下一轮判别器调整策略后,生成器又换成另一个单一骗法。此时从损失曲线看似乎一切正常,但生成的样本高度同质,根本覆盖不了真实分布。
梯度消失则通常发生在训练早期或者判别器收敛过头的情况下。经典 GAN 的判别器使用二进制交叉熵,当判别器能 100% 分出真假时,生成器得到的梯度接近零;而对生成器求导时需要对 log(1 - D(G(z))) 做反向传播,在判别器非常强时这部分梯度也会趋近于零。这就是很多新手第一次跑 GAN 时经常遇到的状况:loss 半天不动,生成图全是噪点。
1.3 训练G之前你必须理解的工程成本
可能有人觉得训练一个 GAN 也就是跑个几步的事,但我的体会是:GAN 的调参工作量是很多传统监督模型的数倍。传统分类模型只需关注训练集损失和验证集指标,而 GAN 需要同时观察生成器损失、判别器损失、生成图片的视觉质量、多样性、是否在反复横跳以及是否出现过拟合。任何单一指标都不能告诉你“模型到底训练好了没有”,必须组合起来综合判断。
硬件成本同样不能小看。同样一张 256×256 的人像,如果用最基础的 DCGAN 尝试,几个小时就能看出端倪;如果换到更高分辨率的模型,一块 24GB 显存的显卡也得按天为单位来训练。这对显存、硬盘空间和散热都提出了额外要求。建议动手前先规划好预算和实验周期,不要心存侥幸,指望第一次实施就一键出大片。
2. Ubuntu 22.04环境搭建:从驱动检查到PyTorch可复现配置
现在进入最劝退新人的环境搭建环节。环境不对,后面一切白搭。我在 Ubuntu 22.04 上重建过多次 GAN 训练机器,为什么选 22.04 LTS?它足够稳定、软件源更新及时,对 CUDA 和深度学习框架的兼容性也更好。下面按顺序把每一步拆开讲,并附带真实会踩到的坑。
2.1 检查GPU可用性与安装NVIDIA驱动
第一步先确认机器上到底有没有 NVIDIA GPU,以及驱动是否已经可用。打开终端执行:
nvidia-smi如果没有任何输出,说明驱动未安装或没有 NVIDIA 显卡。Ubuntu 22.04 上如果之前装过开源驱动 nouveau,建议先把相关内核模块禁用掉,否则会和 NVIDIA 闭源驱动冲突。做法是编辑 /etc/modprobe.d/blacklist-nouveau.conf:
echo "blacklist nouveau" | sudo tee /etc/modprobe.d/blacklist-nouveau.conf echo "options nouveau modeset=0" | sudo tee -a /etc/modprobe.d/blacklist-nouveau.conf保存后执行 sudo update-initramfs -u,重启后再安装驱动。
安装驱动我个人推荐直接用发行版仓库的版本,比如 nvidia-driver-535 或更高版本,命令如下:
sudo apt update sudo apt install nvidia-driver-535安装完成后重启,再次运行 nvidia-smi 应该能看到 GPU 型号、驱动版本和显存信息。如果输出一堆错误,比如 "Unable to determine the device handle",多半是驱动没完全加载,可以运行 dmesg | grep nvidia 查看内核日志定位原因。
注意:不同驱动分支对 CUDA 版本的兼容性不同,先确认显卡架构(比如 Ampere 还是 Ada),再选择驱动分支。没必要追最新版,稳定优先,性能差距通常很小。
2.2 安装CUDA Toolkit和cuDNN
驱动装好后安装 CUDA Toolkit。这里有个常见误解:很多人以为必须手动装 CUDA 才能用 PyTorch,其实 PyTorch 默认自带 CUDA runtime,系统级 CUDA 对常规训练并非必需。不过为了跑一些底层算子或自定义 kernel,装一份也无妨。
最简单的方式是直接用官方 apt 源或仓库包:
sudo apt install nvidia-cuda-toolkit但这个版本可能较旧。如果你想装 CUDA 12.x 的特定版本,官方提供 runfile 安装方式。装完验证:
nvcc --versioncuDNN 的安装需要注意与 CUDA 版本匹配,直接用 pip 安装 nvidia-cudnn-cu12 是省心选择,这样它和 PyTorch 共用一套运行库,基本不用配置额外路径。
2.3 创建Python虚拟环境并安装PyTorch
推荐用 Python 自带的 venv 或 conda 创建独立环境。系统自带的 Python 3.10 + venv 足够:
python3 -m venv gan_env source gan_env/bin/activate pip install --upgrade pip接着安装 PyTorch。我的习惯是用官方指定源安装对应 CUDA 版本,比如 CUDA 12.1:
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121装完后先别急着写模型,跑一段检查:
import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))如果 torch.cuda.is_available() 返回 True,环境基本就通了。这里有一个很常见的坑:PyTorch 能正常 import,但 CUDA 就是不可用,十有八九是装了 CPU 版本,或者 CUDA 与驱动版本不匹配。只要用官方指定 index-url 安装基本就能规避。
2.4 验证环境与常见错误排查
为了确保后续代码不会无声无息地在 CPU 上跑,建议在训练脚本里加一段设备初始化逻辑:
device = "cuda" if torch.cuda.is_available() else "cpu" if device == "cpu": print("警告:当前使用CPU,GAN训练将非常慢")随后跑一个小的卷积运算测试:
x = torch.randn(4, 3, 256, 256).to(device) y = torch.nn.Conv2d(3, 16, 3, padding=1).to(device)(x) print(y.shape)这个测试能快速验证 GPU 上的卷积算子是否正常。如果出现 OOM(Out of Memory),要么是显存不够,要么是 PyTorch 的缓存机制没释放。可以通过 torch.cuda.empty_cache() 手动释放,同时检查批处理大小。
另一个实用技巧是用 nvidia-smi -l 1 持续监控显存占用。GAN 训练比普通分类更吃显存,因为生成器和判别器要同时放在显存里,两者在反向传播时都占用空间。显存小的机器建议先跑 64×64 或 128×128 的小分辨率实验,确认整条管线没问题后再上高分辨率。
3. 数据管道设计:高质量图像生成的“上游”决定一切
很多人在环境装好后第一件事是写网络结构,结果数据没打理好,训练起来怎么调都效果差。在 GAN 的实战里,数据管道的权重可能和网络结构一样大。高质量输入不能保证高质量输出,但垃圾输入几乎必然导致垃圾输出。
3.1 数据集选择:为什么 CelebA 和 FFHQ 适合做人脸生成
如果你是想做人脸生成,经典公开数据集有两个:CelebA(超过 20 万张名人脸部图)和 FFHQ(7 万张 1024×1024 高清人脸图)。CelebA 的优势是数量大、标注丰富、下载方便,缺点是原始分辨率只有 178×218,且包含不少姿态各异的图像;FFHQ 质量更高,构图更整齐,但获取流程相对繁琐。建议新手用 CelebA 或它的对齐版本,先把流程跑通,再考虑 FFHQ 或自建数据。
如果你做的不是人脸,而是特定风格的内容,比如手绘建筑、二次元人物立绘、汽车设计草图,核心要求是一样的:数据要干净、类别要统一、画面构成要相近。比如做二次元人脸,就不要把半身像、全身像、含大量背景的图混在一起,否则生成器会无所适从,最后学会在一个框里平均所有风格,输出一团说不清是什么的东西。
3.2 预处理流程:归一化、尺寸统一和数据增强
图像输入网络前都要转成张量。PyTorch 里通常会做这几步:
- 将所有图片缩放到同一尺寸,用 torchvision.transforms.Resize
- 转换成张量,用 ToTensor()
- 归一化到 [-1, 1],因为生成器最终用 tanh 激活,输入区间保持一致能避免图片失真
一个常见变换组合是:
from torchvision import transforms transform = transforms.Compose([ transforms.Resize((128, 128)), transforms.CenterCrop((112, 112)), transforms.Resize((128, 128)), transforms.RandomHorizontalFlip(p=0.5), transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), ])这里的 Resize-CenterCrop-Resize 组合,用于去除图片中不重要的边缘信息并保留主体,尤其是当原始数据里背景占比差异大时效果明显。如果数据本身构图集中,直接 Resize 到目标尺寸即可。
数据增强在 GAN 里需要谨慎使用。RandomHorizontalFlip 对人脸和自然图像几乎总是有效;而 RandomRotation 和 ColorJitter 则要小心,它们可能改变语义,尤其在纹理敏感的医学图像上容易带来反效果。对人脸生成我通常只用水平翻转,加上少量随机裁剪,避免破坏五官结构。
3.3 DataLoader 配置和内存管理
DataLoader 是 PyTorch 的数据迭代核心。基础配置如下:
from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder dataset = ImageFolder(root="data/train", transform=transform) dataloader = DataLoader(dataset, batch_size=64, shuffle=True, num_workers=4, pin_memory=True)两个容易被忽略但影响较大的参数:
- num_workers:调成 4 到 8 可以充分利用 CPU 多核做图片预处理,但机器 CPU 不强时,设置过高反而会因为线程调度拖慢速度。
- pin_memory:设为 True 能锁定内存页,加快 GPU 传输,如果 RAM 紧张也可以关掉。
数据加载的另一个坑是 IO 瓶颈。我试过把几万张小图放在机械硬盘上直接读,GPU 利用率时常只有 40% 左右。后来把数据集转成 WebDataset 格式,或者直接放到 SSD 上,GPU 利用率能稳定拉到 90% 以上。如果数据量不大,比如 10GB 以内,最简单的做法是把数据提前读入内存,注意别超 RAM 上限即可。
4. GAN架构选型:从DCGAN到WGAN-GP的迭代逻辑
网络结构决定了一个 GAN 的天花板。这一部分我不想把每个模型完整代码都贴一遍,而是重点讲清“我们为什么这样选”“每一步解决了什么问题”,让你在遇到更复杂模型时有能力独立判断。
4.1 为什么从DCGAN开始(以及它的局限性)
DCGAN,深度卷积生成对抗网络,是最经典的把卷积引入 GAN 的架构。它的主要贡献是:
- 用步长卷积替代全连接层和池化,使生成器和判别器能端到端学习空间特征
- 生成器中大量使用 Batch Normalization,稳定深层网络的训练
- 判别器用 LeakyReLU,避免梯度消失
- 生成器输出用 Tanh,把像素值限制在 [-1, 1]
DCGAN 结构清晰,代码量少,在 64×64 分辨率的简单数据集上效果可接受,是新手“抄作业”的首选。但它的短板明显:对超参数极其敏感,在 128×128 以上分辨率容易不稳定;如果不加改进,模式崩塌率较高。所以 DCGAN 适合做流程验证和 baseline,不适合直接用于真正要产出高质量大图的场景。
4.2 WGAN-GP的原理:用梯度惩罚换来稳定训练
我在实际项目里从 DCGAN 换到 WGAN-GP 后,训练过程明显省心不少。WGAN-GP 把损失换成 Wasserstein 距离,这解决了原始 GAN 中“判别器越强梯度越小”的问题。Wasserstein 距离的特点是,即使两个分布完全不重叠,也能给出有意义的梯度方向。
但 WGAN 最初的实现需要用权重裁剪来满足判别器的 Lipschitz 约束,这带来了训练偏差。WGAN-GP 的改进是使用梯度惩罚:在真实样本和生成样本之间随机插值采样,并对这些采样点要求判别器梯度的范数尽可能接近 1。这样既满足约束,又不至于让判别器的权重分布扭曲。
梯度惩罚部分的代码大致如下:
def gradient_penalty(discriminator, real, fake, device): batch_size = real.size(0) alpha = torch.rand(batch_size, 1, 1, 1, device=device) interpolate = (alpha * real + (1 - alpha) * fake).requires_grad_(True) d_interpolate = discriminator(interpolate) fake = torch.ones_like(d_interpolate, requires_grad=False) gradients = torch.autograd.grad( outputs=d_interpolate, inputs=interpolate, grad_outputs=fake, create_graph=True, retain_graph=True, only_inputs=True, )[0] gradient_norm = gradients.view(batch_size, -1).norm(2, dim=1) return ((gradient_norm - 1) ** 2).mean()这一小段是整个 WGAN-GP 的核心。需要注意的点:
- real 和 fake 的尺寸必须完全一致
- interpolate 需要 requires_grad_(True),否则 autograd 拿不到梯度
- 在优化判别器时,保留计算图是必要的,否则下一个 batch 的反向传播会报错
梯度惩罚的系数 lambda_gp 通常取 10,太小约束力不够,太大又容易让训练偏离目标。实际中可以做成超参进行小范围搜索。
4.3 PGGAN/StyleGAN的渐进式训练思路
如果你最终想要高分辨率图像,比如从 256×256 起步甚至到 1024×1024,渐进式生成方法是值得了解的。PGGAN 的思路是:先让网络在低分辨率,例如 4×4、8×8 上稳定训练,再逐步向网络中增加更高分辨率层,让模型由粗到细学习细节,而不是一开始就硬啃 1024×1024 的像素空间。
StyleGAN 系列在 PGGAN 基础上更进一步,把生成过程变成风格注入:不同尺度的特征由不同风格向量控制,提高了生成图像的语义解耦能力,在人脸生成、纹理生成这些需要精细可控风格的任务里优势明显。但代价是显存和耗时开销都非常大。以 1024×1024 人脸为例,在单卡高性能 GPU 上也需要数天迭代才有可用效果。所以日常项目里,我通常建议先弄清任务是“快速验证”还是“极致效果”,再决定是否上渐进式架构,不要一味追求大模型。
4.4 损失函数与训练循环的关键代码实现
以 WGAN-GP 为例,训练循环的简版框架如下:
for epoch in range(num_epochs): for i, (imgs, _) in enumerate(dataloader): imgs = imgs.to(device) real_imgs = imgs batch_size = real_imgs.size(0) # 训练判别器 z = torch.randn(batch_size, latent_dim, 1, 1, device=device) fake_imgs = generator(z) d_real = discriminator(real_imgs) d_fake = discriminator(fake_imgs) gp = gradient_penalty(discriminator, real_imgs, fake_imgs, device) d_loss = -torch.mean(d_real) + torch.mean(d_fake) + lambda_gp * gp d_optimizer.zero_grad() d_loss.backward() d_optimizer.step() # 每隔一定间隔训练生成器 if i % n_critic == 0: z = torch.randn(batch_size, latent_dim, 1, 1, device=device) fake_imgs = generator(z) g_loss = -torch.mean(discriminator(fake_imgs)) g_optimizer.zero_grad() g_loss.backward() g_optimizer.step()两个需要特别留意的细节:
- WGAN-GP 的生成器损失不是让 D(G(z)) 接近 1,而是让 D(G(z)) 的均值尽可能大,即 -mean(D(fake)) 最小化。因为 Wasserstein 距离是无界的,判别器输出不再是概率,而是一种打分,所以不能用交叉熵。
- n_critic 是判别器每更新几次才更新一次生成器。WGAN-GP 的设计里通常判别器多训练几轮,以保证 Wasserstein 距离估计准确,常见取值 1~5。但注意,如果 GPU 显存有限、batch size 偏小,n_critic 过高会让训练变慢。
- 在 WGAN-GP 的判别器里不要用 BatchNorm,因为 batch 维度上的统计量会破坏单样本梯度的连续性。这就是为什么它的判别器通常用 InstanceNorm 或 LayerNorm。
5. 训练优化实战:稳定化技巧与超参数调优
网络结构选好了,距离真正跑出可用结果还差关键一步:把训练过程调稳。大多数 GAN 项目失败,不是架构本身有问题,而是超参数和训练策略的细节处理失当。下面这些优化技巧来自我的实际尝试和多个项目复盘,值得一条条对照检查。
5.1 学习率和优化器选择的不同方案
GAN 的传统共识是:生成器和判别器各自使用独立的学习率,判别器的学习率通常不宜过高。早期 DCGAN 常用 Adam 优化器,learning_rate 取 2e-4,betas 取 (0.5, 0.999)。这里的 beta1 取 0.5 而不是默认的 0.9,是为了减少训练初期的振荡。在对抗训练下,带有强动量的 beta1=0.9 容易让参数在极小值附近左右横跳。
WGAN-GP 最常见的配置是 Adam,生成器和判别器 lr 都取 1e-4,也有人把生成器学习率设为判别器的一半,降低生成器对当前判别器状态的过拟合风险。如果你的 loss 出现剧烈震荡,除了检查梯度惩罚系数,第一件事就是降低学习率,每次减半重新跑。
另一类做法是使用学习率衰减。GAN 里指数衰减并非必须,线性衰减配合固定训练步数通常效果不错,它让训练后期参数变化更小,收敛更平稳。如果用的是渐进式 GAN,通常在切换分辨率时重置学习率,或在过渡阶段用较低的 lr。
5.2 批大小与分辨率的关系
批大小对 GAN 训练有直接影响:批大小太小,判别器对数据分布的估计偏差大,梯度噪声大;批大小太大,显存很快吃满,实际收益却有限。以我的经验:
- 128×128 及以下,batch size 32 到 64 都还不错
- 256×256,如果显存有限,batch size 16 到 24 是可以接受的下限,再小可以用梯度累积来模拟更大 batch
- 512×512 及以上,单纯依赖 batch size 很难同时保证显存和稳定性,需要考虑 patch 式训练、混合精度等
一个实用技巧是:第一次训练不要急着上高 batch,先用显存能放下的最大值跑几十步,确认 loss 正常下降后,再逐步调整。Batch size 对 WGAN-GP 的梯度惩罚影响也很大,过小时梯度惩罚会被少数插值样本带偏,整体训练不稳。
5.3 标签平滑、谱归一化和特征匹配的实践效果
标签平滑是原始 GAN 里常用的技巧:把真实样本标签从 1 平滑为 0.9 或 0.8,避免判别器过度自信导致梯度消失。在 WGAN-GP 中,由于用的是无界打分,标签平滑并不直接适用,但可以在判别器输出上加一点小噪声,效果类似。
谱归一化是另一种稳定判别器的手段,它作用在每个权重矩阵上,通过约束最大奇异值来限制 Lipschitz 常数,和 WGAN-GP 的梯度惩罚本质目标一致,但实现更轻量,不用每次都计算插值梯度。后来的不少高分辨率 GAN 直接选用谱归一化而不是梯度惩罚。两者取舍是:谱归一化更稳但可能稍微限制表达力;梯度惩罚更通用但对超参更敏感。
特征匹配是在判别器中间层加一个额外损失,要求生成器的中间特征与真实图片的中间特征差异最小。这可以让生成器学到更有语义信息的多尺度特征,而不只是盯最后输赢。实际效果因数据集而异,建议自己实验对比,不要盲从网上流传的“最佳配置”。
5.4 训练稳定性监控:判别器损失和生成器损失的解读
监控训练不要只看最终成图,更要看两条损失曲线的相对关系。以 WGAN-GP 为例:
- 判别器在真实样本上的得分 D(real) 应稳定在某个正值附近且缓慢上升。如果剧烈上涨,说明判别器很快记住了训练集,可能有过拟合风险。
- 生成器的得分 D(fake) 应缓慢上升,但不能瞬间逼近 0。如果 D(fake) 一直为负且持续下降,说明生成器严重掉队,训练方向有问题。
- 如果两条分数趋近于同一水平,比如都在 0 附近波动,说明生成分布与真实分布已经比较接近,训练进入平稳阶段。
这里有个小技巧:把训练过程中的生成样本每隔固定迭代次数保存下来,拼成一张纵向时间序列图。这样你能一眼看出模型在第几千步开始出现清晰轮廓、第几万步开始有细节,甚至发现模型在什么时刻开始崩坏。直观程度远胜于只看数字曲线。
6. 模型质量评估:FID/IS指标解析与常见失败模式快速诊断
很多人训练完 GAN,最后评判效果时常只凭“肉眼看像不像”。在较窄的应用场景里这可行,但在需要迭代模型的真实项目里远远不够。GAN 没有真实标签,我们需要一套能衡量生成分布与真实分布接近程度的客观指标。
6.1 FID指标是什么:它能和不能告诉你什么
FID,全称 Fréchet Inception Distance,先用预训练 InceptionV3 网络分别提取真实图片和生成图片的特征向量,再假设这些特征向量服从高斯分布,最后计算两个高斯分布之间的 Fréchet 距离。公式如下:
FID = ||μ_real - μ_gen||^2 + Tr(Σ_real + Σ_gen - 2(Σ_real Σ_gen)^(1/2))
μ 是特征向量的均值,Σ 是协方差矩阵。直观理解:FID 越小,生成分布与真实分布在特征空间上越接近。
FID 的优势是计算稳定,与人的主观感知相关性较强。但它也有盲区:它无法判断类别层面的语义差异。比如生成器生成了一张完全不同的物体,但颜色、纹理接近真实数据,FID 可能也给出不错的分数。所以 FID 是必要参考,但不是唯一标准。
实际操作中,我们一般会收集 5 万或至少 1 万张真实图片和生成图片,分别计算特征统计量,再取平均值。代码实现通常用现成库:
from torchmetrics.image.fid import FrechetInceptionDistance fid = FrechetInceptionDistance(feature=2048) fid.update(real_images, real=True) fid.update(fake_images, real=False) print(fid.compute().item())使用 torchmetrics 时,传入图片的取值范围要与预训练模型兼容,通常归一化到 [0,1] 后按模型要求再处理。
6.2 视觉质量与数值指标的权衡:不能只信肉眼
虽然 FID 是公认的评估指标,但我的经验是:FID 对数据量、图像尺寸、Inception 输入规范等细节敏感,不同配置下的结果可能不可直接比较。要和社区其他实验对比时,务必统一评估数据规模和预处理方式。
另一个常用指标是 Inception Score,IS:它对单张图片计算类别分布的条件熵,同时衡量类别多样性和单张图像的清晰度。IS 越高,图像越多样且越清晰。但 IS 对数据类别均匀度敏感,如果数据集类别本身不均衡,IS 会很有误导性。
在真实项目里,我习惯把 FID、IS 和自建的领域规则一起看。人脸生成可以配合人脸检测器评估是否出现畸形;车辆设计图可以检查对称性和轮廓规范性。领域规则才是最终落地的标尺,通用指标只是帮你快速筛选实验候选。
6.3 三个常见失败模式以及诊断和修复方法
第一个失败模式:生成图像模糊、颜色单一。这通常意味着生成器和判别器都欠拟合,或生成器网络容量太小,学不到足够细节。先检查训练轮数是否太少,再尝试增加模型容量,比如层数和通道数,再调高生成器学习率。如果还不行,看看是不是数据增强过度,目标变得太复杂。
第二个失败模式:生成图像清晰但多样性差。这是很典型的模式崩塌:画面清晰,FID 却高得离谱。修复思路包括:回退到 WGAN-GP 或谱归一化架构、降低判别器学习率、增大噪声向量维度、引入 minibatch discrimination 或 style mixing 等机制。这些方法都能强制生成器不要只锁定到少数几个模式上。
第三个失败模式:训练到中途 loss 突然发散,生成图变成“雪花噪点”。多数情况是某一轮生成器更新跨过了稳定性边界,常见诱因是学习率太高、梯度惩罚系数太小,或生成器判别器更新频率失衡。建议先降低学习率,再把梯度惩罚系数从默认 10 上调到 20 观察,同时把 n_critic 改成 3 或 5,让判别器训练更充分,稳定 Wasserstein 距离估计。
在实际项目里,我还会周期性备份模型权重,并用独立验证集计算测试 FID。如果 FID 连续上百个 epoch 没有下降,说明模型已经卡住,需要果断改变优化策略,而不是继续无脑跑下去。这种“节流止损”意识在 GAN 项目里很重要,因为训练时间很贵。
最后再分享一个我自己坚持的做法:动手训练前,把训练脚本里所有关键超参集中到一个配置文件里,每次只改一个变量做实验,记录每个实验的 FID 和典型样本图。经过十到二十次实验,你会对当前数据集和模型的“脾气”有非常明确的感知。GAN 调优更像是一场持续的博弈,而不是一锤子买卖,耐心和数据管理能力往往比模型本身更值钱。第一次做实验记得把规模控制小一点,在 64×64 或 128×128 的简化数据集上先把训练稳定性搞定,再上完整数据和高分辨率。很多人揪着高分辨率反复折磨,其实只是把低分辨率就能解决的问题复杂化了。