Last active
July 18, 2017 09:20
-
-
Save creotiv/8c735c79a67e1d2568bf1eb8cc20fed1 to your computer and use it in GitHub Desktop.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| import sys | |
| sys.path.append("../libs") | |
| import tfutils | |
| import numpy as np | |
| import tensorflow as tf | |
| import tensorflow.contrib.graph_editor as ge | |
| ########################################################################## | |
| VGG_MEAN = [103.939, 116.779, 123.68] | |
| starter_learning_rate = 0.2 | |
| steps = 5000 | |
| print_steps = 200 | |
| model_path = '../models/vgg16_imagenet/vgg16.tfmodel' | |
| config = tf.ConfigProto( | |
| gpu_options=tf.GPUOptions(per_process_gpu_memory_fraction=0.5) | |
| ) | |
| layers = [u'import/images', | |
| u'import/conv1_1/Relu', u'import/conv1_2/Relu', | |
| u'import/conv2_1/Relu', u'import/conv2_2/Relu', | |
| u'import/conv3_1/Relu', u'import/conv3_2/Relu', u'import/conv3_3/Relu', | |
| u'import/conv4_1/Relu', u'import/conv4_2/Relu', u'import/conv4_3/Relu', | |
| u'import/conv5_1/Relu', u'import/conv5_2/Relu', u'import/conv5_3/Relu', | |
| ] | |
| image = './book_cover.jpg' | |
| style = './starry_night.jpg' | |
| original_layers = ['import/conv3_3/Relu:0'] | |
| style_layers = [ | |
| 'import/conv1_2/Relu:0', | |
| 'import/conv2_2/Relu:0', 'import/conv3_3/Relu:0', 'import/conv4_3/Relu:0' | |
| ] | |
| image = np.array([tfutils.get_image(image, [224, 224, 3]) - VGG_MEAN]) | |
| style = np.array([tfutils.get_image(style, image.shape[1:]) - VGG_MEAN]) | |
| ########################################################################## | |
| def gram_matrix(tensor): | |
| global image | |
| shape = tensor.get_shape() | |
| num_channels = int(shape[3]) | |
| matrix = tf.reshape(tensor, shape=[-1, num_channels]) | |
| gram = tf.matmul(tf.transpose(matrix), matrix) / image.size | |
| return gram | |
| def get_loss(model, mixed_model, image, style, mixed_image, | |
| original_layers=[], style_layers=[]): | |
| with model.as_default() as g: | |
| sess = tf.Session(config=config, graph=g) | |
| _original_layers = [g.get_tensor_by_name(i) for i in original_layers] | |
| original_values = sess.run( | |
| _original_layers, feed_dict={'import/images:0': image}) | |
| _style_layers = [g.get_tensor_by_name(i) for i in style_layers] | |
| _style_layers = [gram_matrix(layer) for layer in _style_layers] | |
| style_values = sess.run(_style_layers, feed_dict={ | |
| 'import/images:0': style}) | |
| with mixed_model.as_default() as g: | |
| original_layers = [g.get_tensor_by_name(i) for i in original_layers] | |
| original_layer_losses = [] | |
| for value0, layer0 in zip(original_values, original_layers): | |
| value_const0 = tf.constant(value0) | |
| loss0 = tf.nn.l2_loss(layer0 - value_const0) / value0.size | |
| original_layer_losses.append(loss0) | |
| original_loss = tf.reduce_sum( | |
| original_layer_losses, name="original_loss") | |
| style_layers = [g.get_tensor_by_name(i) for i in style_layers] | |
| style_layer_losses = [] | |
| for value1, layer1 in zip(style_values, style_layers): | |
| value_const1 = tf.constant(value1) | |
| loss1 = tf.nn.l2_loss(gram_matrix(layer1) - | |
| value_const1) / value1.size | |
| style_layer_losses.append(loss1) | |
| style_loss = tf.reduce_sum(style_layer_losses, name="style_loss") | |
| # Total Variation Denoising | |
| '''total_var_x = sess.run(tf.reduce_prod( | |
| mixed_image[:, 1:, :, :].get_shape())) | |
| total_var_y = sess.run(tf.reduce_prod( | |
| mixed_image[:, :, 1:, :].get_shape())) | |
| second_term_numerator = tf.nn.l2_loss( | |
| mixed_image[:, 1:, :, :] - mixed_image[:, :image.shape[1] - 1, :, :]) | |
| second_term = second_term_numerator / total_var_y | |
| third_term = (tf.nn.l2_loss( | |
| mixed_image[:, :, 1:, :] - mixed_image[:, :, :image.shape[2] - 1, :]) / total_var_x) | |
| total_variation_loss = second_term + third_term''' | |
| total_loss = original_loss + style_loss # + total_variation_loss | |
| return total_loss, style_loss, original_loss | |
| def optimizer(learning_rate, total_loss, global_step=None): | |
| optimizer = tf.train.AdamOptimizer(learning_rate).minimize( | |
| total_loss, global_step=global_step) | |
| return optimizer | |
| ########################################################################## | |
| model = tfutils.get_model_from_binary(model_path) | |
| mixed_model, gdef = tfutils.get_def_from_binary(model_path) | |
| with mixed_model.as_default() as g: | |
| mixed_image = tf.Variable(image, dtype=tf.float32) | |
| tf.import_graph_def(gdef, input_map={'images': mixed_image}) | |
| sess = tf.Session(config=config, graph=g) | |
| ########################################################################## | |
| total_loss, style_loss, original_loss = get_loss( | |
| model, mixed_model, image, style, mixed_image, | |
| original_layers=original_layers, | |
| style_layers=style_layers | |
| ) | |
| with mixed_model.as_default() as g: | |
| global_step = tf.Variable(0, trainable=False) | |
| learning_rate = tf.train.exponential_decay(starter_learning_rate, global_step, | |
| 500, 0.96, staircase=True) | |
| a1 = tf.summary.scalar("total_loss", tf.reduce_mean(total_loss)) | |
| a2 = tf.summary.scalar("style_loss", tf.reduce_mean(style_loss)) | |
| a3 = tf.summary.scalar("original_loss", tf.reduce_mean(original_loss)) | |
| a4 = tf.summary.scalar("learning_rate", learning_rate) | |
| a5 = tf.summary.image("mixed_image", mixed_image, max_outputs=1) | |
| merged = tf.summary.merge([a1, a2, a3, a4]) | |
| immerged = tf.summary.merge([a5]) | |
| writer = tf.summary.FileWriter('./logs', g) | |
| train = optimizer(learning_rate, total_loss, global_step) | |
| sess.run(tf.global_variables_initializer()) | |
| for i in range(steps): | |
| _, _original_loss, _style_loss, _total_loss, _learning_rate = sess.run( | |
| [train, original_loss, style_loss, total_loss, learning_rate]) | |
| print(_total_loss, _style_loss, _original_loss, _learning_rate) | |
| merged_ = sess.run(merged) | |
| writer.add_summary(merged_, i + 1) | |
| writer.flush() | |
| # Print update and save temporary output | |
| if (i + 1) % print_steps == 0: | |
| print('Generation {} out of {}, loss: {}'.format( | |
| i + 1, steps, _total_loss)) | |
| image_eval = sess.run(immerged) | |
| writer.add_summary(image_eval, i + 1) | |
| writer.flush() |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment