diff --git a/FSRCNN.py b/FSRCNN.py index 740df7e..3c0bce2 100644 --- a/FSRCNN.py +++ b/FSRCNN.py @@ -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): diff --git a/gen.py b/gen.py index abcd1b1..9fe8fbb 100644 --- a/gen.py +++ b/gen.py @@ -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')