From 0b8854609158e02335b5c82b7c01a94f419e00c0 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Wed, 19 Jul 2017 10:29:29 -0600 Subject: [PATCH] Fix bug --- async_agent.py | 1 + dqn_agent.py | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/async_agent.py b/async_agent.py index 299b1d4..90cc932 100644 --- a/async_agent.py +++ b/async_agent.py @@ -84,6 +84,7 @@ class AsyncAgent: while True and not self.stop_signal.value: steps, reward = worker.episode() rewards.append(reward) + if len(rewards) > 100: rewards.pop(0) self.logger.debug('worker %d, episode %d, return %f, avg return %f, episode steps %d, total steps %d' % ( id, episode, rewards[-1], np.mean(rewards[-100:]), steps, self.total_steps.value)) diff --git a/dqn_agent.py b/dqn_agent.py index 288d664..6c5929a 100644 --- a/dqn_agent.py +++ b/dqn_agent.py @@ -124,7 +124,7 @@ class DQNAgent: self.logger.info('episode %d, epsilon %f, reward %f, avg reward %f, total steps %d' % ( ep, self.policy.epsilon, reward, avg_reward, self.total_steps)) - if ep % self.test_interval == 0: + if self.test_repetitions and ep % self.test_interval == 0: self.logger.info('Testing...') self.save('data/%sdqn-model-%s.bin' % (self.tag, self.task.name)) test_rewards = []