diff --git a/FSRCNN.py b/FSRCNN.py index 06f89ff..41bf932 100644 --- a/FSRCNN.py +++ b/FSRCNN.py @@ -51,22 +51,27 @@ class Model(object): conv = tf.nn.conv2d(conv, weights, strides=[1,1,1,1], padding='SAME', data_format='NHWC') conv = tf.nn.bias_add(conv, biases, data_format='NHWC') if i == m + 2: + conv = self.prelu(conv, m + 3) + weights = tf.get_variable('w{}'.format(m + 3), shape=[1, 1, s, s], initializer=tf.variance_scaling_initializer(2)) + biases = tf.get_variable('b{}'.format(m + 3), initializer=tf.zeros([s])) + conv = tf.nn.conv2d(conv, weights, strides=[1,1,1,1], padding='SAME', data_format='NHWC') + conv = tf.nn.bias_add(conv, biases, data_format='NHWC') conv = tf.add(conv, features) scope.reuse_variables() conv = self.prelu(conv, 2) # Expanding if self.model_params[1] > 0: - expand_weights = tf.get_variable('w{}'.format(m + 3), shape=[1, 1, s, d], initializer=tf.variance_scaling_initializer(2)) - expand_biases = tf.get_variable('b{}'.format(m + 3), initializer=tf.zeros([d])) + expand_weights = tf.get_variable('w{}'.format(m + 4), shape=[1, 1, s, d], initializer=tf.variance_scaling_initializer(2)) + expand_biases = tf.get_variable('b{}'.format(m + 4), initializer=tf.zeros([d])) conv = tf.nn.conv2d(conv, expand_weights, strides=[1,1,1,1], padding='SAME', data_format='NHWC') conv = tf.nn.bias_add(conv, expand_biases, data_format='NHWC') - conv = self.prelu(conv, m + 3) + conv = self.prelu(conv, m + 4) # Deconvolution deconv_size = self.radius * self.scale * 2 + 1 - deconv_weights = tf.get_variable('w{}'.format(m + 4), shape=[deconv_size, deconv_size, 1, d], initializer=tf.variance_scaling_initializer(0.01)) - deconv_biases = tf.get_variable('b{}'.format(m + 4), initializer=tf.zeros([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') @@ -75,7 +80,7 @@ class Model(object): 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])) + 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) diff --git a/gen.py b/gen.py index 9e12a34..769bc29 100644 --- a/gen.py +++ b/gen.py @@ -53,25 +53,32 @@ def header3(file, r, mi, m, n, s, inp): base_header(file) file.write('//!DESC mapping {}_{}\n'.format(mi + 1, (n//4)%(s//4) + 1)) for i in range(s//4): - file.write('//!BIND {}{}\n'.format(inp if r == 0 and mi == 0 else "MODEL", i+1 + (0 if (r * m + mi) % 2 == 0 else 20))) - if mi == m-1 and (mi+1)*(r+1) > 1: - file.write('//!BIND {}{}\n'.format(inp, (n//4)%(s//4) + 1)) + file.write('//!BIND {}{}\n'.format(inp, i+1 + (0 if (r * m + mi) % 2 == 0 else 20))) file.write('//!SAVE MODEL{}\n'.format((n//4)%(s//4) + 1 + (20 if (r * m + mi) % 2 == 0 else 0))) file.write('//!COMPONENTS 4\n') +def header3_1(file, r, mi, m, n, s, inp): + base_header(file) + file.write('//!DESC sub-band residuals {}\n'.format((n//4)%(s//4) + 1)) + for i in range(s//4): + file.write('//!BIND MODEL{}\n'.format(i + 1 + (20 if (r * m + mi) % 2 == 0 else 0))) + file.write('//!BIND {}{}\n'.format(inp, (n//4)%(s//4) + 1)) + file.write('//!SAVE RES{}\n'.format((n//4)%(s//4) + 1)) + file.write('//!COMPONENTS 4\n') + def header4(file, s, m, r, n, d): base_header(file) file.write('//!DESC expanding {}\n'.format((n//4)%(d//4) + 1)) for i in range(s//4): - file.write('//!BIND MODEL{}\n'.format(i + 1 + (20 if (r * m) % 2 == 1 else 0))) - file.write('//!SAVE EXPANDED{}\n'.format((n//4)%(d//4) + 1 + (20 if (r * m) % 2 == 1 else 0))) + file.write('//!BIND RES{}\n'.format(i + 1)) + file.write('//!SAVE EXPANDED{}\n'.format((n//4)%(d//4) + 1)) file.write('//!COMPONENTS 4\n') def header5(file, m, r, n, 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 + (20 if (r * m) % 2 == 1 else 0))) - file.write('//!SAVE {}{}\n'.format(inp, (n//4)%(d//4) + 1 + (20 if (r * m) % 2 == 1 else 0))) + file.write('//!BIND {}{}\n'.format(inp, (n//4)%(d//4) + 1)) + file.write('//!SAVE {}{}\n'.format(inp, (n//4)%(d//4) + 1)) file.write('//!COMPONENTS 4\n') def header6(file, m, r, d, inp, grl): @@ -82,7 +89,7 @@ def header6(file, m, r, d, inp, grl): if grl: file.write('//!BIND HOOKED\n') for i in range(d//4): - file.write('//!BIND {}{}\n'.format(inp, i+1 + (20 if (r * m) % 2 == 1 else 0))) + file.write('//!BIND {}{}\n'.format(inp, i + 1)) file.write('//!OFFSET -{}.0 -{}.0\n'.format(scale//2, scale//2)) def main(): @@ -137,15 +144,16 @@ def main(): file.write('}\n\n') # Mapping layers + inp = "SHRINKED" if shrinking else "FEATURE" for ri in range(r): for mi in range(m): + tex_name = inp if ri == 0 and mi == 0 else "RES" if ri > 0 and mi == 0 else "MODEL" ln = get_line_number("w{}".format(mi + 3), fname) weights = read_weights(fname, ln, s*9) ln = get_line_number("b{}".format(mi + 3), fname) biases = read_weights(fname, ln) - inp = "SHRINKED" if shrinking else "FEATURE" for n in range(0, s, 4): - header3(file, ri, mi, m, n, s, inp) + header3(file, ri, mi, m, n, s, tex_name) file.write('vec4 hook()\n') file.write('{\n') file.write('vec4 res = vec4({});\n'.format(format_weights(biases[0], n))) @@ -158,26 +166,43 @@ def main(): 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 if ri == 0 and mi == 0 else "MODEL", - idx + 1 + (20 if (ri * m + mi) % 2 == 1 else 0), x, y)) - ln = get_line_number("alpha{}".format(3 if mi == m - 1 else mi + 4), fname) + tex_name, idx + 1 + (20 if (ri * m + mi) % 2 == 1 else 0), x, y)) + ln = get_line_number("alpha{}".format(m + 3 if mi == m - 1 else mi + 4), fname) alphas = read_weights(fname, ln) - if mi == m - 1: - file.write('res += {}{}_texOff(0);\n'.format(inp, (n//4)%(s//4) + 1)) - if ri == r - 1: - ln = get_line_number("alpha2", fname) - alphas = read_weights(fname, ln) file.write('res = max(res, vec4(0.0)) + vec4({}) * min(res, vec4(0.0));\n'.format(format_weights(alphas[0], n))) file.write('return res;\n') file.write('}\n\n') + if mi == m - 1: + ln = get_line_number("w{}".format(m + 3), fname) + weights = read_weights(fname, ln, s*(mi+2)) + ln = get_line_number("b{}".format(m + 3), fname) + biases = read_weights(fname, ln) + for n in range(0, s, 4): + header3_1(file, ri, mi, m, n, s, inp) + file.write('vec4 hook()\n') + file.write('{\n') + file.write('vec4 res = vec4({});\n'.format(format_weights(biases[0], n))) + for l in range(0, s, 4): + file.write('res += mat4({},{},{},{}) * MODEL{}_texOff(0);\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), + l//4 + 1 + (20 if (ri * m + mi) % 2 == 0 else 0))) + file.write('res += {}{}_texOff(0);\n'.format(inp, (n//4)%(s//4) + 1)) + if ri == r - 1: + ln = get_line_number("alpha2", fname) + alphas = read_weights(fname, ln) + file.write('res = max(res, vec4(0.0)) + vec4({}) * min(res, vec4(0.0));\n'.format(format_weights(alphas[0], n))) + file.write('return res;\n') + file.write('}\n\n') + if shrinking: # Expanding layer - ln = get_line_number("w{}".format(m + 3), fname) + ln = get_line_number("w{}".format(m + 4), fname) weights = read_weights(fname, ln, d) - ln = get_line_number("b{}".format(m + 3), fname) + ln = get_line_number("b{}".format(m + 4), fname) biases = read_weights(fname, ln) - ln = get_line_number("alpha{}".format(m + 3), fname) + ln = get_line_number("alpha{}".format(m + 4), fname) alphas = read_weights(fname, ln) for n in range(0, d, 4): header4(file, s, m, r, n, d) @@ -185,14 +210,14 @@ def main(): file.write('{\n') file.write('vec4 res = vec4({});\n'.format(format_weights(biases[0], n))) for l in range(0, s, 4): - file.write('res += mat4({},{},{},{}) * MODEL{}_texOff(vec2(0.0));\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), - l//4 + 1 + (20 if (r * m) % 2 == 1 else 0))) + file.write('res += mat4({},{},{},{}) * RES{}_texOff(vec2(0.0));\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), + l//4 + 1)) file.write('res = max(res, vec4(0.0)) + vec4({}) * min(res, vec4(0.0));\n'.format(format_weights(alphas[0], n))) file.write('return res;\n') file.write('}\n\n') # Sub-pixel convolution - ln = get_line_number("w{}".format(m + 4), fname) + ln = get_line_number("w{}".format(m + 5), fname) weights = read_weights(fname, ln, dsize**2) x=list(reversed(range(scale))) @@ -212,7 +237,7 @@ def main(): weights = list(reversed(weights)) sort = [weights[id[l]].strip(",") for l in range(0, len(id))] - inp = "EXPANDED" if shrinking else "MODEL" + 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') @@ -228,26 +253,27 @@ def main(): 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 + (20 if (r * m) % 2 == 1 else 0), x, y, ';' if l == s1 * s2 - 1 else '+')) + (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 + 4), fname) + ln = get_line_number("b{}".format(m + 5), fname) biases = read_weights(fname, ln) - grl = get_line_number("b{}".format(m + 5), fname) + grl = get_line_number("b{}".format(m + 6), fname) header6(file, m, r, d, inp, grl) file.write('vec4 hook()\n') file.write('{\n') file.write('float res = {};\n'.format(float(biases[0]))) - v = 1 + (20 if (r * m) % 2 == 1 else 0) - file.write('vec2 fcoord = fract({}{}_pos * {}{}_size);\n'.format(inp, v, inp, v)) - file.write('vec2 base = {}{}_pos + (vec2(0.5) - fcoord) * {}{}_pt;\n'.format(inp, v, inp, v)) + 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 += ({}{}_tex(base)'.format(inp, v)) - for i in range(d//4-1): - file.write('+{}{}_tex(base)'.format(inp, i + 1 + v)) + 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: