Support for recursive layers

This commit is contained in:
igv
2017-10-27 14:24:53 +03:00
parent 1a8bb3d9fc
commit bd0e4dc965
2 changed files with 50 additions and 46 deletions
+9 -6
View File
@@ -20,7 +20,7 @@ class Model(object):
def model(self):
d, s, m, _ = self.model_params
d, s, m, r = self.model_params
# Feature Extraction
size = self.padding + 1
@@ -39,11 +39,14 @@ class Model(object):
s = d
# Mapping (# mapping layers = m)
for i in range(3, m + 3):
weights = tf.get_variable('w{}'.format(i), shape=[3, 3, s, s], initializer=tf.variance_scaling_initializer(2))
biases = tf.get_variable('b{}'.format(i), initializer=tf.zeros([s]))
conv = tf.nn.conv2d(conv, weights, strides=[1,1,1,1], padding='SAME', data_format='NHWC')
conv = self.prelu(tf.nn.bias_add(conv, biases, data_format='NHWC'), i)
with tf.variable_scope("mapping_block") as scope:
for ri in range(r):
for i in range(3, m + 3):
weights = tf.get_variable('w{}'.format(i), shape=[3, 3, s, s], initializer=tf.variance_scaling_initializer(2))
biases = tf.get_variable('b{}'.format(i), initializer=tf.zeros([s]))
conv = tf.nn.conv2d(conv, weights, strides=[1,1,1,1], padding='SAME', data_format='NHWC')
conv = self.prelu(tf.nn.bias_add(conv, biases, data_format='NHWC'), i)
scope.reuse_variables()
# Expanding
if self.model_params[1] > 0:
+41 -40
View File
@@ -38,22 +38,22 @@ def header1(file, n, d):
file.write('//!SAVE MODEL{}\n'.format((n//4)%(d//4) + 1))
file.write('//!COMPONENTS 4\n')
def header2(file, w, n, s):
def header2(file, r, mi, m, n, s):
base_header(file)
file.write('//!DESC mapping {}_{}\n'.format(w+1, (n//4)%(s//4) + 1))
file.write('//!DESC mapping {}_{}\n'.format(mi + 1, (n//4)%(s//4) + 1))
for i in range(s//4):
file.write('//!BIND MODEL{}\n'.format(i+1 + (0 if w % 2 == 0 else 20)))
file.write('//!SAVE MODEL{}\n'.format((n//4)%(s//4) + 1 + (20 if w % 2 == 0 else 0)))
file.write('//!BIND MODEL{}\n'.format(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(file, m, n, d):
def header3(file, m, r, n, d):
base_header(file)
file.write('//!DESC sub-pixel convolution {}\n'.format((n//4)%(d//4) + 1))
file.write('//!BIND MODEL{}\n'.format((n//4)%(d//4) + 1 + (20 if m % 2 == 1 else 0)))
file.write('//!SAVE MODEL{}\n'.format((n//4)%(d//4) + 1 + (20 if m % 2 == 1 else 0)))
file.write('//!BIND MODEL{}\n'.format((n//4)%(d//4) + 1 + (20 if (r * m) % 2 == 1 else 0)))
file.write('//!SAVE MODEL{}\n'.format((n//4)%(d//4) + 1 + (20 if (r * m) % 2 == 1 else 0)))
file.write('//!COMPONENTS 4\n')
def header4(file, m, d, grl):
def header4(file, m, r, d, grl):
base_header(file)
file.write('//!WIDTH LUMA.w {} *\n'.format(scale))
file.write('//!HEIGHT LUMA.h {} *\n'.format(scale))
@@ -61,13 +61,13 @@ def header4(file, m, d, grl):
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('//!BIND MODEL{}\n'.format(i+1 + (20 if (r * m) % 2 == 1 else 0)))
file.write('//!OFFSET -{}.0 -{}.0\n'.format(scale//2, scale//2))
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, r = [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")
@@ -96,32 +96,33 @@ def main():
file.write('}\n\n')
# Mapping layers
for w in range(m):
ln = get_line_number("w{}".format(w + 3), fname)
weights = read_weights(fname, ln, s*9)
ln = get_line_number("b{}".format(w + 3), fname)
biases = read_weights(fname, ln)
ln = get_line_number("alpha{}".format(w + 3), fname)
alphas = read_weights(fname, ln)
for n in range(0, s, 4):
header2(file, w, n, s)
file.write('vec4 hook()\n')
file.write('{\n')
file.write('vec4 res = vec4({});\n'.format(",".join(biases[0].strip(",").split(",")[n:n+4])))
p = 0
for l in range(0, len(weights), 4):
if l % s == 0:
y, x = p%3-1, p//3-1
p += 1
file.write('res += mat4({},{},{},{}) * vec4(MODEL{}_texOff(vec2({},{})));\n'.format(
",".join(weights[l].strip(",").split(",")[n:n+4]),
",".join(weights[l+1].strip(",").split(",")[n:n+4]),
",".join(weights[l+2].strip(",").split(",")[n:n+4]),
",".join(weights[l+3].strip(",").split(",")[n:n+4]),
(l//4)%(s//4) + 1 + (20 if w % 2 == 1 else 0), x, y))
file.write('res = mix(res, vec4({}) * res, lessThan(res, vec4(0.0)));\n'.format(",".join(alphas[0].strip(",").split(",")[n:n+4])))
file.write('return res;\n')
file.write('}\n\n')
for ri in range(r):
for mi in range(m):
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)
ln = get_line_number("alpha{}".format(mi + 3), fname)
alphas = read_weights(fname, ln)
for n in range(0, s, 4):
header2(file, mi, n, s)
file.write('vec4 hook()\n')
file.write('{\n')
file.write('vec4 res = vec4({});\n'.format(",".join(biases[0].strip(",").split(",")[n:n+4])))
p = 0
for l in range(0, len(weights), 4):
if l % s == 0:
y, x = p%3-1, p//3-1
p += 1
file.write('res += mat4({},{},{},{}) * vec4(MODEL{}_texOff(vec2({},{})));\n'.format(
",".join(weights[l].strip(",").split(",")[n:n+4]),
",".join(weights[l+1].strip(",").split(",")[n:n+4]),
",".join(weights[l+2].strip(",").split(",")[n:n+4]),
",".join(weights[l+3].strip(",").split(",")[n:n+4]),
(l//4)%(s//4) + 1 + (20 if (ri * m + mi) % 2 == 1 else 0), x, y))
file.write('res = mix(res, vec4({}) * res, lessThan(res, vec4(0.0)));\n'.format(",".join(alphas[0].strip(",").split(",")[n:n+4])))
file.write('return res;\n')
file.write('}\n\n')
# Sub-pixel convolution
ln = get_line_number("w{}".format(m + 4), fname)
@@ -145,7 +146,7 @@ def main():
weights = list(reversed(weights))
sort = [weights[id[l]].strip(",") for l in range(0, len(id))]
for n in range(0, d, 4):
header3(file, m, n, d)
header3(file, m, r, n, d)
file.write('vec4 hook()\n')
file.write('{\n')
file.write('vec4 res = vec4(0);\n')
@@ -159,7 +160,7 @@ 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({}), MODEL{}_texOff(vec2({},{}))){}\n'.format(",".join(sort[l+total].strip(",").split(",")[n:n+4]),
(n//4)%(d//4) + 1 + (20 if m % 2 == 1 else 0), x, y, ';' if l == s1 * s2 - 1 else '+'))
(n//4)%(d//4) + 1 + (20 if (r * m) % 2 == 1 else 0), x, y, ';' if l == s1 * s2 - 1 else '+'))
total = total + l + 1
file.write('return res;\n')
file.write('}\n\n')
@@ -168,11 +169,11 @@ def main():
ln = get_line_number("b{}".format(m + 4), fname)
biases = read_weights(fname, ln)
grl = get_line_number("b{}".format(m + 5), fname)
header4(file, m, d, grl)
header4(file, m, r, d, grl)
file.write('vec4 hook()\n')
file.write('{\n')
file.write('float res = {};\n'.format(float(biases[0])))
v = 1 + (20 if m % 2 == 1 else 0)
v = 1 + (20 if (r * 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))
file.write('ivec2 index = ivec2(fcoord * vec2({}));\n'.format(scale))