mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Update run_iterations
This commit is contained in:
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user