Update logger

This commit is contained in:
Shangtong Zhang
2018-04-12 22:02:16 -06:00
parent aef2e71ed1
commit f255c6ed25
+20 -5
View File
@@ -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)