共计 1713 个字符,预计需要花费 5 分钟才能阅读完成。
传统生成模型对比
生成对抗网络 (GAN) 与变分自编码器(VAE)、自回归模型相比具有显著差异:

- 潜在空间连续性:GAN 的生成器 $G$ 直接学习从随机噪声 $z$ 到数据空间的映射,无需像 VAE 那样受限于变分下界的约束,能生成更清晰的样本
- 对抗训练机制:通过判别器 $D$ 提供的梯度信号,GAN 能捕捉数据分布的细微特征,而自回归模型依赖严格的序列生成顺序
- 生成质量:在图像生成任务中,GAN 生成的样本通常比 VAE 更锐利,避免了后者的模糊效应
核心数学原理
GAN 的优化目标可表示为 minimax 博弈:
$$
\min_G \max_D V(D,G) = \mathbb{E}{x\sim p[\log(1-D(G(z)))]
$$}}[\log D(x)] + \mathbb{E}_{z\sim p_z
- 生成器 $G$:试图生成逼真样本欺骗判别器,目标是最大化 $D(G(z))$
- 判别器 $D$:作为二分类器,目标是准确区分真实样本 $x$ 和生成样本 $G(z)$
DCGAN 实现详解
import torch
import torch.nn as nn
# 生成器网络结构
class Generator(nn.Module):
def __init__(self, latent_dim=100, img_channels=3):
super().__init__()
self.main = nn.Sequential(
# 输入: latent_dim 维噪声
nn.ConvTranspose2d(latent_dim, 512, 4, 1, 0, bias=False),
nn.BatchNorm2d(512),
nn.ReLU(True),
# 输出尺寸: (512,4,4)
nn.ConvTranspose2d(512, 256, 4, 2, 1, bias=False),
nn.BatchNorm2d(256),
nn.ReLU(True),
# 输出尺寸: (256,8,8)
nn.ConvTranspose2d(256, 128, 4, 2, 1, bias=False),
nn.BatchNorm2d(128),
nn.ReLU(True),
# 输出尺寸: (128,16,16)
nn.ConvTranspose2d(128, img_channels, 4, 2, 1, bias=False),
nn.Tanh() # 输出像素值归一化到[-1,1]
)
def forward(self, z):
return self.main(z)
关键实现细节:
- 转置卷积 :通过
ConvTranspose2d实现上采样 - BatchNorm:稳定训练过程,但需注意判别器中不宜使用
- Tanh 激活:将生成图像像素值约束到合理范围
训练过程关键技巧
模式崩溃应对
- 特征表现:生成器只产生有限几种样本模式
- 解决方案:
- 采用 Mini-batch 判别
- 添加多样性正则项
- 使用 Wasserstein GAN 缓解梯度消失
超参数设置
- 学习率:通常设置为 2e-4(Adam 优化器)
- 批量大小:建议 64-256 之间
- 标签平滑:防止判别器过度自信
评估指标实现
Fréchet Inception Distance (FID)计算流程:
- 提取真实图像和生成图像的 Inception-v3 特征
- 计算两个多元高斯分布之间的 Wasserstein- 2 距离
- 分数越低表示生成质量越好
from torchmetrics.image.fid import FrechetInceptionDistance
fid = FrechetInceptionDistance(feature=2048)
# 更新真实图像特征
fid.update(real_images, real=True)
# 更新生成图像特征
fid.update(fake_images, real=False)
print(f"FID score: {fid.compute():.2f}")
开放性问题讨论
- 多样性评估:
- 计算生成样本的最近邻距离
- 使用分类器置信度分布
-
可视化潜在空间插值
-
医疗伦理边界:
- 生成病理图像可能误导诊断
- 需建立数据来源审查机制
- 生成样本应明确标注人工合成属性
工程实践建议
- 使用梯度裁剪(
clip_grad_norm_)稳定训练 - 定期保存模型检查点
- 监控损失函数和生成样本的演化过程
- 尝试 ProGAN 渐进式训练策略提升高分辨率生成质量
正文完
发表至: 未分类
四天前
