From 09d7d07496e9df794c4d32a19af9e6c03b6067b6 Mon Sep 17 00:00:00 2001 From: igv Date: Thu, 23 Aug 2018 15:29:46 +0300 Subject: [PATCH] Replace transposed convolution with sub-pixel convolution --- ESPCN.py | 13 +++--- FSRCNN.py | 25 ++++------- gen.py | 123 ++++++++++++++---------------------------------------- model.py | 6 ++- 4 files changed, 50 insertions(+), 117 deletions(-) diff --git a/ESPCN.py b/ESPCN.py index 509ba0e..f0273fe 100644 --- a/ESPCN.py +++ b/ESPCN.py @@ -34,14 +34,13 @@ class Model(object): conv = tf.nn.bias_add(conv, biases, data_format='NHWC') conv = self.prelu(conv, i) - # Deconvolution - deconv_size = self.radius * self.scale * 2 + 1 - deconv_weights = tf.get_variable('w{}'.format(m+1), shape=[deconv_size, deconv_size, 1, d[-1]], initializer=tf.variance_scaling_initializer(0.01)) - deconv_biases = tf.get_variable('b{}'.format(m+1), initializer=tf.zeros([1])) - deconv_output = [self.batch, self.label_size, self.label_size, self.c_dim] - deconv_stride = [1, self.scale, self.scale, 1] - deconv = tf.nn.conv2d_transpose(conv, deconv_weights, output_shape=deconv_output, strides=deconv_stride, padding='SAME', data_format='NHWC') + # Sub-pixel convolution + size = self.radius * 2 + 1 + deconv_weights = tf.get_variable('deconv_w', shape=[size, size, d[-1], self.scale**2], initializer=tf.variance_scaling_initializer(0.01)) + deconv_biases = tf.get_variable('deconv_b', initializer=tf.zeros([self.scale**2])) + deconv = tf.nn.conv2d(conv, deconv_weights, strides=[1,1,1,1], padding='SAME', data_format='NHWC') deconv = tf.nn.bias_add(deconv, deconv_biases, data_format='NHWC') + deconv = tf.depth_to_space(deconv, self.scale, name='pixel_shuffle', data_format='NHWC') return deconv diff --git a/FSRCNN.py b/FSRCNN.py index 41bf932..7406ca2 100644 --- a/FSRCNN.py +++ b/FSRCNN.py @@ -1,5 +1,5 @@ import tensorflow as tf -from utils import tf_ssim, bilinear_upsample_weights +from utils import tf_ssim class Model(object): @@ -7,7 +7,6 @@ class Model(object): self.name = "FSRCNN" # Different model layer counts and filter sizes for FSRCNN vs FSRCNN-s (fast), (d, s, m) in paper 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 @@ -68,23 +67,13 @@ class Model(object): conv = tf.nn.bias_add(conv, expand_biases, data_format='NHWC') conv = self.prelu(conv, m + 4) - # Deconvolution - deconv_size = self.radius * self.scale * 2 + 1 - deconv_weights = tf.get_variable('w{}'.format(m + 5), shape=[deconv_size, deconv_size, 1, d], initializer=tf.variance_scaling_initializer(0.01)) - deconv_biases = tf.get_variable('b{}'.format(m + 5), initializer=tf.zeros([1])) - deconv_output = [self.batch, self.label_size, self.label_size, self.c_dim] - deconv_stride = [1, self.scale, self.scale, 1] - deconv = tf.nn.conv2d_transpose(conv, deconv_weights, output_shape=deconv_output, strides=deconv_stride, padding='SAME', data_format='NHWC') + # Sub-pixel convolution + size = self.radius * 2 + 1 + deconv_weights = tf.get_variable('deconv_w', shape=[size, size, d, self.scale**2], initializer=tf.variance_scaling_initializer(0.01)) + deconv_biases = tf.get_variable('deconv_b', initializer=tf.zeros([self.scale**2])) + deconv = tf.nn.conv2d(conv, deconv_weights, strides=[1,1,1,1], 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 + 6)] = tf.get_variable('b{}'.format(m + 6), 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') + deconv = tf.depth_to_space(deconv, self.scale, name='pixel_shuffle', data_format='NHWC') return deconv diff --git a/gen.py b/gen.py index 769bc29..c6aeb49 100644 --- a/gen.py +++ b/gen.py @@ -1,13 +1,9 @@ import sys from itertools import islice -from utils import bilinear_upsample_weights -import numpy as np scale = 2 radius = 1 -dsize = radius * scale * 2 + 1 - def get_line_number(phrase, file_name): with open(file_name) as f: for i, line in enumerate(f, 1): @@ -74,23 +70,20 @@ def header4(file, s, m, r, n, d): file.write('//!SAVE EXPANDED{}\n'.format((n//4)%(d//4) + 1)) file.write('//!COMPONENTS 4\n') -def header5(file, m, r, n, d, inp): +def header5(file, d, inp): base_header(file) - file.write('//!DESC sub-pixel convolution {}\n'.format((n//4)%(d//4) + 1)) - file.write('//!BIND {}{}\n'.format(inp, (n//4)%(d//4) + 1)) - file.write('//!SAVE {}{}\n'.format(inp, (n//4)%(d//4) + 1)) + file.write('//!DESC sub-pixel convolution\n') + for i in range(d//4): + file.write('//!BIND {}{}\n'.format(inp, i + 1)) + file.write('//!SAVE {}1\n'.format(inp)) file.write('//!COMPONENTS 4\n') -def header6(file, m, r, d, inp, grl): +def header6(file, inp): 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 {}{}\n'.format(inp, i + 1)) - file.write('//!OFFSET -{}.0 -{}.0\n'.format(scale//2, scale//2)) + file.write('//!BIND {}1\n'.format(inp)) def main(): if len(sys.argv) == 2: @@ -217,89 +210,37 @@ def main(): file.write('}\n\n') # Sub-pixel convolution - ln = get_line_number("w{}".format(m + 5), fname) - weights = read_weights(fname, ln, dsize**2) - - x=list(reversed(range(scale))) - if dsize % 2 == 1: - x=x[-1:]+x[:-1] - xy = [] - for i in x: - for j in x: - xy.append([j, i]) - - id = [] - for i in range(0, len(xy)): - xi, yi = xy[i] - for y in range(yi, dsize, scale): - for x in range(xi, dsize, scale): - id.append(y + x * dsize) - - weights = list(reversed(weights)) - sort = [weights[id[l]].strip(",") for l in range(0, len(id))] - inp = "EXPANDED" if shrinking else "RES" - for n in range(0, d, 4): - header5(file, m, r, n, d, inp) - file.write('vec4 hook()\n') - file.write('{\n') - file.write('vec4 res = vec4(0);\n') - total = 0 - for i in range(scale): - for j in range(scale): - file.write('res[{}] +=\n'.format(i * scale + j)) - s2 = radius*2+1 if i == 0 and dsize % 2 == 1 else radius*2 - for yi, y in enumerate(range(-radius + (0 if i == 0 and dsize % 2 == 1 else 1), radius + 1)): - s1 = radius*2+1 if j == 0 and dsize % 2 == 1 else radius*2 - for xi, x in enumerate(range(-radius + (0 if j == 0 and dsize % 2 == 1 else 1), radius + 1)): - l = yi * s1 + xi - file.write('dot(vec4({}), {}{}_texOff(vec2({},{}))){}\n'.format(format_weights(sort[l+total], n), inp, - (n//4)%(d//4) + 1, x, y, ';' if l == s1 * s2 - 1 else '+')) - total = total + l + 1 - file.write('return res;\n') - file.write('}\n\n') - - # Aggregation - ln = get_line_number("b{}".format(m + 5), fname) + ln = get_line_number("deconv_w", fname) + weights = read_weights(fname, ln, d*(radius*2+1)**2) + ln = get_line_number("deconv_b", fname) biases = read_weights(fname, ln) - grl = get_line_number("b{}".format(m + 6), fname) - header6(file, m, r, d, inp, grl) + inp = "EXPANDED" if shrinking else "RES" + header5(file, d, inp) + file.write('vec4 hook()\n') + file.write('{\n') + file.write('vec4 res = vec4({});\n'.format(biases[0])) + for n in range(0, scale**2, 4): + p = 0 + for l in range(0, len(weights), 4): + if l % d == 0: + y, x = p%(radius*2+1)-radius, p//(radius*2+1)-radius + p += 1 + idx = (l//4)%(d//4) + file.write('res += mat4({},{},{},{}) * {}{}_texOff(vec2({},{}));\n'.format( + format_weights(weights[l], n), format_weights(weights[l+1], n), + format_weights(weights[l+2], n), format_weights(weights[l+3], n), + inp, idx + 1, x, y)) + file.write('return res;\n') + file.write('}\n\n') + + # Aggregation + header6(file, inp) file.write('vec4 hook()\n') file.write('{\n') - file.write('float res = {};\n'.format(float(biases[0]))) file.write('vec2 fcoord = fract({}1_pos * {}1_size);\n'.format(inp, inp)) file.write('vec2 base = {}1_pos + (vec2(0.5) - fcoord) * {}1_pt;\n'.format(inp, inp)) file.write('ivec2 index = ivec2(fcoord * vec2({}));\n'.format(scale)) - file.write('res += (') - for i in range(d//4): - if i > 0: - file.write('+') - file.write('{}{}_tex(base)'.format(inp, i + 1)) - 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('float res = {}1_tex(base)[index.x * {} + index.y];\n'.format(inp, scale)) file.write('return vec4(res, 0, 0, 1);\n') file.write('}\n') diff --git a/model.py b/model.py index ebfcf25..5873636 100644 --- a/model.py +++ b/model.py @@ -69,7 +69,11 @@ class Model(object): self.saver = tf.train.Saver() def run(self): - self.train_op = tf.train.AdamOptimizer(self.learning_rate).minimize(self.loss) + global_step = tf.Variable(0, trainable=False) + optimizer = tf.train.AdamOptimizer(self.learning_rate) + deconv_mult = lambda grads: list(map(lambda x: (x[0] * 1.0, x[1]) if 'deconv' in x[1].name else x, grads)) + grads = deconv_mult(optimizer.compute_gradients(self.loss)) + self.train_op = optimizer.apply_gradients(grads, global_step=global_step) tf.global_variables_initializer().run()