先说我踩的坑
一个月前接了个工业视觉项目:检测注塑件表面缺陷(划痕、麻点、缺料)。产线收集了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,方便你直接跑通),硬件环境:
| GPU | NVIDIA RTX 3080 10GB |
| CPU | Intel i7-12700K |
| PyTorch | 2.1.2 + CUDA 12.1 |
| Python | 3.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。
核心改动三个:
- 去掉判别器的sigmoid,输出不做任何压缩
- loss不用log,直接用拼接打分
- 梯度惩罚项λ·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数:
| 指标 | DCGAN | WGAN-GP |
| FID (100 epoch) | 89.2 | 32.7 |
| 训练时间 | 28min | 35min |
| 模式坍塌 | 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张就开训。