mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Log wall time
This commit is contained in:
@@ -29,6 +29,8 @@ def train(id, config, learning_network, extra):
|
||||
def evaluate(config, task, learning_network, extra):
|
||||
test_rewards = []
|
||||
test_points = []
|
||||
test_wall_times = []
|
||||
initial_time = time.time()
|
||||
worker = config.worker(config, learning_network, extra)
|
||||
# config.logger = Logger('./evaluation_log', gym.logger)
|
||||
while True:
|
||||
@@ -45,10 +47,11 @@ def evaluate(config, task, learning_network, extra):
|
||||
(steps, np.mean(rewards), np.std(rewards) / np.sqrt(config.test_repetitions)))
|
||||
test_rewards.append(np.mean(rewards))
|
||||
test_points.append(steps)
|
||||
test_wall_times.append(time.time() - initial_time)
|
||||
with open('data/%s-%s-statistics-%s.bin' % (
|
||||
config.tag, config.worker.__name__, task.name), 'wb') as f:
|
||||
pickle.dump([test_points, test_rewards], f)
|
||||
if np.mean(rewards) > task.success_threshold:
|
||||
pickle.dump([test_rewards, test_points, test_wall_times], f)
|
||||
if np.mean(rewards) > task.success_threshold or (config.max_steps and steps >= config.max_steps):
|
||||
config.stop_signal.value = True
|
||||
break
|
||||
|
||||
@@ -80,8 +83,8 @@ class AsyncAgent:
|
||||
extra = None
|
||||
args = [(i, config, learning_network, extra) for i in range(config.num_workers)]
|
||||
args.append((config, task, learning_network, extra))
|
||||
procs = [mp.Process(target=train, args=args[i]) for i in range(config.num_workers)]
|
||||
procs.append(mp.Process(target=evaluate, args=args[-1]))
|
||||
procs = [mp.Process(target=evaluate, args=args[-1])]
|
||||
procs.extend([mp.Process(target=train, args=args[i]) for i in range(config.num_workers)])
|
||||
for p in procs: p.start()
|
||||
while True:
|
||||
time.sleep(1)
|
||||
|
||||
@@ -85,6 +85,29 @@ def a3c_pendulum():
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
|
||||
def a3c_lunar_lander():
|
||||
config = Config()
|
||||
config.task_fn = lambda: ContinuousLunarLander()
|
||||
task = config.task_fn()
|
||||
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001)
|
||||
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
config.network_fn = lambda: DisjointActorCriticNet(
|
||||
lambda: GaussianActorNet(task.state_dim, task.action_dim),
|
||||
lambda: GaussianCriticNet(task.state_dim))
|
||||
config.policy_fn = lambda: GaussianPolicy()
|
||||
config.worker = ContinuousAdvantageActorCritic
|
||||
config.discount = 0.99
|
||||
config.max_episode_length = 1000
|
||||
config.num_workers = 8
|
||||
config.update_interval = 5
|
||||
config.test_interval = 1
|
||||
config.test_repetitions = 5
|
||||
config.entropy_weight = 0
|
||||
config.gradient_clip = 40
|
||||
config.logger = Logger('./log', gym.logger)
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
|
||||
def a3c_walker():
|
||||
config = Config()
|
||||
config.task_fn = lambda: BipedalWalker()
|
||||
@@ -331,8 +354,8 @@ def ppo_pendulum():
|
||||
config = Config()
|
||||
config.task_fn = lambda: Pendulum()
|
||||
task = config.task_fn()
|
||||
config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim)
|
||||
config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim)
|
||||
config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim, gpu=False)
|
||||
config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim, gpu=False)
|
||||
config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
|
||||
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
@@ -355,12 +378,40 @@ def ppo_pendulum():
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
|
||||
def ppo_lunar_lander():
|
||||
config = Config()
|
||||
config.task_fn = lambda: ContinuousLunarLander()
|
||||
task = config.task_fn()
|
||||
config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim, gpu=False)
|
||||
config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim, gpu=False)
|
||||
config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
|
||||
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
|
||||
config.policy_fn = lambda: GaussianPolicy()
|
||||
config.replay_fn = lambda: GeneralReplay(memory_size=2048, batch_size=2048)
|
||||
config.worker = ProximalPolicyOptimization
|
||||
config.discount = 0.99
|
||||
config.gae_tau = 0.97
|
||||
config.num_workers = 8
|
||||
config.test_interval = 1
|
||||
config.test_repetitions = 1
|
||||
config.max_episode_length = 1000
|
||||
config.entropy_weight = 0
|
||||
config.gradient_clip = 40
|
||||
config.rollout_length = 10000
|
||||
config.optimize_epochs = 1
|
||||
config.ppo_ratio_clip = 0.2
|
||||
config.logger = Logger('./log', gym.logger)
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
|
||||
def ppo_walker():
|
||||
config = Config()
|
||||
config.task_fn = lambda: BipedalWalker()
|
||||
task = config.task_fn()
|
||||
config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim)
|
||||
config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim)
|
||||
config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim, gpu=False)
|
||||
config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim, gpu=False)
|
||||
config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
|
||||
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
@@ -387,16 +438,18 @@ if __name__ == '__main__':
|
||||
# gym.logger.setLevel(logging.DEBUG)
|
||||
gym.logger.setLevel(logging.INFO)
|
||||
|
||||
dqn_cart_pole()
|
||||
# dqn_cart_pole()
|
||||
# async_cart_pole()
|
||||
# a3c_cart_pole()
|
||||
# a3c_pendulum()
|
||||
# a3c_lunar_lander()
|
||||
# a3c_walker()
|
||||
# ddpg_pendulum()
|
||||
# ddpg_lunar_lander()
|
||||
# ddpg_walker()
|
||||
# ppo_pendulum()
|
||||
# ppo_walker()
|
||||
# ppo_lunar_lander()
|
||||
ppo_walker()
|
||||
|
||||
# dqn_fruit()
|
||||
# hrdqn_fruit()
|
||||
|
||||
@@ -19,6 +19,9 @@ class BasicNet:
|
||||
self.LSTM = LSTM
|
||||
if self.gpu:
|
||||
self.cuda()
|
||||
self.FloatTensor = torch.cuda.FloatTensor
|
||||
else:
|
||||
self.FloatTensor = torch.FloatTensor
|
||||
|
||||
def to_torch_variable(self, x, dtype='float32'):
|
||||
if isinstance(x, Variable):
|
||||
|
||||
@@ -47,3 +47,4 @@ class Config:
|
||||
self.num_heads = 10
|
||||
self.min_epsilon = 0
|
||||
self.save_interval = 0
|
||||
self.max_steps = 0
|
||||
|
||||
Reference in New Issue
Block a user