Keras搭建SRGAN

一. SRGAN介绍

在这里插入图片描述

1.1 相关知识介绍

  1. 相关定义SR(Super Resolution) LR(Low Resolution) HR(High Resolution) MOS(mean opinion score)
  2. 提出新评价标准
    • SRPSNR标准应当换掉,因为高的PSNR并不保证高的感官保真。
    • 提出了新的评价标准 M O S ( M e a n    o p i n i o n    s c o r e ) \color{red}MOS(Mean\;opinion\;score) MOSMeanopinionscore
  3. 提出新Loss
    • Loss用 p e r c e p t u a l    l o s s \color{red}perceptual\;loss perceptualloss: 使用pixel-wise的MSE使得图像变得平滑,而如果先用VGG来抓取到高级特征(feature)表示,再对feature使用MSE,可以更好的抓取不变特征。
      • 该思想来源于[1603.08155] Perceptual Losses for Real-Time Style Transfer and Super-Resolution
      • 具体思想是:一个是Fast Neural Style(快速的画风迁移),另一个是提出了一种单张图像的超分辨率算法。此外,在这篇文章中还提出了一种新的损失Perceptual Loss(感知损失),感知损失由三个部分组成: 感 知 损 失 = 特 征 重 构 损 失 + 风 格 重 构 损 失 + 简 单 损 失 \color{red}感知损失=特征重构损失+风格重构损失+简单损失 =++,不仅考虑到了特征重构后的相似性,也考虑到了低层特征的相似性。
      • 我们来思考一个问题, 为 什 么 超 分 中 丢 失 的 是 高 频 信 息 \color{red}为什么超分中丢失的是高频信息
        我们可以这样考虑,超分问题的本质是通过不同的上采样的方式从一个低分辨率图像恢复到高分辨率图像,从像素级别的角度来看,这是一个一对多或者多对多的问题,那么我们就可以认为这是一个回归问题。既然是回归问题,在拟合的过程中要保证尽量多的信息可以恢复准确,而在图像中,低频信息占大多数,而高频信息占少数,所以在超分问题中高频信息就丢失了
      • 用不同的损失来判定的结果:
        在这里插入图片描述

1.2 SRGAN网络介绍

1.2.1 生成模型(Generater)

生成网络的构成如下图所示:
在这里插入图片描述

1.2.1.1 基础性理解

从左至右来看,SRGAN的生成网络由三个部分组成。

  1. 低分辨率图像进入后会经过一个卷积+RELU函数
  2. 然后经过B个残差网络结构,每个残差网络内部包含两个卷积+标准化+RELU,还有一个残差边
  3. 然后进入上采样部分,将长宽进行放大,两次上采样后,变为原来的4倍,实现提高分辨率。
1.2.1.2 整体理解
  • 生成器: 【 3 x 3 c o n v + B N + P R e L U + 2 s u b − p i x e l c o n v 】 ∗ n \color{red}【3x3 conv + BN + PReLU + 2 sub-pixel conv】 * n 3x3conv+BN+PReLU+2subpixelconvn
  • 生成器是在 S R R e s N e t \color{red}SRResNet SRResNet的基础上做了改进,在生成网络部分(SRResNet)部分包含多个残差块。

    每个残差块中包含两个3×3的卷积层,卷积层后接批规范化层(batch normalization, BN)和LReLU作为激活函数,两个2×亚像素卷积层(sub-pixel convolution layers)被用来增大特征尺寸。

1.2.1.3 代码
  def build_generator(self):
      # 残差函数
      def residual_block(layer_input, filters):
          d = Conv2D(filters=filters, kernel_size=3, strides=1, padding='same')(layer_input)
          d = BatchNormalization(momentum=0.8)(d)
          d = Activation('relu')(d)
          d = Conv2D(filters=filters, kernel_size=3, strides=1, padding='same')(d)
          d = BatchNormalization(momentum=0.8)(d)
          d = Add()([d, layer_input])
          return d

      def deconv2d(layer_input):
          u = UpSampling2D(size=2)(layer_input)
          u = Conv2D(filters=256, kernel_size=3, strides=1, padding='same')(u)
          u = Activation('relu')(u)
          return u

      img_lr = Input(shape=self.lr_shape)
      # 第一部分:低分辨率图像进入后会经过一个Con2D+RELU函数
      c1 = Conv2D(filters=64, kernel_size=9, strides=1, padding='same')(img_lr)
      c1 = Activation('relu')(c1)

      # 第二部分:经过16个残差网络,每个残差网络包括(两个Con2D+标准化+RELU+残差边)
      r = residual_block(c1, 64)
      for _ in range(self.n_residual_blocks - 1):
          r = residual_block(r, 64)

      # 第三部分:上采样,将长宽进行放大,两次上采样,变成原来的4倍,实现提高分辨率的效果
      c2 = Conv2D(filters=64, kernel_size=3, strides=1, padding='same')(r)
      c2 = BatchNormalization(momentum=0.8)(c2)
      c2 = Add()([c2, c1])
      u1 = deconv2d(c2)
      u2 = deconv2d(u1)
      gen_hr = Conv2D(self.channels, kernel_size=9, strides=1, padding='same', activation='tanh')(u2)

      return Model(img_lr, gen_hr)

1.2.2 判别模型(Discriminator)

在这里插入图片描述

1.2.2.1 框架介绍
  1. SRGAN的判别网络由不断重复的 卷积+LeakyRELU和标准化 组成。
  2. 具体地说:
    • 判别器: 【 8 c o n v + L e a k y R e L U + 2 f c + s i g m o i d 】 \color{red}【8 conv + LeakyReLU + 2 fc + sigmoid】 8conv+LeakyReLU+2fc+sigmoid
    • 在判别网络部分包含8个卷积层,随着网络层数加深,特征个数不断增加,特征尺寸不断减小,选取激活函数为LeakyReLU,最终通过两个全连接层和最终的sigmoid激活函数得到预测为自然图像的概率。
1.2.2.2 代码
 def bulid_discriminator(self):

     def d_block(layer_input, filters, strides=1, bn=True):
         d = Conv2D(filters=filters, kernel_size=3, strides=strides, padding='same')(layer_input)
         d = LeakyReLU(alpha=0.2)(d)
         if bn:
             d = BatchNormalization(momentum=0.8)(d)
         return d

     # 判别模型的组成
     d0 = Input(shape=self.hr_shape)

     d1 = d_block(d0, 64, bn=False)
     d1 = d_block(d1, 64, strides=2)
     d1 = d_block(d1, 64 * 2)
     d1 = d_block(d1, 64 * 2, strides=2)
     d2 = d_block(d1, 64 * 4)
     d2 = d_block(d2, 64 * 4, strides=2)
     d3 = d_block(d2, 64 * 8)
     d3 = d_block(d3, 64 * 8, strides=2)

     d4 = Dense(64 * 16)(d3)
     d5 = LeakyReLU(alpha=0.2)(d4)
     validity = Dense(1, activation='sigmoid')(d5)

     return Model(d0, validity)

1.2.3 VGG网络

vgg网络: 【 P r e t r a i n e d    v g g    l o s s 】 \color{red}【Pretrained\;vgg\;loss】 Pretrainedvggloss
本文在生成器结束以后生成的SR图像输送到在ImageNet上已经预训练好的网络,在训练时不训练权重,只参与Loss的计算。

def build_vgg(self):
    # 创建VGG模型,只使用第9层的特征
    vgg = VGG19(weights='imagenet', include_top=False, input_shape=self.hr_shape)

    img_features = [vgg.layers[9].output]
    return Model(vgg.input, img_features)

1.3 训练思路

1.3.1 对判别模型进行训练

  • 将真实的高分辨率图像和虚假的高分辨率图像传入判别模型中
  • 将真实的高分辨率图像的判别结果与1对比得到loss。
  • 将虚假的高分辨率图像的判别结果与0对比得到loss。
  • 利用得到的loss进行训练。

1.3.2 对生成模型进行训练

  • 将低分辨率图像传入生成模型,得到高分辨率图像,利用该高分辨率图像获得判别结果与1进行对比得到loss。
  • 将真实的高分辨率图像和虚假的高分辨率图像传入VGG网络,获得两个图像的特征,通过这两个图像的特征进行比较获得loss。

在这里插入图片描述

二. 公式化理解

2.1 目标函数

SRGAN 目标函数的 GAN 部分公式如下:
在这里插入图片描述
I H R I^{HR} IHR代表高分辨率图像。 I L R I^{LR} ILR代表低分辨率图像(不同于传统 GAN 中的高斯噪声,我们将 低 分 辨 率 图 像 作 为 输 入 传 递 给 生 成 器 \color{red}低分辨率图像作为输入传递给生成器 )。目标函数的其余部分类似于传统的 GAN。

2.2 总损失函数

在这里插入图片描述
P e r c e p t u a l    L o s s ( f o r    V G G    b a s e d    c o n t e n t    l o s s ) = C o n t e n t    L o s s + A d v e r s a r i a l    L o s s (2.1) \color{red}Perceptual\; Loss(for\; VGG\;based\;content\;loss)= Content\; Loss + Adversarial\; Loss\tag{2.1} PerceptualLoss(forVGGbasedcontentloss)=ContentLoss+AdversarialLoss(2.1)
公式(2.1)可以写成:
l S R = l X S R + 1 0 − 3 l G e n S R (2.2) \color{red}l^{SR} = l^{SR}_X + 10^{-3}l^{SR}_{Gen}\tag{2.2} lSR=lXSR+103lGenSR(2.2)
其中:

  • C o n t e n t    L o s s Content\; Loss ContentLoss可以表示成 M S E    L o s s MSE\; Loss MSELoss V G G    L o s s VGG\;Loss VGGLoss之和:
  • l X S R l^{SR}_X lXSR using perceptual similarityinstead of similarity in pixel space.
  • l G e n S R l^{SR}_{Gen} lGenSR push SR image to thenatural image manifold.

2.2.1 Content Loss

  1. MSE
    MSE丢失内容 比 较 I H R 与 生 成 的 图 像 G θ G ( I L R ) , 并 采 取 了 M S E 这 种 差 异 \color{red}比较I^{HR}与生成的图像G_{\theta_G}(I^{LR}),并采取了MSE这种差异 IHRGθG(ILR)MSE在这里插入图片描述
    • W , H W,H W,H:图像的宽度和高度;
    • r r r: 下采样因子;
    • G θ G G_{\theta_{G}} GθG发生器网络的前馈CNN参数;
    • θ G = { W 1 : L ; b 1 : L } \theta_G=\{W_{1:L};b_{1:L}\} θG={W1:Lb1:L}:L层网的权重和偏差.
  2. VGG内容损失
    • 与 MSE 损失相比, V G G 内 容 损 失 对 像 素 空 间 的 变 化 更 具 不 变 性 , 从 而 带 来 更 好 的 感 知 质 量 \color{red}VGG 内容损失对像素空间的变化更具不变性,从而带来更好的感知质量 VGG。因此 SRGAN 使用 VGG 损失来进行内容损失。
    • VGG内容损失是第j个卷积(活化后)取特征矢量( Φ 函 数 \Phi函数 Φ),但VGG19网络内的第i个maxpool层之前。
      在这里插入图片描述
    • Φ i , j \Phi_{i,j} Φi,j : feature map obtained by the j-th convolution afteractivation and before the i-th maxpooling layer withinthe VGG 19(VGG 19 中激活后第 i 个最大池化层之前的第 j 个卷积获得的特征图)。
    • W i , j H i , j W_{i,j}H_{i,j} Wi,jHi,j : 特征图的维度参数.

2.2.2 Adversarial Loss

在这里插入图片描述

三. 代码

3.1 数据加载代码

文件命名为:DataLoader.py

import imageio
from skimage.transform import resize
import scipy
from glob import glob
import numpy as np
import matplotlib.pyplot as plt


class DataLoader():
    def __init__(self, dataset_name, img_res=(128, 128)):
        self.dataset_name = dataset_name
        self.img_res = img_res

    def load_data(self, batch_size=1, is_testing=False):
        data_type = "train" if not is_testing else "test"

        path = glob('./datasets/%s/train/*' % self.dataset_name)

        batch_images = np.random.choice(path, size=batch_size)

        imgs_hr = []
        imgs_lr = []
        for img_path in batch_images:
            img = self.imread(img_path)

            h, w = self.img_res
            low_h, low_w = int(h / 4), int(w / 4)

            img_hr = resize(img, self.img_res)
            img_lr = resize(img, (low_h, low_w))

            # If training => do random flip
            if not is_testing and np.random.random() < 0.5:
                img_hr = np.fliplr(img_hr)
                img_lr = np.fliplr(img_lr)

            imgs_hr.append(img_hr)
            imgs_lr.append(img_lr)

        imgs_hr = np.array(imgs_hr) / 127.5 - 1.
        imgs_lr = np.array(imgs_lr) / 127.5 - 1.

        return imgs_hr, imgs_lr

    def imread(self, path):
        return imageio.imread(path, pilmode='RGB').astype(np.float)

3.2 主函数训练代码

from __future__ import print_function, division
import tensorflow.keras.backend as K
from keras_contrib.layers.normalization.instancenormalization import InstanceNormalization
from tensorflow.keras.layers import Input, Dense, BatchNormalization, Activation, Add
from tensorflow.keras.layers import PReLU, LeakyReLU, UpSampling2D, Conv2D
from tensorflow.keras.applications import VGG19
from tensorflow.keras.models import Model
from tensorflow.keras.optimizers import Adam
import datetime
import matplotlib.pyplot as plt
import sys
from DataLoader import DataLoader
import numpy as np
import os


class SRGAN():
    def __init__(self):
        # 低分辨率图的shape
        self.channels = 3
        self.lr_height = 128
        self.lr_width = 128
        self.lr_shape = (self.lr_height, self.lr_width, self.channels)
        # 高分辨率图的shape
        self.hr_height = self.lr_height * 4
        self.hr_width = self.lr_width * 4
        self.hr_shape = (self.hr_height, self.hr_width, self.channels)

        # 16个残差卷积快
        self.n_residual_blocks = 16
        # optimizer
        optimizer = Adam(0.0002, 0.5)

        # 创建VGG模型
        self.vgg = self.build_vgg()
        self.vgg.trainable = False

        # 导入数据集
        self.dataset_name = 'DIV2K'
        self.data_loader = DataLoader(dataset_name=self.dataset_name,
                                      img_res=(self.hr_height, self.hr_width))
        patch = int(self.hr_height / 2 ** 4)
        self.disc_patch = (patch, patch, 1)

        # 建立判别模型
        self.discriminator = self.bulid_discriminator()
        self.discriminator.compile(loss='binary_crossentropy',
                                   optimizer=optimizer,
                                   metrics=['accuracy'])
        self.discriminator.summary()

        # 建立生成模型
        self.generator = self.build_generator()
        self.generator.summary()

        # 将生成模型和判别模型结合,生成模型训练时候,训练时候不训练判别模型
        img_lr = Input(shape=self.lr_shape)

        fake_hr = self.generator(img_lr)
        fake_features = self.vgg(fake_hr)

        self.discriminator.trainable = False
        validity = self.discriminator(fake_hr)
        self.combined = Model(img_lr, [validity, fake_features])
        self.combined.compile(loss=['binary_crossentropy', 'mse'],
                              loss_weights=[5e-1, 1],
                              optimizer=optimizer)

    def build_vgg(self):
        # 创建VGG模型,只使用第9层的特征
        vgg = VGG19(weights='imagenet', include_top=False, input_shape=self.hr_shape)
     
        img_features = [vgg.layers[9].output]
        return Model(vgg.input, img_features)

    def build_generator(self):
        # 残差函数
        def residual_block(layer_input, filters):
            d = Conv2D(filters=filters, kernel_size=3, strides=1, padding='same')(layer_input)
            d = BatchNormalization(momentum=0.8)(d)
            d = Activation('relu')(d)
            d = Conv2D(filters=filters, kernel_size=3, strides=1, padding='same')(d)
            d = BatchNormalization(momentum=0.8)(d)
            d = Add()([d, layer_input])
            return d

        def deconv2d(layer_input):
            u = UpSampling2D(size=2)(layer_input)
            u = Conv2D(filters=256, kernel_size=3, strides=1, padding='same')(u)
            u = Activation('relu')(u)
            return u

        img_lr = Input(shape=self.lr_shape)
        # 第一部分:低分辨率图像进入后会经过一个Con2D+RELU函数
        c1 = Conv2D(filters=64, kernel_size=9, strides=1, padding='same')(img_lr)
        c1 = Activation('relu')(c1)

        # 第二部分:经过16个残差网络,每个残差网络包括(两个Con2D+标准化+RELU+残差边)
        r = residual_block(c1, 64)
        for _ in range(self.n_residual_blocks - 1):
            r = residual_block(r, 64)

        # 第三部分:上采样,将长宽进行放大,两次上采样,变成原来的4倍,实现提高分辨率的效果
        c2 = Conv2D(filters=64, kernel_size=3, strides=1, padding='same')(r)
        c2 = BatchNormalization(momentum=0.8)(c2)
        c2 = Add()([c2, c1])
        u1 = deconv2d(c2)
        u2 = deconv2d(u1)
        gen_hr = Conv2D(self.channels, kernel_size=9, strides=1, padding='same', activation='tanh')(u2)

        return Model(img_lr, gen_hr)

    def bulid_discriminator(self):

        def d_block(layer_input, filters, strides=1, bn=True):
            d = Conv2D(filters=filters, kernel_size=3, strides=strides, padding='same')(layer_input)
            d = LeakyReLU(alpha=0.2)(d)
            if bn:
                d = BatchNormalization(momentum=0.8)(d)
            return d

        # 判别模型的组成
        d0 = Input(shape=self.hr_shape)

        d1 = d_block(d0, 64, bn=False)
        d1 = d_block(d1, 64, strides=2)
        d1 = d_block(d1, 64 * 2)
        d1 = d_block(d1, 64 * 2, strides=2)
        d2 = d_block(d1, 64 * 4)
        d2 = d_block(d2, 64 * 4, strides=2)
        d3 = d_block(d2, 64 * 8)
        d3 = d_block(d3, 64 * 8, strides=2)

        d4 = Dense(64 * 16)(d3)
        d5 = LeakyReLU(alpha=0.2)(d4)
        validity = Dense(1, activation='sigmoid')(d5)

        return Model(d0, validity)

    def scheduler(self, models, epoch):
        # 学习率下降
        if epoch % 20000 == 0 and epoch != 0:
            for model in models:
                lr = K.get_value(model.optimizer.lr)
                K.set_value(model.optimizer.lr, lr * 0.5)
            print("lr changed to {}".format(lr * 0.5))

    def train(self, epochs, init_epoch=0, batch_size=1, sample_interval=50):
        # 开始计时
        start_time = datetime.datetime.now()
        if init_epoch != 0:
            self.generator.load_weights("weights/%s/gen_epoch%d.h5" % (self.dataset_name, init_epoch),
                                        skip_mismatch=False)
            self.discriminator.load_weights("weights/%s/dis_epoch%d.h5" % (self.dataset_name, init_epoch),
                                            skip_mismatch=False)

        for epoch in range(init_epoch, epochs):
            # 更改学习率
            self.scheduler([self.combined, self.discriminator], epoch)
            # -------------------- #
            # 训练判别器
            # ------------------- #
            # 加载图片
            imgs_hr, imgs_lr = self.data_loader.load_data(batch_size)
            fake_hr = self.generator.predict(imgs_lr)

            valid = np.ones((batch_size,) + self.disc_patch)
            fake = np.zeros((batch_size,) + self.disc_patch)

            d_loss_real = self.discriminator.train_on_batch(imgs_hr, valid)
            d_loss_fake = self.discriminator.train_on_batch(fake_hr, fake)
            d_loss = 0.5 * np.add(d_loss_real, d_loss_fake)

            # -------------------- #
            # 训练生成模型
            # -------------------- #
            # 重新加载图片,为了打乱顺序
            imgs_hr, imgs_lr = self.data_loader.load_data(batch_size)
            # 重新建立标签
            valid = np.ones((batch_size,) + self.disc_patch)
            image_features = self.vgg.predict(imgs_hr)

            g_loss = self.combined.train_on_batch(imgs_lr, [valid, image_features])
            print(d_loss, g_loss)
            elapsed_time = datetime.datetime.now() - start_time
            print("[Epoch %d%d] [D loss: %f, acc: %3d%%] [G loss: %05f, feature loss: %05f] time:%s"
                  % (epoch, epochs,
                     d_loss[0], 100 * d_loss[1],
                     g_loss[1],
                     g_loss[2],
                     elapsed_time))
            if epoch % sample_interval ==0:
                # 显示图片
                self.sample_images(epoch)
                # 保存图片
                if epoch %500 ==0 and epoch != init_epoch:
                    os.makedirs('weights/%s' % self.dataset_name, exist_ok=True)
                    self.generator.save_weights("weights/%s/gen_epoch%d.h5" % (self.dataset_name,epoch))
                    self.discriminator.save_weights("weights/%s/dis_epoch%d.h5" % (self.dataset_name,epoch))

    def sample_images(self, epoch):
        os.makedirs('images/%s' % self.dataset_name, exist_ok=True)
        r, c = 2, 2

        imgs_hr, imgs_lr = self.data_loader.load_data(batch_size=2, is_testing=True)
        fake_hr = self.generator.predict(imgs_lr)

        imgs_lr = 0.5*imgs_lr + 0.5
        fake_hr = 0.5*fake_hr + 0.5
        imgs_hr = 0.5*imgs_hr + 0.5
        titles = ['Generated', 'Original']
        fig, axs = plt.subplots(r, c)
        cnt = 0
        for row in range(r):
            for col,image in enumerate([fake_hr, imgs_hr]):
                axs[row, col].imshow(image[row])
                axs[row, col].set_title(titles[col])
                axs[row, col].axis('off')
            cnt += 1
        fig.savefig("images/%s/%d.png" % (self.dataset_name, epoch))
        plt.close()

        for i in range(r):
            fig = plt.figure()
            plt.imshow(imgs_lr[i])
            fig.savefig('images/%s/%d_lowers%d.png' % (self.dataset_name, epoch, i))
            plt.close()
            fig = plt.figure()
            plt.imshow(imgs_hr[i])
            fig.savefig('images/%s/%d_highs%d.png' % (self.dataset_name, epoch, i))
            plt.close()

if __name__ == '__main__':
    gan = SRGAN()
    gan.train(epochs=60000, init_epoch=1000, batch_size=1, sample_interval=50)

3.2.1 训练结果

训练9300次的结果,高频信息还是不清晰。
在这里插入图片描述

Tip:

  • 训练时间太长,训练一段时间保存,下次再从保存的位置训练;
  • VGG19输入参数需要调整。

四. 参考文献

  1. 【Super Resolution】超分辨率——SRGAN
  2. 好像还挺好玩的GAN8——SRGAN实现图像的分辨率提升
  3. PHOTO-REALISTIC SINGLE IMAGE SUPER-RESOLUTION USING SRGAN
  4. SRGAN-超分辨率图像复原
  5. 激活函数ReLU、Leaky ReLU、PReLU和RReLU
评论 1
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值