diff --git a/utils/tf_logger.py b/utils/tf_logger.py index 927c0aa..08f4339 100644 --- a/utils/tf_logger.py +++ b/utils/tf_logger.py @@ -7,6 +7,7 @@ from tensorboardX import SummaryWriter import os import numpy as np +import torch class Logger(object): def __init__(self, log_dir, vanilla_logger, skip=False): @@ -23,14 +24,28 @@ class Logger(object): self.debug = vanilla_logger.debug self.warning = vanilla_logger.warning self.skip = skip - self.step = 0 + self.all_steps = {} + + def to_numpy(self, v): + if isinstance(v, torch.autograd.Variable): + v = v.data + if isinstance(v, torch.FloatTensor): + v = v.cpu().numpy() + return v + + def get_step(self, tag): + if tag not in self.all_steps: + self.all_steps[tag] = 0 + step = self.all_steps[tag] + self.all_steps[tag] += 1 + return step def scalar_summary(self, tag, value, step=None): if self.skip: return + self.to_numpy(value) if step is None: - step = self.step - self.step += 1 + step = self.get_step(tag) if np.isscalar(value): value = np.asarray([value]) self.writer.add_scalar(tag, value, step) @@ -38,7 +53,7 @@ class Logger(object): def histo_summary(self, tag, values, step=None): if self.skip: return + self.to_numpy(values) if step is None: - step = self.step - self.step += 1 + step = self.get_step(tag) self.writer.add_histogram(tag, values, step, bins=1000) \ No newline at end of file