一. SRGAN介绍

1.1 相关知识介绍
- 相关定义:SR(Super Resolution) LR(Low Resolution) HR(High Resolution) MOS(mean opinion score)
- 提出新评价标准:
- 在SR中PSNR标准应当换掉,因为高的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) MOS(Meanopinionscore)
- 提出新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}为什么超分中丢失的是高频信息
为什么超分中丢失的是高频信息?
我们可以这样考虑,超分问题的本质是通过不同的上采样的方式从一个低分辨率图像恢复到高分辨率图像,从像素级别的角度来看,这是一个一对多或者多对多的问题,那么我们就可以认为这是一个回归问题。既然是回归问题,在拟合的过程中要保证尽量多的信息可以恢复准确,而在图像中,低频信息占大多数,而高频信息占少数,所以在超分问题中高频信息就丢失了。
- 用不同的损失来判定的结果:

- 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,可以更好的抓取不变特征。
1.2 SRGAN网络介绍
1.2.1 生成模型(Generater)
生成网络的构成如下图所示:

1.2.1.1 基础性理解
从左至右来看,SRGAN的生成网络由三个部分组成。
- 低分辨率图像进入后会经过一个卷积+RELU函数
- 然后经过B个残差网络结构,每个残差网络内部包含两个卷积+标准化+RELU,还有一个残差边。
- 然后进入上采样部分,将长宽进行放大,两次上采样后,变为原来的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+2sub−pixelconv】∗n
- 生成器是在
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 框架介绍
- SRGAN的判别网络由不断重复的 卷积+LeakyRELU和标准化 组成。
- 具体地说:
- 判别器: 【 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+10−3lGenSR(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
- MSE
该MSE丢失内容是 比 较 I H R 与 生 成 的 图 像 G θ G ( I L R ) , 并 采 取 了 M S E 这 种 差 异 \color{red}比较I^{HR}与生成的图像G_{\theta_G}(I^{LR}),并采取了MSE这种差异 比较IHR与生成的图像Gθ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:L;b1:L}:L层网的权重和偏差.
- 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输入参数需要调整。

3万+

被折叠的 条评论
为什么被折叠?



