From 85760246e1a2d02c9103e1e74a10d3f6fd99ae25 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Fri, 11 May 2018 16:35:20 -0600 Subject: [PATCH] Update logger --- deep_rl/utils/__init__.py | 6 +--- deep_rl/utils/logger.py | 73 +++++++++++++++++++++++++++++++++++++++ deep_rl/utils/misc.py | 3 ++ examples.py | 40 ++++++++++----------- 4 files changed, 96 insertions(+), 26 deletions(-) create mode 100644 deep_rl/utils/logger.py diff --git a/deep_rl/utils/__init__.py b/deep_rl/utils/__init__.py index 366f00a..5d694a3 100644 --- a/deep_rl/utils/__init__.py +++ b/deep_rl/utils/__init__.py @@ -1,10 +1,6 @@ from .config import * from .normalizer import * from .misc import * -from .tf_logger import Logger +from .logger import get_logger from .plot import Plotter from .schedule import * -import logging -logging.basicConfig(format='%(asctime)s - %(name)s - %(levelname)s: %(message)s') -logger = logging.getLogger('MAIN') -logger.setLevel(logging.INFO) \ No newline at end of file diff --git a/deep_rl/utils/logger.py b/deep_rl/utils/logger.py new file mode 100644 index 0000000..f39b1b9 --- /dev/null +++ b/deep_rl/utils/logger.py @@ -0,0 +1,73 @@ +####################################################################### +# 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 +import logging +logging.basicConfig(format='%(asctime)s - %(name)s - %(levelname)s: %(message)s') +from .misc import * + +def get_logger(name='MAIN', file_name=None, log_dir='./log', skip=False, level=logging.INFO): + logger = logging.getLogger(name) + logger.setLevel(level) + if file_name is not None: + file_name = '%s-%s' % (file_name, get_time_str()) + fh = logging.FileHandler('%s/%s.txt' % (log_dir, file_name)) + fh.setFormatter(logging.Formatter('%(asctime)s - %(name)s - %(levelname)s: %(message)s')) + fh.setLevel(level) + logger.addHandler(fh) + return Logger(log_dir, logger, skip) + +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 diff --git a/deep_rl/utils/misc.py b/deep_rl/utils/misc.py index a4aa07c..8d21b5f 100644 --- a/deep_rl/utils/misc.py +++ b/deep_rl/utils/misc.py @@ -120,3 +120,6 @@ class Batcher: indices = np.arange(self.num_entries) np.random.shuffle(indices) self.data = [d[indices] for d in self.data] + +# def torch_max(tensor, dim): +# return torch.max(tensor, dim=dim, keepdim=True)[0] diff --git a/examples.py b/examples.py index d8ee99c..1e37af0 100644 --- a/examples.py +++ b/examples.py @@ -21,7 +21,7 @@ def dqn_cart_pole(): config.discount = 0.99 config.target_network_update_freq = 200 config.exploration_steps = 1000 - config.logger = Logger('./log', logger) + config.logger = get_logger() config.double_q = True # config.double_q = False run_episodes(DQNAgent(config)) @@ -39,7 +39,7 @@ def a2c_cart_pole(): config.network_fn = lambda state_dim, action_dim: ActorCriticNet(action_dim, FCBody(state_dim)) config.policy_fn = SamplePolicy config.discount = 0.99 - config.logger = Logger('./log', logger) + config.logger = get_logger() config.gae_tau = 1.0 config.entropy_weight = 0.01 config.rollout_length = 5 @@ -58,7 +58,7 @@ def categorical_dqn_cart_pole(): config.discount = 0.99 config.target_network_update_freq = 200 config.exploration_steps = 100 - config.logger = Logger('./log', logger, skip=True) + config.logger = get_logger(skip=True) config.categorical_v_max = 100 config.categorical_v_min = -100 config.categorical_n_atoms = 50 @@ -76,7 +76,7 @@ def quantile_regression_dqn_cart_pole(): config.discount = 0.99 config.target_network_update_freq = 200 config.exploration_steps = 100 - config.logger = Logger('./log', logger, skip=True) + config.logger = get_logger(skip=True) config.num_quantiles = 20 run_episodes(QuantileRegressionDQNAgent(config)) @@ -92,7 +92,7 @@ def n_step_dqn_cart_pole(): config.discount = 0.99 config.target_network_update_freq = 200 config.rollout_length = 5 - config.logger = Logger('./log', logger) + config.logger = get_logger() run_iterations(NStepDQNAgent(config)) def ppo_cart_pole(): @@ -105,7 +105,7 @@ def ppo_cart_pole(): config.network_fn = lambda state_dim, action_dim: \ CategoricalActorCriticWrapper(state_dim, action_dim, network_fn, optimizer_fn) config.discount = 0.99 - config.logger = Logger('./log', logger) + config.logger = get_logger() config.use_gae = True config.gae_tau = 0.95 config.entropy_weight = 0.01 @@ -132,7 +132,7 @@ def option_critic_cart_pole(): config.rollout_length = 5 config.termination_regularizer = 0.01 config.entropy_weight = 0.01 - config.logger = Logger('./log', logger) + config.logger = get_logger() run_iterations(OptionCriticAgent(config)) ## Atari games @@ -152,7 +152,7 @@ def dqn_pixel_atari(name): config.discount = 0.99 config.target_network_update_freq = 10000 config.exploration_steps= 50000 - config.logger = Logger('./log', logger) + config.logger = get_logger() # config.double_q = True config.double_q = False run_episodes(DQNAgent(config)) @@ -175,7 +175,7 @@ def a2c_pixel_atari(name): config.entropy_weight = 0.01 config.rollout_length = 5 config.gradient_clip = 0.5 - config.logger = Logger('./log', logger, skip=True) + config.logger = get_logger(skip=True) run_iterations(A2CAgent(config)) def categorical_dqn_pixel_atari(name): @@ -193,7 +193,7 @@ def categorical_dqn_pixel_atari(name): config.reward_normalizer = SignNormalizer() config.target_network_update_freq = 10000 config.exploration_steps= 50000 - config.logger = Logger('./log', logger) + config.logger = get_logger() config.double_q = False config.categorical_v_max = 10 config.categorical_v_min = -10 @@ -215,7 +215,7 @@ def quantile_regression_dqn_pixel_atari(name): config.discount = 0.99 config.target_network_update_freq = 10000 config.exploration_steps= 50000 - config.logger = Logger('./log', logger) + config.logger = get_logger() config.double_q = False config.num_quantiles = 200 run_episodes(QuantileRegressionDQNAgent(config)) @@ -236,7 +236,7 @@ def n_step_dqn_pixel_atari(name): config.target_network_update_freq = 10000 config.rollout_length = 5 config.gradient_clip = 5 - config.logger = Logger('./log', logger) + config.logger = get_logger() run_iterations(NStepDQNAgent(config)) def ppo_pixel_atari(name): @@ -253,7 +253,7 @@ def ppo_pixel_atari(name): config.state_normalizer = ImageNormalizer() config.reward_normalizer = SignNormalizer() config.discount = 0.99 - config.logger = Logger('./log', logger) + config.logger = get_logger() config.use_gae = True config.gae_tau = 0.95 config.entropy_weight = 0.01 @@ -284,7 +284,7 @@ def option_ciritc_pixel_atari(name): config.max_steps = 1e8 config.entropy_weight = 0.01 config.termination_regularizer = 0.01 - config.logger = Logger('./log', logger) + config.logger = get_logger() run_iterations(OptionCriticAgent(config)) def dqn_ram_atari(name): @@ -301,7 +301,7 @@ def dqn_ram_atari(name): config.target_network_update_freq = 10000 config.max_episode_length = 0 config.exploration_steps= 100 - config.logger = Logger('./log', logger) + config.logger = get_logger() config.double_q = True # config.double_q = False run_episodes(DQNAgent(config)) @@ -333,7 +333,7 @@ def ppo_continuous(): config.num_mini_batches = 32 config.ppo_ratio_clip = 0.2 config.iteration_log_interval = 1 - config.logger = Logger('./log', logger) + config.logger = get_logger() run_iterations(PPOAgent(config)) def ddpg_continuous(): @@ -360,7 +360,7 @@ def ddpg_continuous(): config.random_process_fn = lambda action_dim: GaussianProcess(action_dim, LinearSchedule(0.3, 0, 1e6)) config.min_memory_size = 64 config.target_network_mix = 1e-3 - config.logger = Logger('./log', logger) + config.logger = get_logger() run_episodes(DDPGAgent(config)) def plot(): @@ -397,11 +397,9 @@ if __name__ == '__main__': mkdir('dataset') mkdir('log') set_one_thread() - # logger.setLevel(logging.DEBUG) - logger.setLevel(logging.INFO) # dqn_cart_pole() - # a2c_cart_pole() + a2c_cart_pole() # categorical_dqn_cart_pole() # quantile_regression_dqn_cart_pole() # n_step_dqn_cart_pole() @@ -417,7 +415,7 @@ if __name__ == '__main__': # option_ciritc_pixel_atari('BreakoutNoFrameskip-v4') # dqn_ram_atari('Breakout-ramNoFrameskip-v4') - ddpg_continuous() + # ddpg_continuous() # ppo_continuous() # action_conditional_video_prediction()