GAN生成式对抗网络(四)——SRGAN超高分辨率图片重构

论坛 期权论坛     
选择匿名的用户   2021-5-30 11:16   184   0
<div class="blogpost-body cnblogs-markdown" id="cnblogs_post_body">
<p>论文pdf 地址:<a class="uri" href="https://arxiv.org/pdf/1609.04802v1.pdf">https://arxiv.org/pdf/1609.04802v1.pdf</a></p>
<h2 id="我的实际效果">我的实际效果</h2>
<p>清晰度距离我的期待有距离。<br> 颜色上面存在差距。<br> 解决想法<br> 增加一个颜色判别器。将颜色值反馈给生成器</p>
<p><img alt="1545753-20181128120704364-961028862.png" src="https://beijingoptbbs.oss-cn-beijing.aliyuncs.com/cs/5606289-3e634975af79935695cde588597a9284.png"></p>
<p>srgan论文是建立在gan基础上的,利用gan生成式对抗网络,将图片重构为高清分辨率的图片。<br> github上有开源的srgan项目。由于开源者,开发时考虑的问题更丰富,技巧更为高明,导致其代码都比较难以阅读和理解。<br> 在为了充分理解这个论文。这里结合论文,开源代码,和自己的理解重新写了个srgan高清分辨率模型。</p>
<h2 id="gan原理">GAN原理</h2>
<p>在一个不断提高判断能力的判断器的持续反馈下,不断改善生成器的生成参数,直到生成器生成的结果能够通过判断器的判断。(见本博客其他文章)</p>
<h2 id="srgan用到的模块及其关系">SRGAN用到的模块,及其关系</h2>
<p>损失值,根据的这个关系结构计算的。<br><img alt="1545753-20181127163533885-1386223271.png" src="https://beijingoptbbs.oss-cn-beijing.aliyuncs.com/cs/5606289-767187abb678c11db3cfa04f915a2d0c.png"><br> 注意:vgg19是使用已经训练好的模型,这里只是拿来提取特征使用,</p>
<p>对于生成器,根据三个运算结果数据,进行随机梯度的优化调整<br> ①判定器生成数据的鉴定结果<br> ②vgg19的特征比较情况<br> ③生成图形与理想图形的mse差距</p>
<h2 id="论文中生成器和判别器的模型图">论文中,生成器和判别器的模型图</h2>
<p><img alt="1545753-20181127170109742-1985386475.png" src="https://beijingoptbbs.oss-cn-beijing.aliyuncs.com/cs/5606289-fe0850aa29fd3f55defe95fa8cd80537.png"><br> 生成器结构为:一层卷积,16层残差卷积,再将第一层卷积结果&#43;16层残差结,卷积&#43;2倍反卷积,卷积&#43;2倍反卷积,tanh缩放,产生生成结果。<br> 判别器结构为:8层卷积&#43;reshape,全连接。(论文中,用了两层。我这里只用了一层全连接,参数量太大,我6G 的gpu内存不够用)<br> vgg19结构:在vgg19的第四层,返回获取到的特征结果,进行MSE对比<br> 注意:BN处理,leaky relu等等处理技巧</p>
<h2 id="代码解释">代码解释</h2>
<pre class="blockcode"><code>import numpy as np
import os
import tensorlayer as tl
import tensorflow as tf

#获取vgg9.npy中vgg19的参数,
vgg19_npy_path &#61; &#34;./vgg19.npy&#34;
if not os.path.isfile(vgg19_npy_path):
    print(&#34;Please download vgg19.npz from : https://github.com/machrisaa/tensorflow-vgg&#34;)
    exit()
npz &#61; np.load(vgg19_npy_path, encoding&#61;&#39;latin1&#39;).item()
w_params &#61; []
b_params &#61; []
for val in sorted(npz.items()):
    W &#61; np.asarray(val[1][0])
    b &#61; np.asarray(val[1][1])
    # print(&#34;  Loading %s: %s, %s&#34; % (val[0], W.shape, b.shape))
    w_params.append(W, )
    b_params.extend(b)


#tensorlayer加载图片时,用于处理图片。随机获取图片中 192*192的矩阵, 内存不足时,可以优化这里
def crop_sub_imgs_fn(x, is_random&#61;True):
    x &#61; tl.prepro.crop(x, wrg&#61;192, hrg&#61;192, is_random&#61;is_random)
    x &#61; x / (255. / 2.)
    x &#61; x - 1.
    return x
#resize矩阵 内存不足时,可以优化这里
def downsample_fn(x):
    x &#61; tl.prepro.imresize(x, size&#61;[48, 48], interp&#61;&#39;bicubic&#39;, mode&#61;None)
    x &#61; x / (255. / 2.)
    x &#61; x - 1.
    return x

# 参数
config &#61; {
    &#34;epoch&#34;: 5,
}

# 内存不够时,可以减小这个
batch_size &#61; 10


class SRGAN(object):
    def __init__(self):
        # with tf.device(&#39;/gpu:0&#39;):
        #占位变量,存储需要重构的图片
        self.x &#61; tf.placeholder(tf.float32, shape&#61;[batch_size, 48, 48, 3], name&#61;&#39;train_bechanged&#39;)
        #占位变量,存储需要学习的理想中的图片
        self.y &#61; tf.placeholder(tf.float32, shape&#61;[batch_size, 192, 192, 3], name&#61;&#39;train_target&#39;)
        self.init_fake_y &#61; self.generator(self.x)  # 预训练时生成的假照片
        self.fake_y &#61; self.generator(self.x, reuse&#61;True)  # 全部训练时生成的假照片

         #占位变量,存储需要重构的测试图片
        self.test_x &#61; tf.placeholder(tf.float32, shape&#61;[1, None, None, 3], name&#61;&#39;test_generator&#39;)
        #占位变量,存储重构后的测试图片
        self.test_fake_y &#61; self.generator(self.test_x, reuse&#61;True)  # 生成的假照片

        #占位变量,将生成图片resize
        self.fake_y_vgg &#61; tf.image.resize_images(
            self.fake_y, size&#61;[224, 224], method&#61;
分享到 :
0 人收藏
您需要登录后才可以回帖 登录 | 立即注册

本版积分规则

积分:3875789
帖子:775174
精华:0
期权论坛 期权论坛
发布
内容

下载期权论坛手机APP