diff --git a/deep_rl/utils/logger.py b/deep_rl/utils/logger.py index f39b1b9..bba1fdf 100644 --- a/deep_rl/utils/logger.py +++ b/deep_rl/utils/logger.py @@ -41,9 +41,7 @@ class Logger(object): self.all_steps = {} def to_numpy(self, v): - if isinstance(v, torch.autograd.Variable): - v = v.data - if isinstance(v, torch.FloatTensor): + if isinstance(v, torch.Tensor): v = v.cpu().detach().numpy() return v diff --git a/deep_rl/utils/tf_logger.py b/deep_rl/utils/tf_logger.py deleted file mode 100644 index 3d3bc93..0000000 --- a/deep_rl/utils/tf_logger.py +++ /dev/null @@ -1,59 +0,0 @@ -####################################################################### -# 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 -import os -import numpy as np -import torch - -class Logger(object): - def __init__(self, log_dir, vanilla_logger, skip=False): - try: - for f in os.listdir(log_dir): - if not f.startswith('events'): - continue - os.remove('%s/%s' % (log_dir, f)) - except IOError: - os.mkdir(log_dir) - if not skip: - self.writer = SummaryWriter(log_dir) - self.info = vanilla_logger.info - self.debug = vanilla_logger.debug - self.warning = vanilla_logger.warning - self.skip = skip - 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().detach().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.get_step(tag) - if np.isscalar(value): - value = np.asarray([value]) - self.writer.add_scalar(tag, value, step) - - def histo_summary(self, tag, values, step=None): - if self.skip: - return - self.to_numpy(values) - if step is None: - step = self.get_step(tag) - self.writer.add_histogram(tag, values, step, bins=1000) \ No newline at end of file