mirror of
https://github.com/igv/FSRCNN-TensorFlow.git
synced 2026-10-06 17:04:12 +08:00
Replace transposed convolution with sub-pixel convolution
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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')
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user