Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

GAN (生成对抗网络) Demo 详解

项目概述

本项目实现了一个完整的生成对抗网络(GAN)模型,用于生成手写数字图像。通过详细的代码注释和可视化,帮助理解GAN的工作原理和训练过程。

文件结构

Gan-demo/
├── gan_demo.py      # 主要的GAN实现代码
├── requirements.txt # 项目依赖
├── README.md       # 项目说明文档
└── data/
    └── mnist.npz   # MNIST数据集(Keras格式)

GAN 核心概念

什么是GAN?

生成对抗网络(Generative Adversarial Network)由两个神经网络组成:

  • 生成器(Generator): 从随机噪声生成假图像
  • 判别器(Discriminator): 判断图像是真实的还是生成的

训练原理

生成器和判别器进行对抗训练:

  • 生成器试图生成能够欺骗判别器的图像
  • 判别器试图正确区分真实图像和生成图像
  • 两者在训练过程中不断改进,最终生成器能够生成逼真的图像

数据加载流程

1. 数据集准备

# 从本地data/mnist.npz加载Keras格式MNIST数据集
mnist_data = np.load('data/mnist.npz')
x_train = mnist_data['x_train']  # (60000, 28, 28)

2. 数据预处理

# 归一化到[-1, 1],并增加通道维度
x_train = (x_train.astype(np.float32) / 255.0 - 0.5) / 0.5  # [0,1] -> [-1,1]
x_train = np.expand_dims(x_train, 1)  # (60000, 1, 28, 28)

3. 数据加载器

train_tensor = torch.tensor(x_train)
dataset = TensorDataset(train_tensor)
dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)

模型架构

生成器(Generator)

class Generator(nn.Module):
    def __init__(self, z_dim, img_shape):
        super(Generator, self).__init__()
        self.model = nn.Sequential(
            nn.Linear(z_dim, 128),           # 输入层:噪声 -> 128维
            nn.ReLU(True),
            nn.Linear(128, 256),             # 隐藏层1:128 -> 256维
            nn.BatchNorm1d(256),             # 批归一化
            nn.ReLU(True),
            nn.Linear(256, 512),             # 隐藏层2:256 -> 512维
            nn.BatchNorm1d(512),
            nn.ReLU(True),
            nn.Linear(512, int(np.prod(img_shape))),  # 输出层:512 -> 784维(28*28)
            nn.Tanh()                        # 激活函数,输出范围[-1, 1]
        )

工作原理

  • 输入:100维随机噪声向量
  • 通过全连接层逐步扩展维度
  • 使用BatchNorm和ReLU激活函数
  • 最终输出784维向量,重塑为28×28图像

判别器(Discriminator)

class Discriminator(nn.Module):
    def __init__(self, img_shape):
        super(Discriminator, self).__init__()
        self.model = nn.Sequential(
            nn.Linear(int(np.prod(img_shape)), 512),  # 输入层:784 -> 512维
            nn.LeakyReLU(0.2, inplace=True),
            nn.Linear(512, 256),                      # 隐藏层:512 -> 256维
            nn.LeakyReLU(0.2, inplace=True),
            nn.Linear(256, 1),                        # 输出层:256 -> 1维
            nn.Sigmoid()                              # 输出概率[0, 1]
        )

工作原理

  • 输入:28×28图像展平为784维向量
  • 通过全连接层逐步降维
  • 使用LeakyReLU激活函数
  • 最终输出一个概率值,表示图像为真实的概率

训练流程

1. 初始化

# 超参数设置
batch_size = 64
lr = 0.0002
z_dim = 100
epochs = 50

# 损失函数和优化器
criterion = nn.BCELoss()  # 二分类交叉熵损失
optimizer_G = optim.Adam(G.parameters(), lr=lr, betas=(0.5, 0.999))
optimizer_D = optim.Adam(D.parameters(), lr=lr, betas=(0.5, 0.999))

2. 训练循环

for epoch in range(1, epochs+1):
    for i, (imgs,) in enumerate(dataloader):
        real_imgs = imgs.to(device)
        batch_size_curr = real_imgs.size(0)
        
        # 标签设置
        valid = torch.ones(batch_size_curr, 1, device=device)   # 真实图像标签为1
        fake = torch.zeros(batch_size_curr, 1, device=device)   # 生成图像标签为0
        
        # ---------------------
        #  训练生成器G
        # ---------------------
        optimizer_G.zero_grad()
        z = torch.randn(batch_size_curr, z_dim, device=device)  # 生成随机噪声
        gen_imgs = G(z)                                         # 生成假图像
        g_loss = criterion(D(gen_imgs), valid)                  # 希望判别器认为生成图像为真
        g_loss.backward()
        optimizer_G.step()
        
        # ---------------------
        #  训练判别器D
        # ---------------------
        optimizer_D.zero_grad()
        real_loss = criterion(D(real_imgs), valid)              # 真实图像损失
        fake_loss = criterion(D(gen_imgs.detach()), fake)       # 生成图像损失
        d_loss = (real_loss + fake_loss) / 2                    # 总损失
        d_loss.backward()
        optimizer_D.step()

3. 训练策略详解

生成器训练目标

  • 目标:让判别器将生成的图像误判为真实图像
  • 损失函数criterion(D(gen_imgs), valid)
  • 含义:希望判别器对生成图像输出接近1的概率

判别器训练目标

  • 目标:正确区分真实图像和生成图像
  • 损失函数
    • real_loss = criterion(D(real_imgs), valid) - 真实图像应输出高概率
    • fake_loss = criterion(D(gen_imgs.detach()), fake) - 生成图像应输出低概率
  • 总损失(real_loss + fake_loss) / 2

可视化功能

1. 生成图像可视化

# 每隔10个epoch显示生成的图像
if epoch % 10 == 0 or epoch == 1:
    G.eval()
    with torch.no_grad():
        fake_imgs = G(fixed_noise).detach().cpu()  # 使用固定噪声生成图像
    # 显示5×5网格的生成图像

2. 损失曲线可视化

# 绘制生成器和判别器的损失曲线
plt.plot(G_losses, label='Generator Loss')
plt.plot(D_losses, label='Discriminator Loss')

关键超参数说明

参数 说明
batch_size 64 批次大小,影响训练稳定性和内存使用
lr 0.0002 学习率,影响模型收敛速度
z_dim 100 噪声向量维度,影响生成图像的多样性
epochs 50 训练轮数,影响模型性能

训练过程观察要点

1. 损失变化

  • 理想情况:生成器和判别器损失都应该逐渐下降
  • 判别器过强:判别器损失过低,生成器难以学习
  • 生成器过强:生成器损失过低,可能出现模式崩塌

2. 生成质量

  • 初期:生成图像为随机噪声
  • 中期:开始出现数字轮廓
  • 后期:生成清晰的数字图像

3. 训练平衡

  • 生成器和判别器需要保持平衡
  • 如果一方过强,另一方将无法有效学习
  • 可以通过调整学习率或训练频率来平衡

运行方法

  1. 安装依赖:
pip install -r requirements.txt
  1. 运行训练:
python gan_demo.py
  1. 观察输出:
  • 控制台显示训练进度和损失
  • 弹出窗口显示生成的图像
  • 训练结束后显示损失曲线

扩展建议

  1. 改进网络结构:使用卷积神经网络(CNN)替代全连接层
  2. 调整超参数:尝试不同的学习率、批次大小等
  3. 添加正则化:使用梯度惩罚、谱归一化等技术
  4. 更换数据集:尝试其他图像数据集如CIFAR-10
  5. 实现其他GAN变体:如DCGAN、WGAN、StyleGAN等

总结

本demo通过完整的GAN实现,展示了:

  • 如何从本地数据加载和预处理数据
  • 如何构建生成器和判别器网络
  • 如何实现对抗训练过程
  • 如何可视化和监控训练过程

通过运行这个demo,你可以直观地理解GAN的工作原理,观察生成图像从噪声到清晰数字的演化过程,为深入学习更复杂的生成模型打下基础。

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages