共计 2530 个字符,预计需要花费 7 分钟才能阅读完成。
传统 GAN 的三大痛点
在开始介绍 BEGAN 之前,我们先来看看传统 GAN 模型面临的三个主要问题:

-
模式崩溃(Mode Collapse):生成器倾向于生成有限几种样本,无法覆盖全部数据分布。比如在生成数字时,可能只生成 ”1″ 和 ”7″ 而忽略其他数字。
-
梯度不稳定:判别器训练得太好会导致生成器梯度消失,而判别器训练不足又会导致生成器收到无意义的梯度信号。
-
评估指标不可靠:传统的 GAN 缺乏可靠的量化评估指标,难以客观比较不同模型的性能。
BEGAN 核心原理
BEGAN(Boundary Equilibrium Generative Adversarial Networks)通过引入均衡概念和 Wasserstein 距离的改进,有效解决了上述问题。
边界平衡理论
BEGAN 的关键在于维持生成器和判别器之间的平衡。其目标函数可以表示为:
$$\mathcal{L}D = \mathcal{L}(x) – k_t \mathcal{L}(G(z))$$
$$\mathcal{L}_G = \mathcal{L}(G(z))$$
$$k(G(z)))$$} = k_t + \lambda(\gamma\mathcal{L}(x) – \mathcal{L
其中:
– $\mathcal{L}$ 是自编码器的重构损失
– $k_t$ 是控制平衡的比例因子
– $\gamma$ 是目标多样性比例(通常设为 0.5)
– $\lambda$ 是学习率(通常设为 0.001)
与传统 GAN/WGAN 对比
| 特性 | 传统 GAN | WGAN | BEGAN |
|---|---|---|---|
| 损失函数 | JS 散度 | Wasserstein 距离 | 自编码器损失 |
| 训练稳定性 | 低 | 中 | 高 |
| 模式崩溃 | 严重 | 较轻 | 很轻 |
| 评估指标 | 不可靠 | 较可靠 | 可靠 |
PyTorch 实现
以下是 BEGAN 的核心实现代码:
import torch
import torch.nn as nn
import torch.optim as optim
class Generator(nn.Module):
def __init__(self, z_dim=64, hidden_dim=128):
super().__init__()
self.net = nn.Sequential(nn.Linear(z_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, 784),
nn.Tanh() # 输出在 [-1,1] 之间
)
def forward(self, z):
return self.net(z)
class Discriminator(nn.Module):
def __init__(self, x_dim=784, hidden_dim=128):
super().__init__()
self.encoder = nn.Sequential(nn.Linear(x_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim),
)
self.decoder = nn.Sequential(nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, x_dim),
)
def forward(self, x):
h = self.encoder(x)
return self.decoder(h)
# 初始化模型和优化器
G = Generator().to(device)
D = Discriminator().to(device)
opt_G = optim.Adam(G.parameters(), lr=0.0001)
opt_D = optim.Adam(D.parameters(), lr=0.0001)
# 训练循环
k = 0.0 # 初始平衡系数
gamma = 0.5 # 多样性比例
lambda_k = 0.001 # k 的学习率
for epoch in range(epochs):
for real_data, _ in dataloader:
# 训练判别器
z = torch.randn(batch_size, z_dim).to(device)
fake_data = G(z)
D_real = D(real_data)
D_fake = D(fake_data.detach())
loss_real = torch.mean(torch.abs(D_real - real_data))
loss_fake = torch.mean(torch.abs(D_fake - fake_data))
loss_D = loss_real - k * loss_fake
opt_D.zero_grad()
loss_D.backward()
opt_D.step()
# 训练生成器
fake_data = G(z)
D_fake = D(fake_data)
loss_G = torch.mean(torch.abs(D_fake - fake_data))
opt_G.zero_grad()
loss_G.backward()
opt_G.step()
# 更新平衡系数 k
diff = gamma * loss_real - loss_fake
k = k + lambda_k * diff
k = min(max(k, 0.0), 1.0) # 限制在 [0,1] 范围内
实战调优指南
超参数设置
-
学习率:建议从 0.0001 开始尝试,可设置在 0.00005 到 0.0002 之间
-
batch size:通常 64-256 效果较好,太大会降低生成多样性
-
平衡系数 γ :控制生成多样性与质量的权衡,一般设为 0.5
训练监控
-
损失曲线:应同时监控判别器和生成器的损失,理想情况下两者应该保持动态平衡
-
样本可视化:每隔一定 epoch 保存生成样本,观察质量变化
-
平衡系数 k :应保持在 [0,1] 范围内波动,如果持续增大或减小说明训练失衡
开放性思考
- BEGAN 在非图像领域的应用:
- 如何修改网络结构以适应时间序列数据生成?
-
文本生成任务中如何设计合适的重构损失函数?
-
生成器性能突降排查:
- 首先检查平衡系数 k 是否超出合理范围
- 检查学习率是否设置过高导致训练不稳定
- 确认输入噪声 z 的分布是否发生变化
结语
BEGAN 通过引入边界平衡机制,显著提高了 GAN 训练的稳定性。在实践中,我发现在人脸生成任务上,BEGAN 相比传统 GAN 能更快收敛,且生成的样本多样性更好。希望这篇笔记能帮助你快速上手 BEGAN,避开 GAN 训练中的常见陷阱。
