共计 1534 个字符,预计需要花费 4 分钟才能阅读完成。
GAN 基本原理介绍
生成对抗网络 (GAN) 由生成器 (Generator) 和判别器 (Discriminator) 组成,它们就像两个互相博弈的对手:

- 生成器负责伪造数据,目标是产生足够逼真的假样本
- 判别器则是鉴伪专家,需要区分输入数据是真实样本还是生成器伪造的
它们的对抗训练过程可以类比假币制造者与警察的较量:
- 生成器不断改进假币制作工艺
- 判别器持续升级验钞技术
- 最终达到纳什均衡时,生成器能产生以假乱真的样本
常见问题分析
初学者常遇到这些典型问题:
- 模式崩溃(Mode Collapse):生成器只学会产生有限几种样本,比如生成人脸时只有几个固定表情
- 训练不稳定:损失函数剧烈波动,难以收敛
- 梯度消失:判别器过早变得太强,导致生成器无法获得有效梯度
DCGAN 实现详解
下面用 PyTorch 实现深度卷积 GAN(DCGAN),这是 GAN 的经典改进版本:
import torch
import torch.nn as nn
# 生成器网络结构
class Generator(nn.Module):
def __init__(self, latent_dim=100):
super().__init__()
self.main = 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),
# 最终输出 3 通道 RGB 图像
nn.ConvTranspose2d(128, 3, 4, 2, 1, bias=False),
nn.Tanh() # 输出值归一化到[-1,1]
)
# 判别器网络结构
class Discriminator(nn.Module):
def __init__(self):
super().__init__()
self.main = nn.Sequential(
# 输入 3 通道图像
nn.Conv2d(3, 128, 4, 2, 1, bias=False),
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, 512, 4, 2, 1, bias=False),
nn.BatchNorm2d(512),
nn.LeakyReLU(0.2, inplace=True),
# 最终输出一个概率值
nn.Conv2d(512, 1, 4, 1, 0, bias=False),
nn.Sigmoid())
训练结果展示与分析
训练过程中建议监控这些指标:
- 生成器和判别器的损失值变化曲线
- 定期保存生成的样本图像
- 使用 Inception Score 等量化指标
调参技巧与避坑指南
通过实践总结的实用建议:
- 学习率设置:通常使用较小的学习率(如 0.0002)
- 批量大小:不宜过大,64-128 是常用范围
- 标签平滑:真实样本标签用 0.9 代替 1.0
- 交替训练:适当调整生成器和判别器的训练频次比例
总结与展望
掌握 GAN 的基础实现后,可以尝试这些进阶方向:
- 条件 GAN:控制生成样本的特定属性
- 风格迁移:将图像转换为不同艺术风格
- 超分辨率重建:提升图像分辨率
GAN 正在重塑内容生成领域,期待你创造出有趣的应用!
正文完
