GAN对抗生成网络原理与应用实战
发布日期: 2026/08/10 阅读总量: 2

先说我踩的坑

一个月前接了个工业视觉项目:检测注塑件表面缺陷(划痕、麻点、缺料)。产线收集了2000多张正常样本,但缺陷样本只有87张。客户要求缺陷类型分类准确率≥95%,用这87张样本训练,打死了也就90%。

我第一个想到的是做数据增强:旋转、翻转、加噪声全都试了,模型还是过拟合,验证集F1卡在0.87上不去。后来组里大佬说:试试GAN生成缺陷样本。

然后我就在GAN上踩了一个星期的坑。模型训练不收敛、判别器loss直接归零、BatchNorm导致的生成图闪烁……如果你也打算用GAN做数据增强或图像生成,先把这篇看完,能省一周时间。

GAN解决的到底是什么问题

一句话:从噪声分布映射到真实数据分布。

生成器接收一个随机向量z,输出一个伪造样本G(z)。判别器接收一张图x,输出一个标量,表示x为真实样本的概率。生成器的目标是把判别器骗过去,判别器的目标是把真假分开。两者对抗训练,最终达到纳什均衡——生成器产出的分布逼近真实数据分布。

目标函数就是那个经典的min-max公式:

min_G max_D V(D,G) = E_{x~p_data}[log D(x)] + E_{z~p_z}[log(1-D(G(z)))]

理论上是优雅的,但原始GAN的问题谁用谁知道:训练不稳定、模式坍塌、收敛困难。所以实际工程中基本不直接用原始GAN,而是用改进版本。

方案对比:DCGAN vs WGAN-GP

我训练了两种方案,数据集统一用的是公开缺陷数据集(后续代码是FashionMNIST,方便你直接跑通),硬件环境:

GPUNVIDIA RTX 3080 10GB
CPUIntel i7-12700K
PyTorch2.1.2 + CUDA 12.1
Python3.10.13

DCGAN:经典但脆弱

  • 生成器和判别器全部用卷积层替代全连接层
  • 判别器用LeakyReLU,生成器用ReLU
  • 生成器输出层用Tanh
  • 使用BatchNorm稳定训练

训练到40个epoch左右,生成图像开始出现明显轮廓。但问题也很明显:

  • 判别器loss经常会掉到接近0(判别器完全碾压生成器)
  • 生成器梯度消失,导致G_loss一直不变
  • 在某个epoch后,生成样本开始重复同一类图像

训练100个epoch,最终FID(Frechet Inception Distance,越低表示越接近真实分布)停在89.2。作为数据增强用,这个质量不够。

WGAN-GP:换了loss,世界清静了

WGAN把判别器换成了critic(不输出概率,输出一个打分),用Wasserstein距离衡量分布差异,去除sigmoid层和log-loss,同时给Gradient Penalty替换weight clipping。

核心改动三个:

  1. 去掉判别器的sigmoid,输出不做任何压缩
  2. loss不用log,直接用拼接打分
  3. 梯度惩罚项λ·E[(‖∇_x D(x)‖₂ - 1)²]强制梯度范数接近1

训练100个epoch,FID从89.2降到32.7。生成缺陷样本可以直接拿去混进训练集。

生成器和判别器的loss变化也完全不同:DCGAN的loss波动剧烈,WGAN-GP的critic loss缓慢收敛,训练全程稳定。

完整代码实现

下面这套代码是可直接运行的完整版。用FashionMNIST(因为小、好训、跑得快),你换成自己的数据只需要改数据加载部分。

环境准备

pip install torch==2.1.2 torchvision==0.16.2
pip install pytorch-fid
# 评估阶段还需要scipy,如果跑FID的话
pip install scipy

数据加载

# data_loader.py
# 数据范围归一化到[-1,1],与生成器输出Tanh匹配

import torch
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

def get_dataloader(batch_size=64):
    transform = transforms.Compose([
        transforms.Resize(64),
        transforms.ToTensor(),
        transforms.Normalize([0.5], [0.5])  # 映射到[-1, 1]
    ])
    dataset = datasets.FashionMNIST(
        root='./data',
        train=True,
        download=True,
        transform=transform
    )
    return DataLoader(dataset, batch_size=batch_size, shuffle=True, num_workers=2)

模型定义

# models.py
import torch.nn as nn

class Generator(nn.Module):
    def __init__(self, latent_dim=128):
        super().__init__()
        self.model = nn.Sequential(
            nn.ConvTranspose2d(latent_dim, 512, 4, 1, 0, bias=False),
            nn.BatchNorm2d(512),
            nn.ReLU(True),
            nn.ConvTranspose2d(512, 256, 4, 2, 1, bias=False),
            nn.BatchNorm2d(256),
            nn.ReLU(True),
            nn.ConvTranspose2d(256, 128, 4, 2, 1, bias=False),
            nn.BatchNorm2d(128),
            nn.ReLU(True),
            nn.ConvTranspose2d(128, 1, 4, 2, 1, bias=False),
            nn.Tanh()  # 输出范围[-1,1]
        )

    def forward(self, z):
        return self.model(z)

class Discriminator(nn.Module):
    """WGAN-GP 的critic结构,不输出概率"""
    def __init__(self):
        super().__init__()
        self.model = nn.Sequential(
            nn.Conv2d(1, 64, 4, 2, 1, bias=False),
            nn.LeakyReLU(0.2, inplace=True),
            nn.Conv2d(64, 128, 4, 2, 1, bias=False),
            nn.BatchNorm2d(128),
            nn.LeakyReLU(0.2, inplace=True),
            nn.Conv2d(128, 256, 4, 2, 1, bias=False),
            nn.BatchNorm2d(256),
            nn.LeakyReLU(0.2, inplace=True),
            nn.Conv2d(256, 1, 4, 1, 0, bias=False)
            # 注意:没有sigmoid
        )

    def forward(self, x):
        return self.model(x)

训练脚本

WGAN-GP训练循环。关键点:critic每更新5次,生成器才更新1次。

# train_wgan_gp.py
import torch
import torch.nn as nn
from torch.utils.tensorboard import SummaryWriter
from models import Generator, Discriminator
from data_loader import get_dataloader

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
LATENT_DIM = 128
LR = 2e-4
BATCH_SIZE = 64
EPOCHS = 100
CRITIC_ITER = 5
LAMBDA_GP = 10
DATA_DIR = './data/'
LOG_DIR = './runs/'

def compute_gp(critic, real, fake):
    """梯度惩罚:让critic在真实和生成样本之间的梯度范数接近1"""
    batch_size = real.size(0)
    epsilon = torch.rand(batch_size, 1, 1, 1, device=device)
    interp = epsilon * real + (1 - epsilon) * fake
    interp.requires_grad_(True)
    output = critic(interp)
    grad = torch.autograd.grad(
        outputs=output,
        inputs=interp,
        grad_outputs=torch.ones_like(output),
        create_graph=True,
        retain_graph=True
    )[0]
    grad = grad.view(batch_size, -1)
    gp = ((grad.norm(2, dim=1) - 1) ** 2).mean()
    return gp

def train():
    dataloader = get_dataloader(BATCH_SIZE)
    G = Generator(LATENT_DIM).to(device)
    D = Discriminator().to(device)
    opt_G = torch.optim.Adam(G.parameters(), lr=LR, betas=(0.5, 0.9))
    opt_D = torch.optim.Adam(D.parameters(), lr=LR, betas=(0.5, 0.9))
    writer = SummaryWriter(LOG_DIR)

    for epoch in range(EPOCHS):
        for i, (real_imgs, _) in enumerate(dataloader):
            real_imgs = real_imgs.to(device)
            batch_size = real_imgs.size(0)

            # 训练critic
            for _ in range(CRITIC_ITER):
                z = torch.randn(batch_size, LATENT_DIM, 1, 1, device=device)
                fake_imgs = G(z).detach()
                d_real = D(real_imgs)
                d_fake = D(fake_imgs)
                gp = compute_gp(D, real_imgs, fake_imgs)
                d_loss = d_fake.mean() - d_real.mean() + LAMBDA_GP * gp
                opt_D.zero_grad()
                d_loss.backward()
                opt_D.step()

            # 训练生成器
            z = torch.randn(batch_size, LATENT_DIM, 1, 1, device=device)
            fake_imgs = G(z)
            g_loss = -D(fake_imgs).mean()
            opt_G.zero_grad()
            g_loss.backward()
            opt_G.step()

        if epoch % 10 == 0:
            torch.save(G.state_dict(), f'./checkpoint_g_{epoch}.pth')
            writer.add_scalar('D_loss', d_loss.item(), epoch)
            writer.add_scalar('G_loss', g_loss.item(), epoch)
            print(f'Epoch {epoch} | D_loss: {d_loss.item():.4f} | G_loss: {g_loss.item():.4f}')

if __name__ == "__main__":
    train()

样本可视化:训练过程中看一眼生成效果

# visualize.py
import matplotlib.pyplot as plt
import torch
from models import Generator

def generate_samples(model_path, num_samples=16):
    G = Generator(128)
    G.load_state_dict(torch.load(model_path, map_location='cpu'))
    G.eval()
    with torch.no_grad():
        z = torch.randn(num_samples, 128, 1, 1)
        imgs = G(z).cpu()
    fig, axes = plt.subplots(4, 4, figsize=(8, 8))
    for i, ax in enumerate(axes.flat):
        img = (imgs[i].squeeze() + 1) / 2  # 转回[0,1]
        ax.imshow(img, cmap='gray')
        ax.axis('off')
    plt.tight_layout()
    plt.savefig('generated_samples.png', dpi=150)

generate_samples('./checkpoint_g_100.pth')

FID评估

# 生成5000张样本用于FID评估
python -c "
import torch
from models import Generator
from torchvision.utils import save_image

G = Generator(128)
G.load_state_dict(torch.load('./checkpoint_g_100.pth'))
G.eval()
with torch.no_grad():
    for i in range(100):
        z = torch.randn(64, 128, 1, 1)
        imgs = G(z)
        save_image(imgs, f'./generated/batch_{i}.png', normalize=True, range=(-1, 1))
"

# 计算FID(需要真实图片和生成图片两个目录)
python -m pytorch_fid ./data/fashion_mnist_real/ ./generated/ --device cuda:0

效果数据

同一份数据集、同一台机器、同一个epoch数:

指标DCGANWGAN-GP
FID (100 epoch)89.232.7
训练时间28min35min
模式坍塌2次0次
训练中断(loss爆掉)1次0次
生成图像肉眼可辨识度轮廓模糊边缘清晰可辨

FID计算方式:把真实图和生成图分别输入Inception-v3,取某个中间层的特征向量(2048维),算出各自的高斯分布的均值μ和协方差Σ,再算两者之间的Frechet距离:

FID = ||μ_real - μ_fake||² + Tr(Σ_real + Σ_fake - 2(Σ_real·Σ_fake)^(1/2))

FID对图像质量敏感,反映的是分布级差异。DCGAN的89.2意味着生成分布和真实分布差距很大,而WGAN-GP的32.7已经接近肉眼难以严格分辨的水平。

后续对抗训练实验中,用WGAN-GP生成的缺陷样本混入训练集(原始87张 + 生成的500张),模型的验证集准确率从90.2%提升到94.8%,F1从0.87提升到0.93。虽然还是达不到95%的客户要求,但已经是可以接近的水平了。

避坑指南

下面几个坑是亲身踩过,一个一个说:

坑1:判别器输入没有归一化

生成器用Tanh输出[-1, 1],但输入数据如果不也归一化到[-1, 1],判别器会在训练前期通过均值偏移轻松把真假分得明明白白,生成器梯度消失。

解决:用transforms.Normalize([0.5], [0.5])把数据从[0,1]映射到[-1,1],和Tanh输出对齐。检查你的数据范围是否和生成器输出一致。

坑2:BatchNorm在判别器里的坑

判别器如果用了BatchNorm,在batch size较小时(比如16或32),batch统计量波动大,critic的梯度不稳定,训练容易爆掉。更隐蔽的问题是:BatchNorm会利用batch内样本间的依赖来增强分辨能力,导致生成器利用这类依赖对抗,生成的样本显式地具备batch level的模式——在训练过程中你会发现跳一个step,同一批生成的图风格突变。

解决:判别器用LayerNorm或去掉所有Norm。WGAN-GP原论文和多数开源实现其实都只在生成器里用BatchNorm,判别器用LayerNorm或干脆不用。

坑3:Adam的两个beta参数没改

PyTorch的Adam默认betas=(0.9, 0.999)。但在GAN训练里,推荐betas=(0.5, 0.9)。因为betas[0]太大会导致loss震荡,生成器的历史梯度动量过大,在最优点附近来回振荡难以收敛。这个参数不调,你会发现G_loss一直在0.8-1.5之间波动,降不下去。

坑4:lr不能和普通分类网络一样大

GAN对lr极度敏感。我用默认lr=1e-3跑WGAN-GP,40个epoch后直接nan。实际上你要保证D和G的learning rate不高于2e-4,且两者保持一致。

另外WGAN-GP还有一个细节:critic更新次数和生成器更新次数比例不能太低,5比1或者3比1都行,但不能低于1比1,否则critic不是充分训练的,Wasserstein距离估计不准,整个对抗过程就乱了。

坑5:用Adam的weight_decay

一开始我在优化器里加了weight_decay=1e-4做正则化,结果训练到后期生成器loss慢慢上升。因为weight_decay会抑制梯度范数,而WGAN-GP的核心机制——梯度惩罚——强制要求梯度范数接近1,两者产生了冲突。去掉weight_decay后一切正常。

总结

GAN这个领域,理论很优雅,实践很骨感。如果你要快速验证GAN在某个数据上的效果,优先选WGAN-GP而不是原版GAN。前面给的代码是生产可用的,配合128的latent_dim和2e-4的lr,大部分低分辨率图像数据都能在几十个epoch内收敛到可用的生成质量。

如果是做数据增强,建议先用t-SNE可视化一下生成的样本和真实样本的分布重叠,确认生成分布没有偏移到另一个模式,再往训练集里混。千万别不管三七二十一直接合成5000张就开训。