Global residual learning

This commit is contained in:
igv
2017-10-24 11:15:09 +03:00
parent cc7d2722d2
commit 1a8bb3d9fc
2 changed files with 50 additions and 8 deletions
+15 -4
View File
@@ -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):
+35 -4
View File
@@ -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')