From 5274a73ebe03976b9e2d041794963507ab576e72 Mon Sep 17 00:00:00 2001 From: igv Date: Sun, 26 Nov 2017 12:02:37 +0200 Subject: [PATCH] gen.py: support for shrinking/expanding layers --- gen.py | 97 +++++++++++++++++++++++++++++++++++++++++++++++----------- 1 file changed, 79 insertions(+), 18 deletions(-) diff --git a/gen.py b/gen.py index 861fcf8..9e12a34 100644 --- a/gen.py +++ b/gen.py @@ -41,24 +41,40 @@ def header1(file, n, d): file.write('//!SAVE FEATURE{}\n'.format((n//4)%(d//4) + 1)) file.write('//!COMPONENTS 4\n') -def header2(file, r, mi, m, n, s): +def header2(file, d, n, s): + base_header(file) + file.write('//!DESC shrinking {}\n'.format((n//4)%(s//4) + 1)) + for i in range(d//4): + file.write('//!BIND {}{}\n'.format("FEATURE", i + 1)) + file.write('//!SAVE SHRINKED{}\n'.format((n//4)%(s//4) + 1)) + file.write('//!COMPONENTS 4\n') + +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("FEATURE" if r == 0 and mi == 0 else "MODEL", i+1 + (0 if (r * m + mi) % 2 == 0 else 20))) + 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 FEATURE{}\n'.format((n//4)%(s//4) + 1)) + file.write('//!BIND {}{}\n'.format(inp, (n//4)%(s//4) + 1)) 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, r, n, d): +def header4(file, s, 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 (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('//!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('//!COMPONENTS 4\n') -def header4(file, m, r, d, grl): +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('//!COMPONENTS 4\n') + +def header6(file, m, r, d, inp, grl): base_header(file) file.write('//!WIDTH LUMA.w {} *\n'.format(scale)) file.write('//!HEIGHT LUMA.h {} *\n'.format(scale)) @@ -66,7 +82,7 @@ def header4(file, m, r, 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 (r * m) % 2 == 1 else 0))) + file.write('//!BIND {}{}\n'.format(inp, i+1 + (20 if (r * m) % 2 == 1 else 0))) file.write('//!OFFSET -{}.0 -{}.0\n'.format(scale//2, scale//2)) def main(): @@ -75,6 +91,9 @@ def main(): d, s, m, r = [int(i) for i in fname[7:fname.index('.')].split("_")] if s == 0: s = d + shrinking = False + else: + shrinking = True dst = fname.replace("_", "-").replace("weights", "FSRCNNX_x{}_".format(scale)).replace("txt", "glsl") with open(dst, 'w') as file: @@ -94,9 +113,29 @@ def main(): y, x = p%(feature_radius*2+1)-feature_radius, p//(feature_radius*2+1)-feature_radius p += 1 file.write('res += vec4({}) * float(LUMA_texOff(vec2({},{})));\n'.format(format_weights(weights[l], n), x, y)) + if shrinking: + ln = get_line_number("alpha1", 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: + # Shrinking layer + ln = get_line_number("w2", fname) + weights = read_weights(fname, ln, d) + ln = get_line_number("b2", fname) + biases = read_weights(fname, ln) + for n in range(0, s, 4): + header2(file, d, n, s) + 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, d, 4): + file.write('res += mat4({},{},{},{}) * FEATURE{}_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('return res;\n') + file.write('}\n\n') + # Mapping layers for ri in range(r): for mi in range(m): @@ -104,8 +143,9 @@ def main(): 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): - header2(file, ri, mi, m, n, s) + header3(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))) @@ -118,7 +158,7 @@ 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), - "FEATURE" if ri == 0 and mi == 0 else "MODEL", + 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) alphas = read_weights(fname, ln) @@ -131,6 +171,26 @@ def main(): file.write('return res;\n') file.write('}\n\n') + if shrinking: + # Expanding layer + ln = get_line_number("w{}".format(m + 3), fname) + weights = read_weights(fname, ln, d) + ln = get_line_number("b{}".format(m + 3), fname) + biases = read_weights(fname, ln) + ln = get_line_number("alpha{}".format(m + 3), fname) + alphas = read_weights(fname, ln) + for n in range(0, d, 4): + header4(file, s, m, r, n, d) + 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(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 = 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) weights = read_weights(fname, ln, dsize**2) @@ -152,8 +212,9 @@ def main(): weights = list(reversed(weights)) sort = [weights[id[l]].strip(",") for l in range(0, len(id))] + inp = "EXPANDED" if shrinking else "MODEL" for n in range(0, d, 4): - header3(file, m, r, n, d) + header5(file, m, r, n, d, inp) file.write('vec4 hook()\n') file.write('{\n') file.write('vec4 res = vec4(0);\n') @@ -166,7 +227,7 @@ def main(): 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({}), MODEL{}_texOff(vec2({},{}))){}\n'.format(format_weights(sort[l+total], n), + 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 '+')) total = total + l + 1 file.write('return res;\n') @@ -176,17 +237,17 @@ 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, r, d, grl) + 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(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('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('ivec2 index = ivec2(fcoord * vec2({}));\n'.format(scale)) - file.write('res += (MODEL{}_tex(base)'.format(v)) + file.write('res += ({}{}_tex(base)'.format(inp, v)) for i in range(d//4-1): - file.write('+MODEL{}_tex(base)'.format(i + 1 + v)) + file.write('+{}{}_tex(base)'.format(inp, i + 1 + v)) file.write(')[index.y * {} + index.x];\n'.format(scale)) if grl: