Replace transposed convolution with sub-pixel convolution

This commit is contained in:
igv
2018-08-23 15:29:46 +03:00
parent 1b35ecb2cd
commit 09d7d07496
4 changed files with 50 additions and 117 deletions
+6 -7
View File
@@ -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
+7 -18
View File
@@ -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
+32 -91
View File
@@ -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')
+5 -1
View File
@@ -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()