From 9bf07e302007742d5164db99f6257781c546bc26 Mon Sep 17 00:00:00 2001 From: Javier Barbero Date: Thu, 30 Jan 2020 12:38:39 +0100 Subject: [PATCH] Added GAN parameter for logging generated images to Tensorboard --- niftynet/application/gan_application.py | 21 +++++++-- niftynet/engine/application_variables.py | 9 ++-- niftynet/io/misc_io.py | 49 ++++++++++++++++++++ niftynet/utilities/user_parameters_custom.py | 6 +++ 4 files changed, 79 insertions(+), 6 deletions(-) diff --git a/niftynet/application/gan_application.py b/niftynet/application/gan_application.py index 5ddfc4a0..31dd29e9 100755 --- a/niftynet/application/gan_application.py +++ b/niftynet/application/gan_application.py @@ -179,13 +179,11 @@ def switch_sampler(for_training): stddev=1.0, dtype=tf.float32) conditioning = data_dict['conditioning'] - net_output = self.net( + fake_image, real_logits, fake_logits = self.net( noise, images, conditioning, self.is_training) loss_func = LossFunction( loss_type=self.action_param.loss_type) - real_logits = net_output[1] - fake_logits = net_output[2] lossG, lossD = loss_func(real_logits, fake_logits) if self.net_param.decay > 0: reg_losses = tf.get_collection( @@ -220,6 +218,23 @@ def switch_sampler(for_training): outputs_collector.add_to_collection( var=lossG, name='lossG', average_over_devices=False, collection=TF_SUMMARIES) + # images to display in tensorboard + if self.gan_param.tensorboard_n_fake_images > 0: + outputs_collector.add_to_collection( + var=fake_image[:self.gan_param.tensorboard_n_fake_images], + name='fake_image_sagittal', + collection=TF_SUMMARIES, summary_type='image3_sagittal_n') + + outputs_collector.add_to_collection( + var=fake_image[:self.gan_param.tensorboard_n_fake_images], + name='fake_image_coronal', + collection=TF_SUMMARIES, summary_type='image3_coronal_n') + + outputs_collector.add_to_collection( + var=fake_image[:self.gan_param.tensorboard_n_fake_images], + name='fake_image_axial', + collection=TF_SUMMARIES, summary_type='image3_axial_n') + with tf.name_scope('Optimiser'): optimiser_class = OptimiserFactory.create( diff --git a/niftynet/engine/application_variables.py b/niftynet/engine/application_variables.py index 588b5ef0..5f427ac7 100755 --- a/niftynet/engine/application_variables.py +++ b/niftynet/engine/application_variables.py @@ -9,7 +9,8 @@ from tensorflow.contrib.framework import list_variables from niftynet.io.misc_io import \ - image3_axial, image3_coronal, image3_sagittal, resolve_checkpoint + image3_axial, image3_coronal, image3_sagittal, \ + image3_axial_n, image3_coronal_n, image3_sagittal_n, resolve_checkpoint from niftynet.utilities import util_common as util from niftynet.utilities.restore_initializer import restore_initializer @@ -22,8 +23,10 @@ 'image': tf.summary.image, 'image3_sagittal': image3_sagittal, 'image3_coronal': image3_coronal, - 'image3_axial': image3_axial} - + 'image3_axial': image3_axial, + 'image3_sagittal_n': image3_sagittal_n, + 'image3_coronal_n': image3_coronal_n, + 'image3_axial_n': image3_axial_n} class GradientsCollector(object): """ diff --git a/niftynet/io/misc_io.py b/niftynet/io/misc_io.py index f3ced819..bd594ce0 100755 --- a/niftynet/io/misc_io.py +++ b/niftynet/io/misc_io.py @@ -840,6 +840,55 @@ def image3_axial(name, return image3(name, tensor, max_outputs, collections, [3], [1, 2]) +def image3_sagittal_n(name, + tensor, + collections=(tf.GraphKeys.SUMMARIES, )): + """ + Create 2D image summary in the sagittal view. An image will be generated for + every element in axis 0 of the tensor. + + :param name: + :param tensor: + :param max_outputs: + :param collections: + :return: + """ + return image3(name, tensor, tensor.shape.as_list()[0], collections, [1], [2, 3]) + + +def image3_coronal_n(name, + tensor, + collections=(tf.GraphKeys.SUMMARIES, )): + """ + Create 2D image summary in the coronal view. An image will be generated for + every element in axis 0 of the tensor. + + :param name: + :param tensor: + :param max_outputs: + :param collections: + :return: + """ + return image3(name, tensor, tensor.shape.as_list()[0], collections, [2], [1, 3]) + + +def image3_axial_n(name, + tensor, + max_outputs=3, + collections=(tf.GraphKeys.SUMMARIES, )): + """ + Create 2D image summary in the axial view. An image will be generated for + every element in axis 0 of the tensor. + + :param name: + :param tensor: + :param max_outputs: + :param collections: + :return: + """ + return image3(name, tensor, tensor.shape.as_list()[0], collections, [3], [1, 2]) + + def set_logger(file_name=None): """ Writing logs to a file if file_name, diff --git a/niftynet/utilities/user_parameters_custom.py b/niftynet/utilities/user_parameters_custom.py index df0e22e0..ef4d888e 100755 --- a/niftynet/utilities/user_parameters_custom.py +++ b/niftynet/utilities/user_parameters_custom.py @@ -211,6 +211,12 @@ def __add_gan_args(parser): type=int, default=10) + parser.add_argument( + "--tensorboard_n_fake_images", + help="the number of fake images to log to Tensorboard in every update", + type=int, + default=0) + from niftynet.application.gan_application import SUPPORTED_INPUT parser = add_input_name_args(parser, SUPPORTED_INPUT) return parser