Update run_iterations

This commit is contained in:
Shangtong Zhang committed 2018-03-13 09:38:48 -06:00
1 parent 174cc50440
commit 8f75029d69
3 files changed
+14 -5

No files matched your search

+10 -1
View File
@@ -58,16 +58,25 @@ def run_episodes(agent):
def run_iterations(agent):
config = agent.config
agent_type = agent.__class__.__name__
agent_name = agent.__class__.__name__
iteration = 0
steps = []
rewards = []
while True:
agent.iteration()
steps.append(agent.total_steps)
rewards.append(np.mean(agent.last_episode_rewards))
if iteration % config.iteration_log_interval == 0:
config.logger.info('total steps %d, mean/max/min reward %f/%f/%f' % (
agent.total_steps, np.mean(agent.last_episode_rewards),
np.max(agent.last_episode_rewards),
np.min(agent.last_episode_rewards)
))
if iteration % (config.iteration_log_interval * 100) == 0:
with open('data/%s-%s-online-stats-%s.bin' % (agent_name, config.tag, agent.task.name), 'wb') as f:
pickle.dump({'rewards': rewards,
'steps': steps}, f)
agent.save('data/%s-%s-model-%s.bin' % (agent_name, config.tag, agent.task.name))
iteration += 1
def sync_grad(target_network, src_network):
+1 -1
View File
@@ -13,7 +13,7 @@ class Logger(object):
try:
for f in os.listdir(log_dir):
os.remove('%s/%s' % (log_dir, f))
except FileNotFoundError:
except IOError:
os.mkdir(log_dir)
if not skip:
self.writer = SummaryWriter(log_dir)