本项目实现了一个完整的生成对抗网络(GAN)模型,用于生成手写数字图像。通过详细的代码注释和可视化,帮助理解GAN的工作原理和训练过程。
Gan-demo/
├── gan_demo.py # 主要的GAN实现代码
├── requirements.txt # 项目依赖
├── README.md # 项目说明文档
└── data/
└── mnist.npz # MNIST数据集(Keras格式)
生成对抗网络(Generative Adversarial Network)由两个神经网络组成:
- 生成器(Generator): 从随机噪声生成假图像
- 判别器(Discriminator): 判断图像是真实的还是生成的
生成器和判别器进行对抗训练:
- 生成器试图生成能够欺骗判别器的图像
- 判别器试图正确区分真实图像和生成图像
- 两者在训练过程中不断改进,最终生成器能够生成逼真的图像
# 从本地data/mnist.npz加载Keras格式MNIST数据集
mnist_data = np.load('data/mnist.npz')
x_train = mnist_data['x_train'] # (60000, 28, 28)# 归一化到[-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)train_tensor = torch.tensor(x_train)
dataset = TensorDataset(train_tensor)
dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)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图像
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激活函数
- 最终输出一个概率值,表示图像为真实的概率
# 超参数设置
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))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()- 目标:让判别器将生成的图像误判为真实图像
- 损失函数:
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
# 每隔10个epoch显示生成的图像
if epoch % 10 == 0 or epoch == 1:
G.eval()
with torch.no_grad():
fake_imgs = G(fixed_noise).detach().cpu() # 使用固定噪声生成图像
# 显示5×5网格的生成图像# 绘制生成器和判别器的损失曲线
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 | 训练轮数,影响模型性能 |
- 理想情况:生成器和判别器损失都应该逐渐下降
- 判别器过强:判别器损失过低,生成器难以学习
- 生成器过强:生成器损失过低,可能出现模式崩塌
- 初期:生成图像为随机噪声
- 中期:开始出现数字轮廓
- 后期:生成清晰的数字图像
- 生成器和判别器需要保持平衡
- 如果一方过强,另一方将无法有效学习
- 可以通过调整学习率或训练频率来平衡
- 安装依赖:
pip install -r requirements.txt- 运行训练:
python gan_demo.py- 观察输出:
- 控制台显示训练进度和损失
- 弹出窗口显示生成的图像
- 训练结束后显示损失曲线
- 改进网络结构:使用卷积神经网络(CNN)替代全连接层
- 调整超参数:尝试不同的学习率、批次大小等
- 添加正则化:使用梯度惩罚、谱归一化等技术
- 更换数据集:尝试其他图像数据集如CIFAR-10
- 实现其他GAN变体:如DCGAN、WGAN、StyleGAN等
本demo通过完整的GAN实现,展示了:
- 如何从本地数据加载和预处理数据
- 如何构建生成器和判别器网络
- 如何实现对抗训练过程
- 如何可视化和监控训练过程
通过运行这个demo,你可以直观地理解GAN的工作原理,观察生成图像从噪声到清晰数字的演化过程,为深入学习更复杂的生成模型打下基础。