From d32627d1b5238ae456584c0829e69552d345ba07 Mon Sep 17 00:00:00 2001 From: Nadav Bhonker Date: Tue, 27 Feb 2018 11:30:54 +0200 Subject: [PATCH] fix dataset config, handle no log dir --- dataset.py | 5 +++-- main.py | 2 +- utils/tf_logger.py | 7 +++++-- 3 files changed, 9 insertions(+), 5 deletions(-) diff --git a/dataset.py b/dataset.py index ef47d41..e751a76 100644 --- a/dataset.py +++ b/dataset.py @@ -17,7 +17,8 @@ PREFIX = '/local/data' def dqn_pixel_atari(name): config = Config() config.history_length = 4 - config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False) + config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False, + history_length=config.history_length) action_dim = config.task_fn().action_dim config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) config.network_fn = lambda: NatureConvNet(config.history_length, action_dim) @@ -26,7 +27,7 @@ def dqn_pixel_atari(name): config.discount = 0.99 config.target_network_update_freq = 10000 config.max_episode_length = 0 - config.exploration_steps= 50000 + config.exploration_steps = 50000 config.logger = Logger('./log', logger) config.test_interval = 10 config.test_repetitions = 1 diff --git a/main.py b/main.py index 293a4dd..9cfce40 100644 --- a/main.py +++ b/main.py @@ -99,7 +99,7 @@ def dqn_pixel_atari(name): action_dim = config.task_fn().action_dim config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) config.network_fn = lambda: NatureConvNet(config.history_length, action_dim, gpu=0) - # config.network_fn = lambda: DuelingNatureConvNet(config.history_length, n_actions) + # config.network_fn = lambda: DuelingNatureConvNet(config.history_length, action_dim) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=32, dtype=np.uint8) config.reward_shift_fn = lambda r: np.sign(r) diff --git a/utils/tf_logger.py b/utils/tf_logger.py index 6c2c79b..a608a80 100644 --- a/utils/tf_logger.py +++ b/utils/tf_logger.py @@ -10,8 +10,11 @@ import numpy as np class Logger(object): def __init__(self, log_dir, vanilla_logger, skip=False): - for f in os.listdir(log_dir): - os.remove('%s/%s' % (log_dir, f)) + try: + for f in os.listdir(log_dir): + os.remove('%s/%s' % (log_dir, f)) + except FileNotFoundError: + os.mkdir(log_dir) if not skip: self.writer = SummaryWriter(log_dir) self.info = vanilla_logger.info