• 生成对抗网络入门案例


    前言

    生成对抗网络(Generative Adversarial Networks,简称GANs)是一种用于生成新样本的机器学习模型。它由两个主要组件组成:生成器(Generator)和判别器(Discriminator)。生成器尝试生成与训练数据相似的新样本,而判别器则试图区分生成器生成的样本和真实训练数据。

    下面是一个简单的对抗生成网络的入门例子,用于生成手写数字图像:

    实现过程

    1、导入必要的库和模块

    1. import numpy as np
    2. import matplotlib.pyplot as plt
    3. from tensorflow.keras.datasets import mnist
    4. from tensorflow.keras.models import Sequential
    5. from tensorflow.keras.layers import Dense, Flatten, Reshape
    6. from tensorflow.keras.layers import Conv2D, Conv2DTranspose
    7. from tensorflow.keras.optimizers import Adam

    2、加载MNIST数据集

    1. (x_train, _), (_, _) = mnist.load_data()
    2. x_train = x_train / 255.0
    3. x_train = np.expand_dims(x_train, axis=3)

    3、定义生成器模型

    1. generator = Sequential()
    2. generator.add(Dense(7*7*128, input_shape=(100,), activation='relu'))
    3. generator.add(Reshape((7, 7, 128)))
    4. generator.add(Conv2DTranspose(64, (3, 3), strides=(2, 2), padding='same', activation='relu'))
    5. generator.add(Conv2DTranspose(1, (3, 3), strides=(2, 2), padding='same', activation='sigmoid'))

    4、定义判别器模型

    1. discriminator = Sequential()
    2. discriminator.add(Conv2D(64, (3, 3), strides=(2, 2), padding='same', input_shape=(28, 28, 1), activation='relu'))
    3. discriminator.add(Conv2D(128, (3, 3), strides=(2, 2), padding='same', activation='relu'))
    4. discriminator.add(Flatten())
    5. discriminator.add(Dense(1, activation='sigmoid'))

    5、编译判别器模型

    discriminator.compile(loss='binary_crossentropy', optimizer=Adam(learning_rate=0.0002, beta_1=0.5), metrics=['accuracy'])

    6、冻结判别器模型的权重

    discriminator.trainable = False

    7、定义GAN模型

    1. gan = Sequential()
    2. gan.add(generator)
    3. gan.add(discriminator)

    8、编译GAN模型

    gan.compile(loss='binary_crossentropy', optimizer=Adam(learning_rate=0.0002, beta_1=0.5))

    9、定义训练函数

    1. def train_gan(epochs, batch_size, sample_interval):
    2. for epoch in range(epochs):
    3. # 生成随机噪声作为输入
    4. noise = np.random.normal(0, 1, (batch_size, 100))
    5. # 生成假样本
    6. generated_images = generator.predict(noise)
    7. # 从真实样本中随机选择一批样本
    8. real_images = x_train[np.random.randint(0, x_train.shape[0], batch_size)]
    9. # 训练判别器
    10. discriminator_loss_real = discriminator.train_on_batch(real_images, np.ones((batch_size, 1)))
    11. discriminator_loss_fake = discriminator.train_on_batch(generated_images, np.zeros((batch_size, 1)))
    12. discriminator_loss = 0.5 * np.add(discriminator_loss_real, discriminator_loss_fake)
    13. # 训练生成器
    14. noise = np.random.normal(0, 1, (batch_size, 100))
    15. generator_loss = gan.train_on_batch(noise, np.ones((batch_size, 1)))
    16. # 打印损失
    17. if epoch % sample_interval == 0:
    18. print(f"Epoch {epoch}/{epochs}, Discriminator Loss: {discriminator_loss[0]}, Generator Loss: {generator_loss}")
    19. # 保存生成的图像
    20. save_images(epoch)

    10、保存生成的图像

    1. def save_images(epoch):
    2. rows, cols = 5, 5
    3. noise = np.random.normal(0, 1, (rows * cols, 100))
    4. generated_images = generator.predict(noise)
    5. generated_images = 0.5 * generated_images + 0.5
    6. fig, axs = plt.subplots(rows, cols)
    7. idx = 0
    8. for i in range(rows):
    9. for j in range(cols):
    10. axs[i, j].imshow(generated_images[idx, :, :, 0], cmap='gray')
    11. axs[i, j].axis('off')
    12. idx += 1
    13. fig.savefig(f"gan_images/mnist_{epoch}.png")
    14. plt.close()

    11、训练GAN模型

    1. epochs = 10000
    2. batch_size = 128
    3. sample_interval = 1000

    完整代码

    1. import numpy as np
    2. import matplotlib.pyplot as plt
    3. from tensorflow.keras.datasets import mnist
    4. from tensorflow.keras.models import Sequential
    5. from tensorflow.keras.layers import Dense, Flatten, Reshape
    6. from tensorflow.keras.layers import Conv2D, Conv2DTranspose
    7. from tensorflow.keras.optimizers import Adam
    8. # 加载MNIST数据集
    9. (x_train, _), (_, _) = mnist.load_data()
    10. x_train = x_train / 255.0
    11. x_train = np.expand_dims(x_train, axis=3)
    12. # 定义生成器模型
    13. generator = Sequential()
    14. generator.add(Dense(7*7*128, input_shape=(100,), activation='relu'))
    15. generator.add(Reshape((7, 7, 128)))
    16. generator.add(Conv2DTranspose(64, (3, 3), strides=(2, 2), padding='same', activation='relu'))
    17. generator.add(Conv2DTranspose(1, (3, 3), strides=(2, 2), padding='same', activation='sigmoid'))
    18. # 定义判别器模型
    19. discriminator = Sequential()
    20. discriminator.add(Conv2D(64, (3, 3), strides=(2, 2), padding='same', input_shape=(28, 28, 1), activation='relu'))
    21. discriminator.add(Conv2D(128, (3, 3), strides=(2, 2), padding='same', activation='relu'))
    22. discriminator.add(Flatten())
    23. discriminator.add(Dense(1, activation='sigmoid'))
    24. # 编译判别器模型
    25. discriminator.compile(loss='binary_crossentropy', optimizer=Adam(learning_rate=0.0002, beta_1=0.5), metrics=['accuracy'])
    26. # 冻结判别器模型的权重
    27. discriminator.trainable = False
    28. # 定义GAN模型
    29. gan = Sequential()
    30. gan.add(generator)
    31. gan.add(discriminator)
    32. # 编译GAN模型
    33. gan.compile(loss='binary_crossentropy', optimizer=Adam(learning_rate=0.0002, beta_1=0.5))
    34. # 定义训练函数
    35. def train_gan(epochs, batch_size, sample_interval):
    36. for epoch in range(epochs):
    37. # 生成随机噪声作为输入
    38. noise = np.random.normal(0, 1, (batch_size, 100))
    39. # 生成假样本
    40. generated_images = generator.predict(noise)
    41. # 从真实样本中随机选择一批样本
    42. real_images = x_train[np.random.randint(0, x_train.shape[0], batch_size)]
    43. # 训练判别器
    44. discriminator_loss_real = discriminator.train_on_batch(real_images, np.ones((batch_size, 1)))
    45. discriminator_loss_fake = discriminator.train_on_batch(generated_images, np.zeros((batch_size, 1)))
    46. discriminator_loss = 0.5 * np.add(discriminator_loss_real, discriminator_loss_fake)
    47. # 训练生成器
    48. noise = np.random.normal(0, 1, (batch_size, 100))
    49. generator_loss = gan.train_on_batch(noise, np.ones((batch_size, 1)))
    50. # 打印损失
    51. if epoch % sample_interval == 0:
    52. print(f"Epoch {epoch}/{epochs}, Discriminator Loss: {discriminator_loss[0]}, Generator Loss: {generator_loss}")
    53. # 保存生成的图像
    54. save_images(epoch)
    55. # 保存生成的图像
    56. def save_images(epoch):
    57. rows, cols = 5, 5
    58. noise = np.random.normal(0, 1, (rows * cols, 100))
    59. generated_images = generator.predict(noise)
    60. generated_images = 0.5 * generated_images + 0.5
    61. fig, axs = plt.subplots(rows, cols)
    62. idx = 0
    63. for i in range(rows):
    64. for j in range(cols):
    65. axs[i, j].imshow(generated_images[idx, :, :, 0], cmap='gray')
    66. axs[i, j].axis('off')
    67. idx += 1
    68. fig.savefig(f"gan_images/mnist_{epoch}.png")
    69. plt.close()
    70. # 训练GAN模型
    71. epochs = 10000
    72. batch_size = 128
    73. sample_interval = 1000
    74. train_gan(epochs, batch_size, sample_interval)

    训练结果:

    这个例子使用了MNIST数据集,生成手写数字图像。生成器和判别器模型使用了卷积神经网络的结构。在训练过程中,生成器试图生成逼真的手写数字图像,而判别器则试图区分真实图像和生成图像。通过反复迭代训练生成器和判别器,GAN模型能够逐渐生成更逼真的手写数字图像。生成的图像会保存在gan_images文件夹中。

  • 相关阅读:
    kubernetes-Service详解
    【框架】Spring Framework :SpringBoot
    【开源】基于Vue.js的校园二手交易系统的设计和实现
    [附源码]计算机毕业设计JAVA校园一卡通管理信息系统台
    linux系统iptables的操作
    2022 极术通讯-Arm体系结构的同步概述和案例研究
    Degrade is Upgrade: Learning Degradation for Low-light Image Enhancement论文阅读笔记
    上周热点回顾(10.23-10.29)
    机器学习 | 模型评估和选择 各种评估指标总结——错误率精度-查准率查全率-真正例率假正例率 PR曲线ROC曲线
    基于密码芯片的 DDR 加速器的设计与实现
  • 原文地址:https://blog.csdn.net/qq_39312146/article/details/133559186