Update run_iterations

This commit is contained in:
Shangtong Zhang
2018-03-13 09:38:48 -06:00
parent 174cc50440
commit 8f75029d69
3 changed files with 14 additions and 5 deletions
+3 -3
View File
@@ -402,7 +402,7 @@ def n_step_dqn_pixel_atari(name):
config.num_workers = 8
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
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, task.action_dim, gpu=1)
config.network_fn = lambda: NatureConvNet(config.history_length, task.action_dim, gpu=0)
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1)
config.reward_shift_fn = lambda r: np.sign(r)
config.discount = 0.99
@@ -423,7 +423,7 @@ if __name__ == '__main__':
# categorical_dqn_cart_pole()
# async_cart_pole()
# a3c_cart_pole()
# a2c_cart_pole()
a2c_cart_pole()
# a3c_continuous()
# p3o_continuous()
# d3pg_continuous()
@@ -432,7 +432,7 @@ if __name__ == '__main__':
# dqn_pixel_atari('PongNoFrameskip-v4')
# categorical_dqn_pixel_atari('PongNoFrameskip-v4')
n_step_dqn_pixel_atari('PongNoFrameskip-v4')
# n_step_dqn_pixel_atari('PongNoFrameskip-v4')
# async_pixel_atari('PongNoFrameskip-v4')
# a3c_pixel_atari('PongNoFrameskip-v4')
# a2c_pixel_atari('PongNoFrameskip-v4')
+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)