mirror of
https://github.com/igv/FSRCNN-TensorFlow.git
synced 2026-08-18 00:57:28 +08:00
Global residual learning
This commit is contained in:
@@ -1,24 +1,26 @@
|
||||
import tensorflow as tf
|
||||
from utils import tf_ssim
|
||||
from utils import tf_ssim, bilinear_upsample_weights
|
||||
|
||||
class Model(object):
|
||||
|
||||
def __init__(self, config):
|
||||
self.name = "FSRCNN"
|
||||
# Different model layer counts and filter sizes for FSRCNN vs FSRCNN-s (fast), (d, s, m) in paper
|
||||
model_params = [[56, 12, 4], [32, 8, 1]]
|
||||
self.model_params = model_params[config.fast]
|
||||
model_params = [32, 0, 4, 1]
|
||||
self.GRL = True # global residual learning
|
||||
self.model_params = model_params
|
||||
self.scale = config.scale
|
||||
self.radius = config.radius
|
||||
self.padding = config.padding
|
||||
self.images = config.images
|
||||
self.batch = config.batch
|
||||
self.image_size = config.image_size - self.padding
|
||||
self.label_size = config.label_size
|
||||
self.c_dim = config.c_dim
|
||||
|
||||
def model(self):
|
||||
|
||||
d, s, m = self.model_params
|
||||
d, s, m, _ = self.model_params
|
||||
|
||||
# Feature Extraction
|
||||
size = self.padding + 1
|
||||
@@ -59,6 +61,15 @@ class Model(object):
|
||||
deconv = tf.nn.conv2d_transpose(conv, deconv_weights, output_shape=deconv_output, strides=deconv_stride, padding='SAME', data_format='NHWC')
|
||||
deconv = tf.nn.bias_add(deconv, deconv_biases, data_format='NHWC')
|
||||
|
||||
if self.GRL:
|
||||
# Deconvolution 2
|
||||
upsample_filter = bilinear_upsample_weights(self.scale, self.c_dim)
|
||||
self.biases['b{}'.format(m + 5)] = tf.get_variable('b{}'.format(m + 5), initializer=tf.constant(1, shape=[1]))
|
||||
deconv_output = [self.batch, self.label_size, self.label_size, self.c_dim]
|
||||
deconv_stride = [1, self.scale, self.scale, 1]
|
||||
img = tf.image.resize_image_with_crop_or_pad(self.images, self.image_size, self.image_size)
|
||||
deconv += tf.nn.conv2d_transpose(img, upsample_filter, output_shape=deconv_output, strides=deconv_stride, padding='SAME')
|
||||
|
||||
return deconv
|
||||
|
||||
def prelu(self, _x, i):
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
import sys
|
||||
from itertools import islice
|
||||
from utils import bilinear_upsample_weights
|
||||
import numpy as np
|
||||
|
||||
scale = 2
|
||||
radius = 1
|
||||
@@ -11,6 +13,7 @@ def get_line_number(phrase, file_name):
|
||||
for i, line in enumerate(f, 1):
|
||||
if phrase in line:
|
||||
return i
|
||||
return False
|
||||
|
||||
def read_weights(file_name, ln, size=1):
|
||||
content = []
|
||||
@@ -50,11 +53,13 @@ def header3(file, m, n, d):
|
||||
file.write('//!SAVE MODEL{}\n'.format((n//4)%(d//4) + 1 + (20 if m % 2 == 1 else 0)))
|
||||
file.write('//!COMPONENTS 4\n')
|
||||
|
||||
def header4(file, m, d):
|
||||
def header4(file, m, d, grl):
|
||||
base_header(file)
|
||||
file.write('//!WIDTH LUMA.w {} *\n'.format(scale))
|
||||
file.write('//!HEIGHT LUMA.h {} *\n'.format(scale))
|
||||
file.write('//!DESC aggregation\n')
|
||||
if grl:
|
||||
file.write('//!BIND HOOKED\n')
|
||||
for i in range(d//4):
|
||||
file.write('//!BIND MODEL{}\n'.format(i+1 + (20 if m % 2 == 1 else 0)))
|
||||
file.write('//!OFFSET -{}.0 -{}.0\n'.format(scale//2, scale//2))
|
||||
@@ -62,7 +67,7 @@ def header4(file, m, d):
|
||||
def main():
|
||||
if len(sys.argv) == 2:
|
||||
fname=sys.argv[1]
|
||||
d, s, m = [int(i) for i in fname[7:fname.index('.')].split("_")]
|
||||
d, s, m, _ = [int(i) for i in fname[7:fname.index('.')].split("_")]
|
||||
if s == 0:
|
||||
s = d
|
||||
dst = fname.replace("_", "-").replace("weights", "FSRCNNX_x{}_".format(scale)).replace("txt", "glsl")
|
||||
@@ -162,10 +167,11 @@ def main():
|
||||
# Aggregation
|
||||
ln = get_line_number("b{}".format(m + 4), fname)
|
||||
biases = read_weights(fname, ln)
|
||||
header4(file, m, d)
|
||||
grl = get_line_number("b{}".format(m + 5), fname)
|
||||
header4(file, m, d, grl)
|
||||
file.write('vec4 hook()\n')
|
||||
file.write('{\n')
|
||||
file.write('float res = {};\n'.format(biases[0]))
|
||||
file.write('float res = {};\n'.format(float(biases[0])))
|
||||
v = 1 + (20 if m % 2 == 1 else 0)
|
||||
file.write('vec2 fcoord = fract(MODEL{}_pos * MODEL{}_size);\n'.format(v, v))
|
||||
file.write('vec2 base = MODEL{}_pos + (vec2(0.5) - fcoord) * MODEL{}_pt;\n'.format(v, v))
|
||||
@@ -174,6 +180,31 @@ def main():
|
||||
for i in range(d//4-1):
|
||||
file.write('+MODEL{}_tex(base)'.format(i + 1 + v))
|
||||
file.write(')[index.y * {} + index.x];\n'.format(scale))
|
||||
|
||||
if grl:
|
||||
weights = bilinear_upsample_weights(scale, 1).ravel(order='F')
|
||||
x=list(reversed(range(scale)))
|
||||
xy = []
|
||||
for i in x:
|
||||
for j in x:
|
||||
xy.append([j, i])
|
||||
id = []
|
||||
for i in range(len(xy)):
|
||||
xi, yi = xy[i]
|
||||
for y in range(yi, scale * 2, scale):
|
||||
for x in range(xi, scale * 2, scale):
|
||||
id.append(y + x * (scale * 2))
|
||||
sort = [weights[id[l]] for l in range(0, len(id))]
|
||||
file.write('vec4 img = vec4(0);\n')
|
||||
for i in range(scale):
|
||||
for j in range(scale):
|
||||
idx = i * scale + j
|
||||
file.write('img[{}] =\n'.format(idx))
|
||||
file.write('{} * HOOKED_tex(base + HOOKED_pt * vec2(0,0)).r+\n'.format(sort[idx * 4]))
|
||||
file.write('{} * HOOKED_tex(base + HOOKED_pt * vec2(1,0)).r+\n'.format(sort[idx * 4 + 1]))
|
||||
file.write('{} * HOOKED_tex(base + HOOKED_pt * vec2(0,1)).r+\n'.format(sort[idx * 4 + 2]))
|
||||
file.write('{} * HOOKED_tex(base + HOOKED_pt * vec2(1,1)).r;\n'.format(sort[idx * 4 + 3]))
|
||||
file.write('res += img[index.y * {} + index.x];\n'.format(scale))
|
||||
file.write('return vec4(res, 0, 0, 1);\n')
|
||||
file.write('}\n')
|
||||
|
||||
|
||||
Reference in New Issue
Block a user