ShkShahid/Auto-encoder_For_Image_Reconstruction
1
1 2 3 4import tensorflow as tf5import tensorlayer as tl6import numpy as np7 8# The HDR reconstruction autoencoder fully convolutional neural network9def model(x, batch_size=1, is_training=False):10 11 # Encoder network (VGG16, until pool5)12 x_in = tf.scalar_mul(255.0, x)13 net_in = tl.layers.InputLayer(x_in, name='input_layer')14 conv_layers, skip_layers = encoder(net_in)15 16 # Fully convolutional layers on top of VGG16 conv layers17 network = tl.layers.Conv2dLayer(conv_layers,18 act = tf.identity,19 shape = [3, 3, 512, 512],20 strides = [1, 1, 1, 1],21 padding='SAME',22 name ='encoder/h6/conv')23 network = tl.layers.BatchNormLayer(network, is_train=is_training, name='encoder/h6/batch_norm')24 network.outputs = tf.nn.relu(network.outputs, name='encoder/h6/relu')25 26 # Decoder network27 network = decoder(network, skip_layers, batch_size, is_training)28 29 if is_training:30 return network, conv_layers31 32 return network33 34 35# Final prediction of the model, including blending with input36def get_final(network, x_in):37 sb, sy, sx, sf = x_in.get_shape().as_list()38 y_predict = network.outputs39 40 # Highlight mask41 thr = 0.0542 alpha = tf.reduce_max(x_in, reduction_indices=[3])43 alpha = tf.minimum(1.0, tf.maximum(0.0, alpha-1.0+thr)/thr)44 alpha = tf.reshape(alpha, [-1, sy, sx, 1])45 alpha = tf.tile(alpha, [1,1,1,3])46 47 # Linearied input and prediction48 x_lin = tf.pow(x_in, 2.0)49 y_predict = tf.exp(y_predict)-1.0/255.050 51 # Alpha blending52 y_final = (1-alpha)*x_lin + alpha*y_predict53 54 return y_final55 56 57# Convolutional layers of the VGG16 model used as encoder network58def encoder(input_layer):59 60 VGG_MEAN = [103.939, 116.779, 123.68]61 62 # Convert RGB to BGR63 red, green, blue = tf.split(input_layer.outputs, 3, 3)64 bgr = tf.concat([ blue - VGG_MEAN[0], green - VGG_MEAN[1], red - VGG_MEAN[2] ], axis=3)65 66 network = tl.layers.InputLayer(bgr, name='encoder/input_layer_bgr')67 68 # Convolutional layers size 169 network = conv_layer(network, [ 3, 64], 'encoder/h1/conv_1')70 beforepool1 = conv_layer(network, [64, 64], 'encoder/h1/conv_2')71 network = pool_layer(beforepool1, 'encoder/h1/pool')72 73 # Convolutional layers size 274 network = conv_layer(network, [64, 128], 'encoder/h2/conv_1')75 beforepool2 = conv_layer(network, [128, 128], 'encoder/h2/conv_2')76 network = pool_layer(beforepool2, 'encoder/h2/pool')77 78 # Convolutional layers size 379 network = conv_layer(network, [128, 256], 'encoder/h3/conv_1')80 network = conv_layer(network, [256, 256], 'encoder/h3/conv_2')81 beforepool3 = conv_layer(network, [256, 256], 'encoder/h3/conv_3')82 network = pool_layer(beforepool3, 'encoder/h3/pool')83 84 # Convolutional layers size 485 network = conv_layer(network, [256, 512], 'encoder/h4/conv_1')86 network = conv_layer(network, [512, 512], 'encoder/h4/conv_2')87 beforepool4 = conv_layer(network, [512, 512], 'encoder/h4/conv_3')88 network = pool_layer(beforepool4, 'encoder/h4/pool')89 90 # Convolutional layers size 591 network = conv_layer(network, [512, 512], 'encoder/h5/conv_1')92 network = conv_layer(network, [512, 512], 'encoder/h5/conv_2')93 beforepool5 = conv_layer(network, [512, 512], 'encoder/h5/conv_3')94 network = pool_layer(beforepool5, 'encoder/h5/pool')95 96 return network, (input_layer, beforepool1, beforepool2, beforepool3, beforepool4, beforepool5)97 98 99# Decoder network100def decoder(input_layer, skip_layers, batch_size=1, is_training=False):101 sb, sx, sy, sf = input_layer.outputs.get_shape().as_list()102 alpha = 0.0103 104 # Upsampling 1105 network = deconv_layer(input_layer, (batch_size,sx,sy,sf,sf), 'decoder/h1/decon2d', alpha, is_training)106 107 # Upsampling 2108 network = skip_connection_layer(network, skip_layers[5], 'decoder/h2/fuse_skip_connection', is_training)109 network = deconv_layer(network, (batch_size,2*sx,2*sy,sf,sf), 'decoder/h2/decon2d', alpha, is_training)110 111 # Upsampling 3112 network = skip_connection_layer(network, skip_layers[4], 'decoder/h3/fuse_skip_connection', is_training)113 network = deconv_layer(network, (batch_size,4*sx,4*sy,sf,sf/2), 'decoder/h3/decon2d', alpha, is_training)114 115 # Upsampling 4116 network = skip_connection_layer(network, skip_layers[3], 'decoder/h4/fuse_skip_connection', is_training)117 network = deconv_layer(network, (batch_size,8*sx,8*sy,sf/2,sf/4), 'decoder/h4/decon2d', alpha, is_training)118 119 # Upsampling 5120 network = skip_connection_layer(network, skip_layers[2], 'decoder/h5/fuse_skip_connection', is_training)121 network = deconv_layer(network, (batch_size,16*sx,16*sy,sf/4,sf/8), 'decoder/h5/decon2d', alpha, is_training)122 123 # Skip-connection at full size124 network = skip_connection_layer(network, skip_layers[1], 'decoder/h6/fuse_skip_connection', is_training)125 126 # Final convolution127 network = tl.layers.Conv2dLayer(network,128 act = tf.identity,129 shape = [1, 1, int(sf/8), 3],130 strides=[1, 1, 1, 1],131 padding='SAME',132 W_init = tf.contrib.layers.xavier_initializer(uniform=False),133 b_init = tf.constant_initializer(value=0.0),134 name ='decoder/h7/conv2d')135 136 # Final skip-connection137 network = tl.layers.BatchNormLayer(network, is_train=is_training, name='decoder/h7/batch_norm')138 network.outputs = tf.maximum(alpha*network.outputs, network.outputs, name='decoder/h7/leaky_relu')139 network = skip_connection_layer(network, skip_layers[0], 'decoder/h7/fuse_skip_connection')140 141 return network142 143 144# Load weights for VGG16 encoder convolutional layers145# Weights are from a .npy file generated with the caffe-tensorflow tool146def load_vgg_weights(network, weight_file, session):147 params = []148 149 if weight_file.lower().endswith('.npy'):150 npy = np.load(weight_file, encoding='latin1')151 for key, val in sorted(npy.item().items()):152 if(key[:4] == "conv"):153 print(" Loading %s" % (key))154 print(" weights with size %s " % str(val['weights'].shape))155 print(" and biases with size %s " % str(val['biases'].shape))156 params.append(val['weights'])157 params.append(val['biases'])158 else:159 print('No weights in suitable .npy format found for path ', weight_file)160 161 print('Assigning loaded weights..')162 tl.files.assign_params(session, params, network)163 164 return network165 166 167# === Layers ==================================================================168 169# Convolutional layer170def conv_layer(input_layer, sz, str):171 network = tl.layers.Conv2dLayer(input_layer,172 act = tf.nn.relu,173 shape = [3, 3, sz[0], sz[1]],174 strides = [1, 1, 1, 1],175 padding = 'SAME',176 name = str)177 178 return network179 180 181# Max-pooling layer182def pool_layer(input_layer, str):183 network = tl.layers.PoolLayer(input_layer,184 ksize=[1, 2, 2, 1],185 strides=[1, 2, 2, 1],186 padding='SAME',187 pool = tf.nn.max_pool,188 name = str)189 190 return network191 192 193# Concatenating fusion of skip-connections194def skip_connection_layer(input_layer, skip_layer, str, is_training=False):195 _, sx, sy, sf = input_layer.outputs.get_shape().as_list()196 _, sx_, sy_, sf_ = skip_layer.outputs.get_shape().as_list()197 198 assert (sx_,sy_,sf_) == (sx,sy,sf)199 200 # skip-connection domain transformation, from LDR encoder to log HDR decoder201 skip_layer.outputs = tf.log(tf.pow(tf.scalar_mul(1.0/255, skip_layer.outputs), 2.0)+1.0/255.0)202 203 # specify weights for fusion of concatenation, so that it performs an element-wise addition204 weights = np.zeros((1, 1, sf+sf_, sf))205 for i in range(sf):206 weights[0, 0, i, i] = 1207 weights[:, :, i+sf_, i] = 1208 add_init = tf.constant_initializer(value=weights, dtype=tf.float32)209 210 # concatenate layers211 network = tl.layers.ConcatLayer([input_layer,skip_layer], concat_dim=3, name ='%s/skip_connection'%str)212 213 # fuse concatenated layers using the specified weights for initialization214 network = tl.layers.Conv2dLayer(network,215 act = tf.identity,216 shape = [1, 1, sf+sf_, sf],217 strides = [1, 1, 1, 1],218 padding = 'SAME',219 W_init = add_init,220 b_init = tf.constant_initializer(value=0.0),221 name = str)222 223 return network224 225 226# Deconvolution layer227def deconv_layer(input_layer, sz, str, alpha, is_training=False):228 scale = 2229 230 filter_size = (2 * scale - scale % 2)231 num_in_channels = int(sz[3])232 num_out_channels = int(sz[4])233 234 # create bilinear weights in numpy array235 bilinear_kernel = np.zeros([filter_size, filter_size], dtype=np.float32)236 scale_factor = (filter_size + 1) // 2237 if filter_size % 2 == 1:238 center = scale_factor - 1239 else:240 center = scale_factor - 0.5241 for x in range(filter_size):242 for y in range(filter_size):243 bilinear_kernel[x,y] = (1 - abs(x - center) / scale_factor) * \244 (1 - abs(y - center) / scale_factor)245 weights = np.zeros((filter_size, filter_size, num_out_channels, num_in_channels))246 for i in range(num_out_channels):247 weights[:, :, i, i] = bilinear_kernel248 249 init_matrix = tf.constant_initializer(value=weights, dtype=tf.float32)250 251 network = tl.layers.DeConv2dLayer(input_layer,252 shape = [filter_size, filter_size, num_out_channels, num_in_channels],253 output_shape = [sz[0], sz[1]*scale, sz[2]*scale, num_out_channels],254 strides=[1, scale, scale, 1],255 W_init=init_matrix,256 padding='SAME',257 act=tf.identity,258 name=str)259 260 network = tl.layers.BatchNormLayer(network, is_train=is_training, name='%s/batch_norm_dc'%str)261 network.outputs = tf.maximum(alpha*network.outputs, network.outputs, name='%s/leaky_relu_dc'%str)262 263 return network264 