(args)
| 38 | |
| 39 | |
| 40 | def train(args): |
| 41 | |
| 42 | |
| 43 | input_photo = tf.placeholder(tf.float32, [args.batch_size, |
| 44 | args.patch_size, args.patch_size, 3]) |
| 45 | input_superpixel = tf.placeholder(tf.float32, [args.batch_size, |
| 46 | args.patch_size, args.patch_size, 3]) |
| 47 | input_cartoon = tf.placeholder(tf.float32, [args.batch_size, |
| 48 | args.patch_size, args.patch_size, 3]) |
| 49 | |
| 50 | output = network.unet_generator(input_photo) |
| 51 | output = guided_filter(input_photo, output, r=1) |
| 52 | |
| 53 | |
| 54 | blur_fake = guided_filter(output, output, r=5, eps=2e-1) |
| 55 | blur_cartoon = guided_filter(input_cartoon, input_cartoon, r=5, eps=2e-1) |
| 56 | |
| 57 | gray_fake, gray_cartoon = utils.color_shift(output, input_cartoon) |
| 58 | |
| 59 | d_loss_gray, g_loss_gray = loss.lsgan_loss(network.disc_sn, gray_cartoon, gray_fake, |
| 60 | scale=1, patch=True, name='disc_gray') |
| 61 | d_loss_blur, g_loss_blur = loss.lsgan_loss(network.disc_sn, blur_cartoon, blur_fake, |
| 62 | scale=1, patch=True, name='disc_blur') |
| 63 | |
| 64 | |
| 65 | vgg_model = loss.Vgg19('vgg19_no_fc.npy') |
| 66 | vgg_photo = vgg_model.build_conv4_4(input_photo) |
| 67 | vgg_output = vgg_model.build_conv4_4(output) |
| 68 | vgg_superpixel = vgg_model.build_conv4_4(input_superpixel) |
| 69 | h, w, c = vgg_photo.get_shape().as_list()[1:] |
| 70 | |
| 71 | photo_loss = tf.reduce_mean(tf.losses.absolute_difference(vgg_photo, vgg_output))/(h*w*c) |
| 72 | superpixel_loss = tf.reduce_mean(tf.losses.absolute_difference\ |
| 73 | (vgg_superpixel, vgg_output))/(h*w*c) |
| 74 | recon_loss = photo_loss + superpixel_loss |
| 75 | tv_loss = loss.total_variation_loss(output) |
| 76 | |
| 77 | g_loss_total = 1e4*tv_loss + 1e-1*g_loss_blur + g_loss_gray + 2e2*recon_loss |
| 78 | d_loss_total = d_loss_blur + d_loss_gray |
| 79 | |
| 80 | all_vars = tf.trainable_variables() |
| 81 | gene_vars = [var for var in all_vars if 'gene' in var.name] |
| 82 | disc_vars = [var for var in all_vars if 'disc' in var.name] |
| 83 | |
| 84 | |
| 85 | tf.summary.scalar('tv_loss', tv_loss) |
| 86 | tf.summary.scalar('photo_loss', photo_loss) |
| 87 | tf.summary.scalar('superpixel_loss', superpixel_loss) |
| 88 | tf.summary.scalar('recon_loss', recon_loss) |
| 89 | tf.summary.scalar('d_loss_gray', d_loss_gray) |
| 90 | tf.summary.scalar('g_loss_gray', g_loss_gray) |
| 91 | tf.summary.scalar('d_loss_blur', d_loss_blur) |
| 92 | tf.summary.scalar('g_loss_blur', g_loss_blur) |
| 93 | tf.summary.scalar('d_loss_total', d_loss_total) |
| 94 | tf.summary.scalar('g_loss_total', g_loss_total) |
| 95 | |
| 96 | update_ops = tf.get_collection(tf.GraphKeys.UPDATE_OPS) |
| 97 | with tf.control_dependencies(update_ops): |
no test coverage detected