GAN实战:用PyTorch在MNIST数据集上生成逼真手写数字
GAN实战用PyTorch在MNIST数据集上生成逼真手写数字当Ian Goodfellow在2014年首次提出生成对抗网络(GAN)时他可能没想到这个框架会彻底改变计算机视觉领域。今天我们将用PyTorch这个强大的深度学习框架在经典的MNIST数据集上实现一个能够生成逼真手写数字的GAN模型。这不仅是一次技术实践更是一场创造力的展示——教会机器如何像人类一样书写数字。1. 环境准备与数据加载在开始之前确保你已经安装了PyTorch和相关的依赖库。如果你有NVIDIA GPU强烈建议安装CUDA版本的PyTorch以加速训练过程。pip install torch torchvision matplotlib numpyMNIST数据集包含60,000张28x28像素的手写数字灰度图像每张图像都标注了对应的数字(0-9)。让我们先加载并预处理这些数据import torch import torchvision from torchvision import transforms # 定义数据转换 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean(0.5,), std(0.5,)) # 将像素值从[0,1]归一化到[-1,1] ]) # 加载训练集 train_dataset torchvision.datasets.MNIST( root./data, trainTrue, downloadTrue, transformtransform ) # 创建数据加载器 batch_size 128 train_loader torch.utils.data.DataLoader( datasettrain_dataset, batch_sizebatch_size, shuffleTrue )提示如果下载速度慢可以手动下载MNIST数据集并放在指定目录下避免重复下载。2. GAN模型架构设计GAN由两个相互对抗的神经网络组成生成器(Generator)和判别器(Discriminator)。生成器的任务是创造逼真的假图像而判别器则需要区分真实图像和生成器产生的假图像。2.1 生成器网络生成器接收一个随机噪声向量(通常来自正态分布)作为输入输出一张伪造的手写数字图像。我们使用全连接层构建生成器import torch.nn as nn class Generator(nn.Module): def __init__(self, latent_dim100, img_shape(28, 28)): super(Generator, self).__init__() self.img_shape img_shape self.img_size img_shape[0] * img_shape[1] self.model nn.Sequential( nn.Linear(latent_dim, 256), nn.LeakyReLU(0.2), nn.Linear(256, 512), nn.LeakyReLU(0.2), nn.Linear(512, 1024), nn.LeakyReLU(0.2), nn.Linear(1024, self.img_size), nn.Tanh() # 输出值在[-1,1]之间 ) def forward(self, z): img self.model(z) img img.view(img.size(0), *self.img_shape) return img2.2 判别器网络判别器接收一张图像(真实或生成的)作为输入输出一个标量表示图像为真的概率class Discriminator(nn.Module): def __init__(self, img_shape(28, 28)): super(Discriminator, self).__init__() self.img_size img_shape[0] * img_shape[1] self.model nn.Sequential( nn.Linear(self.img_size, 1024), nn.LeakyReLU(0.2), nn.Dropout(0.3), nn.Linear(1024, 512), nn.LeakyReLU(0.2), nn.Dropout(0.3), nn.Linear(512, 256), nn.LeakyReLU(0.2), nn.Dropout(0.3), nn.Linear(256, 1), nn.Sigmoid() # 输出概率值[0,1] ) def forward(self, img): img_flat img.view(img.size(0), -1) validity self.model(img_flat) return validity3. 训练过程与技巧GAN的训练是一个动态平衡过程需要精心调整超参数和训练策略。以下是训练GAN时的关键步骤和技巧3.1 初始化与损失函数首先初始化模型、优化器和损失函数# 设备配置 device torch.device(cuda if torch.cuda.is_available() else cpu) # 初始化生成器和判别器 latent_dim 100 generator Generator(latent_dim).to(device) discriminator Discriminator().to(device) # 定义优化器 lr 0.0002 beta1 0.5 g_optimizer torch.optim.Adam(generator.parameters(), lrlr, betas(beta1, 0.999)) d_optimizer torch.optim.Adam(discriminator.parameters(), lrlr, betas(beta1, 0.999)) # 损失函数 adversarial_loss nn.BCELoss()3.2 训练循环GAN的训练分为两个交替进行的阶段训练判别器和训练生成器。import time import matplotlib.pyplot as plt num_epochs 200 sample_interval 400 # 每隔多少步保存一次生成样本 for epoch in range(num_epochs): start_time time.time() for i, (imgs, _) in enumerate(train_loader): # 真实图像 real_imgs imgs.to(device) batch_size real_imgs.size(0) # 真实和假标签 real_labels torch.ones(batch_size, 1).to(device) fake_labels torch.zeros(batch_size, 1).to(device) # --------------------- # 训练判别器 # --------------------- d_optimizer.zero_grad() # 真实图像的损失 real_outputs discriminator(real_imgs) d_loss_real adversarial_loss(real_outputs, real_labels) # 生成假图像 z torch.randn(batch_size, latent_dim).to(device) fake_imgs generator(z) # 假图像的损失 fake_outputs discriminator(fake_imgs.detach()) d_loss_fake adversarial_loss(fake_outputs, fake_labels) # 总判别器损失 d_loss d_loss_real d_loss_fake d_loss.backward() d_optimizer.step() # --------------------- # 训练生成器 # --------------------- g_optimizer.zero_grad() # 生成器希望判别器将假图像分类为真 gen_outputs discriminator(fake_imgs) g_loss adversarial_loss(gen_outputs, real_labels) g_loss.backward() g_optimizer.step() # 打印训练状态 if i % 100 0: print( f[Epoch {epoch}/{num_epochs}] [Batch {i}/{len(train_loader)}] f[D loss: {d_loss.item():.4f}] [G loss: {g_loss.item():.4f}] ) # 每个epoch结束后保存生成样本 with torch.no_grad(): z torch.randn(10, latent_dim).to(device) gen_imgs generator(z).cpu() fig, axs plt.subplots(1, 10, figsize(20, 2)) for j in range(10): axs[j].imshow(gen_imgs[j].squeeze(), cmapgray) axs[j].axis(off) plt.show() epoch_time time.time() - start_time print(fEpoch {epoch} completed in {epoch_time:.2f} seconds)3.3 训练技巧与常见问题GAN训练过程中可能会遇到以下问题及解决方案模式崩溃(Mode Collapse): 生成器只生成有限的几种样本解决方法使用小批量判别(minibatch discrimination)、增加噪声或尝试不同的架构如WGAN判别器过强: 判别器学习过快导致生成器无法进步解决方法降低判别器的学习率或减少其更新频率梯度消失: 当判别器太完美时生成器梯度会消失解决方法使用改进的损失函数如Wasserstein损失4. 结果评估与改进训练完成后我们需要评估生成图像的质量并考虑可能的改进方向。4.1 生成样本可视化让我们生成一些样本并与真实MNIST图像进行对比import numpy as np # 生成100个样本 n_rows 10 n_cols 10 z torch.randn(n_rows * n_cols, latent_dim).to(device) gen_imgs generator(z).cpu().detach() # 绘制生成图像 fig, axes plt.subplots(n_rows, n_cols, figsize(10, 10)) for i, ax in enumerate(axes.flatten()): ax.imshow(gen_imgs[i].squeeze(), cmapgray) ax.axis(off) plt.tight_layout() plt.show() # 绘制真实图像 real_imgs, _ next(iter(train_loader)) fig, axes plt.subplots(n_rows, n_cols, figsize(10, 10)) for i, ax in enumerate(axes.flatten()): ax.imshow(real_imgs[i].squeeze(), cmapgray) ax.axis(off) plt.tight_layout() plt.show()4.2 定量评估指标虽然视觉检查很重要但我们也需要一些定量指标来评估生成质量Inception Score (IS): 衡量生成图像的多样性和可识别性Fréchet Inception Distance (FID): 比较生成图像和真实图像的统计特性注意对于MNIST这样的简单数据集这些指标可能不如视觉检查直观有效。4.3 模型改进方向如果对当前结果不满意可以考虑以下改进架构改进:使用DCGAN(深度卷积GAN)替代全连接网络尝试更先进的架构如ProGAN、StyleGAN训练技巧:使用标签平滑(label smoothing)实现渐进式训练(progressive growing)尝试不同的损失函数如Wasserstein损失超参数优化:调整学习率和优化器参数增加模型容量或调整噪声维度# DCGAN生成器示例 class DCGenerator(nn.Module): def __init__(self, latent_dim100): super(DCGenerator, self).__init__() self.main nn.Sequential( # 输入是Z, 进入全连接 nn.ConvTranspose2d(latent_dim, 512, 4, 1, 0, biasFalse), nn.BatchNorm2d(512), nn.ReLU(True), # 状态大小: (512) x 4 x 4 nn.ConvTranspose2d(512, 256, 4, 2, 1, biasFalse), nn.BatchNorm2d(256), nn.ReLU(True), # 状态大小: (256) x 8 x 8 nn.ConvTranspose2d(256, 128, 4, 2, 1, biasFalse), nn.BatchNorm2d(128), nn.ReLU(True), # 状态大小: (128) x 16 x 16 nn.ConvTranspose2d(128, 1, 4, 2, 1, biasFalse), nn.Tanh() # 状态大小: (1) x 32 x 32 ) def forward(self, input): return self.main(input)在实际项目中我发现使用卷积转置层(ConvTranspose2d)的DCGAN通常比全连接网络能生成更清晰的图像特别是当处理更复杂的数据集时。不过对于MNIST这样的简单数据集全连接网络通常已经足够。