Team Ai
Apppublic

ShkShahid/Auto-encoder_For_Image_Reconstruction

sourceHugging Faceapache-2.0updated 4y agoView on Hugging Face
1likes
network.py264 linesDownload Raw Back to root
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