Skip to content

训练一个 GAN

在探索了生成对抗网络(Generative Adversarial Network, GAN)的架构和工作原理后,本章将提供一个实现和训练 GAN 的实际示例。我们将使用 Python 和 TensorFlow 及其 Keras API 构建一个 GAN,用于生成类似于著名 MNIST 数据集中的手写数字。

训练 GAN 涉及对两个神经网络(生成器和判别器)进行迭代优化。典型的训练过程分解如下:

  • 系统包含两个神经网络:生成器网络 (G) 和判别器网络 (D)。它们的权重通常随机初始化。
  • 生成器 (G) 接收一个随机噪声向量(隐向量,latent vector)作为输入,并旨在生成合成数据样本。
  • 判别器 (D) 接收数据样本(来自数据集的真实样本或来自 G 的合成样本)作为输入,并尝试将其分类为“真实”(real)或“虚假”(fake)。
  • 将随机噪声向量(例如,来自高斯分布)输入到生成器网络中。
  • 生成器通过其各层处理此噪声,输出旨在模拟真实数据的合成数据样本。
  • 从训练数据集中抽取一批真实数据样本。
  • 生成器使用随机噪声生成一批虚假数据样本。
  • 判别器在这批真实样本和虚假样本上进行训练。其目标是正确地将真实样本识别为真实,虚假样本识别为虚假。更新其权重以最小化其分类误差(例如,使用二元交叉熵损失,binary cross-entropy loss)。
  • 生成器生成一批新的虚假数据样本。
  • 将这些虚假样本通过判别器(在此阶段其权重保持冻结)。
  • 根据判别器的输出计算生成器的损失(loss)。生成器旨在“欺骗”(fool)判别器,即让判别器将其虚假样本分类为真实。更新其权重以最小化此损失(例如,通过最小化判别器判断正确的负对数概率,或最大化判别器被欺骗的对数概率)。
  • 重复步骤 2(生成虚假数据)、3(判别器训练)和 4(生成器训练),进行多次迭代(iterations)或周期(epochs)。
  • 在每次迭代中,生成器和判别器交替(或有时以不同频率)训练,不断尝试超越对方。
  • 这种对抗(adversarial)过程理想情况下会达到一个平衡点(equilibrium),此时生成器产生高度逼真的数据,而判别器已无法可靠地区分真实样本和虚假样本(其准确率徘徊在 50% 左右)。

现在,我们将逐步介绍使用 Python、TensorFlow 和 MNIST 手写数字数据集构建和训练 GAN 的过程。

首先,确保你的 Python 环境安装了必要的库。你主要需要 TensorFlow(包含 Keras)和 Matplotlib。如果尚未安装,可以使用 pip 进行安装:

pip install tensorflow matplotlib numpy

首先在 Python 脚本中导入所需的模块:

import numpy as np
import tensorflow as tf
from tensorflow.keras import layers, models, optimizers, losses
from tensorflow.keras.datasets import mnist
import matplotlib.pyplot as plt

步骤 3:加载并预处理 MNIST 数据集

Section titled “步骤 3:加载并预处理 MNIST 数据集”

MNIST 数据集包含 60,000 张训练图像和 10,000 张测试图像,都是 28x28 像素的手写数字 (0-9)。我们将像素值归一化到 [-1, 1] 范围,这通常对 GAN 训练有利,尤其是在生成器的输出层使用 tanh 激活函数时。

# Load the dataset (we only need the training images for GANs)
(x_train, _), (_, _) = mnist.load_data()
# Normalize the images to the range [-1, 1]
x_train = (x_train.astype('float32') - 127.5) / 127.5
# Add a channel dimension (for grayscale)
x_train = np.expand_dims(x_train, axis=-1)
# Define constants
BUFFER_SIZE = 60000 # Equal to the number of training examples
BATCH_SIZE = 256
NOISE_DIM = 100 # Dimensionality of the random noise vector

步骤 4:创建生成器和判别器模型

Section titled “步骤 4:创建生成器和判别器模型”

生成器将从随机噪声创建虚假数字图像,而判别器将尝试区分真实的 MNIST 数字和这些虚假图像。

生成器接收随机噪声向量(隐向量,latent vector)作为输入,并经过多个层进行变换以生成 28x28x1 的图像。我们使用 Dense 层,后跟 BatchNormalization 和 LeakyReLU 激活函数,最后使用 tanh 激活函数将输出值映射到 [-1, 1] 范围。

def build_generator():
model = models.Sequential(name='Generator')
model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(NOISE_DIM,)))
model.add(layers.BatchNormalization())
model.add(layers.LeakyReLU())
model.add(layers.Reshape((7, 7, 256)))
# assert model.output_shape == (None, 7, 7, 256) # Note: None is the batch size
model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False))
# assert model.output_shape == (None, 7, 7, 128)
model.add(layers.BatchNormalization())
model.add(layers.LeakyReLU())
model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False))
# assert model.output_shape == (None, 14, 14, 64)
model.add(layers.BatchNormalization())
model.add(layers.LeakyReLU())
model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh'))
# assert model.output_shape == (None, 28, 28, 1)
return model
generator = build_generator()

判别器是一个基于 CNN 的分类器。它接收 28x28x1 的图像作为输入(真实或生成的图像),并输出一个标量值,表示该图像是真实的概率(或 logit)。我们使用 Conv2D 层、LeakyReLU 激活、Dropout 进行正则化,以及一个不带激活的最终 Dense 层(输出 logit)。

def build_discriminator():
model = models.Sequential(name='Discriminator')
model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same', input_shape=[28, 28, 1]))
model.add(layers.LeakyReLU())
model.add(layers.Dropout(0.3))
model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same'))
model.add(layers.LeakyReLU())
model.add(layers.Dropout(0.3))
model.add(layers.Flatten())
model.add(layers.Dense(1)) # Output logits (no sigmoid activation here)
return model
discriminator = build_discriminator()

我们将使用二元交叉熵损失(binary cross-entropy loss)。由于我们的判别器输出的是 logit,我们将 from_logits 设置为 True。生成器旨在使判别器将其生成的虚假图像输出为“真实”(标签 1)。判别器旨在将真实图像正确分类为“真实”(标签 1),将虚假图像分类为“虚假”(标签 0)。Adam 是 GAN 常用的优化器(optimizers)。

# Binary cross-entropy loss function (expects logits)
cross_entropy = losses.BinaryCrossentropy(from_logits=True)
def discriminator_loss(real_output, fake_output):
real_loss = cross_entropy(tf.ones_like(real_output), real_output)
fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output)
total_loss = real_loss + fake_loss
return total_loss
def generator_loss(fake_output):
# Generator wants discriminator to classify fake images as real (label 1)
return cross_entropy(tf.ones_like(fake_output), fake_output)
# Optimizers (Adam is commonly used for GANs)
generator_optimizer = optimizers.Adam(learning_rate=1e-4)
discriminator_optimizer = optimizers.Adam(learning_rate=1e-4)

步骤 6:定义训练循环(训练步骤)

Section titled “步骤 6:定义训练循环(训练步骤)”

训练步骤涉及生成虚假图像、计算两个网络的损失,以及使用反向传播(backpropagation)和定义的优化器更新它们的权重。我们使用 tf.GradientTape 来记录操作以便进行自动微分。

@tf.function
def train_step(images):
noise = tf.random.normal([BATCH_SIZE, NOISE_DIM])
with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:
generated_images = generator(noise, training=True)
real_output = discriminator(images, training=True)
fake_output = discriminator(generated_images, training=True)
gen_loss = generator_loss(fake_output)
disc_loss = discriminator_loss(real_output, fake_output)
gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables)
gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)
generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))
discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))
return gen_loss, disc_loss
# Main training function
def train(dataset, epochs):
for epoch in range(epochs):
epoch_gen_loss_avg = tf.keras.metrics.Mean()
epoch_disc_loss_avg = tf.keras.metrics.Mean()
for image_batch in dataset:
gen_loss, disc_loss = train_step(image_batch)
epoch_gen_loss_avg.update_state(gen_loss)
epoch_disc_loss_avg.update_state(disc_loss)
print(f'Epoch {epoch + 1}, Gen Loss: {epoch_gen_loss_avg.result():.4f}, Disc Loss: {epoch_disc_loss_avg.result():.4f}')
# Optionally, generate and save some images periodically
if (epoch + 1) % 10 == 0:
generate_and_save_images(generator, epoch + 1, seed_for_generating_images)

接下来,通过打乱(shuffling)和批量处理(batching)图像来准备 MNIST 数据集,然后开始训练过程。我们还定义一个固定的种子(seed)用于生成样本图像,以便观察训练进度。

# Prepare the dataset for training
train_dataset = tf.data.Dataset.from_tensor_slices(x_train).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)
# A fixed noise seed for generating sample images during training to visualize progress
NUM_EXAMPLES_TO_GENERATE = 16
seed_for_generating_images = tf.random.normal([NUM_EXAMPLES_TO_GENERATE, NOISE_DIM])
# Define training epochs
EPOCHS = 50 # Adjust as needed; GANs often require many epochs
# Start training
# train(train_dataset, EPOCHS) # Uncomment to run training

训练后(或训练期间定期),可以使用训练好的生成器生成新图像并显示它们。这包括创建随机噪声,将其输入到生成器中,并绘制生成的图像。

def generate_and_save_images(model, epoch, test_input):
# `training=False` so all layers run in inference mode (e.g., batchnorm).
predictions = model(test_input, training=False)
fig = plt.figure(figsize=(4, 4))
for i in range(predictions.shape[0]):
plt.subplot(4, 4, i + 1)
# Denormalize pixel values from [-1, 1] to [0, 1] for imshow if using grayscale
# or [0, 255] if needed. Here, * 0.5 + 0.5 maps to [0,1]
plt.imshow(predictions[i, :, :, 0] * 0.5 + 0.5, cmap='gray')
plt.axis('off')
# plt.savefig(f'image_at_epoch_{epoch:04d}.png') # Optional: save images
plt.show()
# Example of generating images after training (assuming 'generator' is trained):
# generate_and_save_images(generator, EPOCHS, seed_for_generating_images)
# Or if you load a pre-trained generator:
# loaded_generator = tf.keras.models.load_model('path_to_your_saved_generator_model')
# generate_and_save_images(loaded_generator, 0, seed_for_generating_images)

运行上述代码(取消对 train 调用和可能的最后 generate_and_save_images 调用的注释后),你将观察到训练期间生成器和判别器的损失(loss)。如果训练成功,generate_and_save_images 函数生成的图像将从嘈杂的图案逐渐改进为可识别的手写数字。输出通常是一个 16 张生成的数字图像组成的网格,随着周期的增加,它们会越来越像 MNIST 数据集中的图像。

训练 GAN 涉及几个关键步骤:设置环境、定义健壮的生成器和判别器模型、选择合适的损失函数(loss functions)和优化器(optimizers),以及精心实现对抗训练循环。通过遵循这些步骤,你可以训练自己的 GAN 来生成新的数据,正如我们在 MNIST 数据集上演示的那样。

本章提供了使用 Python、TensorFlow 和 Keras 构建和训练 GAN 的详细指南。GAN 训练对超参数(hyperparameters)和模型架构(architecture)很敏感,因此实验往往是取得良好结果的关键。本示例可作为探索生成对抗网络及其应用的有趣世界的入门基础。为了进一步学习,可以考虑探索更先进的 GAN 架构和技术,以稳定训练并提高样本质量,例如 Wasserstein GANs (WGANs) 或 Spectral Normalization。

GAN 训练中的常见挑战包括模式崩溃(mode collapse,生成器产生有限多样性的样本)和训练不稳定(instability)。TensorFlow 教程 (tensorflow.org/tutorials/generative) 和关于 GAN 改进的研究论文等资源可以提供更深入的见解和解决方案。