diff --git a/utils/__init__.py b/utils/__init__.py index af5f185..e71c201 100644 --- a/utils/__init__.py +++ b/utils/__init__.py @@ -1,12 +1,7 @@ from .config import * from .normalizer import * from .misc import * - -try: - from .tf_logger import Logger -except: - from .vanilla_logger import Logger - +from .tf_logger import Logger import logging logging.basicConfig(format='%(asctime)s - %(name)s - %(levelname)s: %(message)s') logger = logging.getLogger('MAIN') diff --git a/utils/tf_logger.py b/utils/tf_logger.py index 5d9b506..600dd0a 100644 --- a/utils/tf_logger.py +++ b/utils/tf_logger.py @@ -1,20 +1,14 @@ -# Adapted from https://github.com/yunjey/pytorch-tutorial/blob/master/tutorials/04-utils/tensorboard/logger.py -# Code referenced from https://gist.github.com/gyglim/1f8dfb1b5c82627ae3efcfbbadb9f514 -import tensorflow as tf -import numpy as np -import scipy.misc -import logging - -try: - from StringIO import StringIO # Python 2.7 -except ImportError: - from io import BytesIO # Python 3.x +####################################################################### +# Copyright (C) 2017 Shangtong Zhang(zhangshangtong.cpp@gmail.com) # +# Permission given to modify the code as long as you keep this # +# declaration at the top # +####################################################################### +from tensorboardX import SummaryWriter class Logger(object): def __init__(self, log_dir, vanilla_logger, skip=False): - """Create a summary writer logging to log_dir.""" - self.writer = tf.summary.FileWriter(log_dir) + self.writer = SummaryWriter(log_dir) self.info = vanilla_logger.info self.debug = vanilla_logger.debug self.warning = vanilla_logger.warning @@ -23,61 +17,9 @@ class Logger(object): def scalar_summary(self, tag, value, step): if self.skip: return - """Log a scalar variable.""" - summary = tf.Summary(value=[tf.Summary.Value(tag=tag, simple_value=value)]) - self.writer.add_summary(summary, step) + self.writer.add_scalar(tag, value, step) - def image_summary(self, tag, images, step): + def histo_summary(self, tag, values, step): if self.skip: return - """Log a list of images.""" - - img_summaries = [] - for i, img in enumerate(images): - # Write the image to a string - try: - s = StringIO() - except: - s = BytesIO() - scipy.misc.toimage(img).save(s, format="png") - - # Create an Image object - img_sum = tf.Summary.Image(encoded_image_string=s.getvalue(), - height=img.shape[0], - width=img.shape[1]) - # Create a Summary value - img_summaries.append(tf.Summary.Value(tag='%s/%d' % (tag, i), image=img_sum)) - - # Create and write Summary - summary = tf.Summary(value=img_summaries) - self.writer.add_summary(summary, step) - - def histo_summary(self, tag, values, step, bins=1000): - if self.skip: - return - """Log a histogram of the tensor of values.""" - - # Create a histogram using numpy - counts, bin_edges = np.histogram(values, bins=bins) - - # Fill the fields of the histogram proto - hist = tf.HistogramProto() - hist.min = float(np.min(values)) - hist.max = float(np.max(values)) - hist.num = int(np.prod(values.shape)) - hist.sum = float(np.sum(values)) - hist.sum_squares = float(np.sum(values ** 2)) - - # Drop the start of the first bin - bin_edges = bin_edges[1:] - - # Add bin edges and counts - for edge in bin_edges: - hist.bucket_limit.append(edge) - for c in counts: - hist.bucket.append(c) - - # Create and write Summary - summary = tf.Summary(value=[tf.Summary.Value(tag=tag, histo=hist)]) - self.writer.add_summary(summary, step) - self.writer.flush() \ No newline at end of file + self.writer.add_histogram(tag, values, step) \ No newline at end of file diff --git a/utils/vanilla_logger.py b/utils/vanilla_logger.py deleted file mode 100644 index a21e785..0000000 --- a/utils/vanilla_logger.py +++ /dev/null @@ -1,22 +0,0 @@ -import numpy as np -import logging - -class Logger(object): - def __init__(self, log_dir, vanilla_logger, skip=False): - """Create a summary writer logging to log_dir.""" - self.info = vanilla_logger.info - self.debug = vanilla_logger.debug - self.warning = vanilla_logger.warning - self.skip = skip - - def scalar_summary(self, tag, value, step): - if self.skip: - return - - def image_summary(self, tag, images, step): - if self.skip: - return - - def histo_summary(self, tag, values, step, bins=1000): - if self.skip: - return