GAN入门实战:从零开始构建你的第一个生成对抗网络

1次阅读
没有评论

共计 2811 个字符,预计需要花费 8 分钟才能阅读完成。

image.webp

生成对抗网络(Generative Adversarial Networks,GAN)在图像生成领域展现了惊人的能力,从风格迁移到超分辨率重建,再到如今火热的 AI 绘画,背后都有 GAN 的身影。对于初学者来说,理解 GAN 的核心思想并不复杂——它通过让两个神经网络相互博弈来学习数据分布。今天我们就用 PyTorch 来实现一个最简单的 GAN 模型,生成 MNIST 手写数字。

GAN 入门实战:从零开始构建你的第一个生成对抗网络

一、GAN 的核心原理

GAN 由生成器(Generator)和判别器(Discriminator)组成。生成器负责生成假数据,判别器则要判断输入是真实数据还是生成器造的假数据。两者就像警察和小偷,在对抗中不断进步。

  1. 对抗训练机制
  2. 生成器 G 试图生成足以乱真的假数据
  3. 判别器 D 努力区分真假数据
  4. 两者通过反向传播交替优化

  5. 数学表达 :训练过程可以表示为最小最大博弈:
    $$\min_G \max_D V(D,G) = \mathbb{E}{x\sim p[\log(1-D(G(z)))]$$}(x)}[\log D(x)] + \mathbb{E}_{z\sim p_z(z)

  6. 损失函数 :实际使用时会对生成器采用改进的交叉熵损失:

  7. 原始公式可能导致梯度消失
  8. 常见变体是将生成器目标改为最大化 D(G(z)) 而非最小化 1 -D(G(z))

二、PyTorch 代码实现

以下是一个完整的 MNIST 生成示例(Python 3.8+,PyTorch 1.10+):

import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

# 1. 数据准备
transform = transforms.Compose([transforms.ToTensor(),
    transforms.Normalize((0.5,), (0.5,))  # MNIST 单通道
])
train_set = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
train_loader = DataLoader(train_set, batch_size=64, shuffle=True)

# 2. 网络定义
class Generator(nn.Module):
    def __init__(self):
        super().__init__()
        self.main = nn.Sequential(nn.Linear(100, 256),  # 输入噪声维度 100
            nn.LeakyReLU(0.2),
            nn.Linear(256, 512),
            nn.LeakyReLU(0.2),
            nn.Linear(512, 784),  # MNIST 28x28=784
            nn.Tanh()  # 输出归一化到 [-1,1]
        )

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

class Discriminator(nn.Module):
    def __init__(self):
        super().__init__()
        self.main = nn.Sequential(nn.Linear(784, 512),
            nn.LeakyReLU(0.2),
            nn.Linear(512, 256),
            nn.LeakyReLU(0.2),
            nn.Linear(256, 1),
            nn.Sigmoid()  # 输出概率值)

    def forward(self, x):
        x = x.view(x.size(0), -1)  # 展平
        return self.main(x)

# 3. 训练循环
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
G = Generator().to(device)
D = Discriminator().to(device)

criterion = nn.BCELoss()
optim_G = optim.Adam(G.parameters(), lr=0.0002, betas=(0.5, 0.999))
optim_D = optim.Adam(D.parameters(), lr=0.0002, betas=(0.5, 0.999))

for epoch in range(50):
    for i, (real_imgs, _) in enumerate(train_loader):
        # 真实数据
        real_imgs = real_imgs.to(device)
        real_labels = torch.ones(real_imgs.size(0), 1).to(device)

        # 生成假数据
        noise = torch.randn(real_imgs.size(0), 100).to(device)
        fake_imgs = G(noise)
        fake_labels = torch.zeros(real_imgs.size(0), 1).to(device)

        # 训练判别器
        optim_D.zero_grad()
        real_loss = criterion(D(real_imgs), real_labels)
        fake_loss = criterion(D(fake_imgs.detach()), fake_labels)
        d_loss = real_loss + fake_loss
        d_loss.backward()
        optim_D.step()

        # 训练生成器
        optim_G.zero_grad()
        g_loss = criterion(D(fake_imgs), real_labels)  # 让生成的图片被判别为真
        g_loss.backward()
        optim_G.step()

关键点说明:
– 第 18、32 行:使用 LeakyReLU 防止梯度消失
– 第 45 行:detach() 切断生成器梯度传播
– 第 53 行:生成器目标是让 D(G(z)) 接近 1

三、实战避坑指南

  1. 模式崩溃(Mode Collapse)
  2. 现象:生成器只产出有限几种样本
  3. 解决方案:

    • 增加噪声维度(如从 100 调整到 256)
    • 尝试 Wasserstein GAN 等改进结构
  4. 参数调整

  5. 学习率通常在 0.0001-0.0005 之间
  6. 批大小建议 64/128,过大可能导致训练不稳定

  7. 可视化监控

  8. 每 100 次迭代保存生成样本
  9. 使用 TensorBoard 记录损失曲线
  10. 观察判别器准确率应保持在 50-60% 左右

四、延伸思考

  1. 判别器过强的表现是训练早期准确率就接近 100%,这时可以:
  2. 降低判别器学习率
  3. 减少判别器层数
  4. 添加 Dropout 层

  5. DCGAN 通过以下改进提升生成质量:

  6. 使用转置卷积替代全连接
  7. 引入 Batch Normalization
  8. 移除池化层改用步长卷积

通过这个基础实现,你应该已经感受到了 GAN 的奇妙之处。建议尝试修改网络结构或调整超参数,观察对生成结果的影响。当你能稳定生成清晰的手写数字时,就迈出了掌握生成模型的第一步!

正文完
 0
评论(没有评论)