Finalize async methods

This commit is contained in:
Shangtong Zhang
2017-06-10 21:08:40 -06:00
parent 9c73abfbff
commit 8ea108e4bb
6 changed files with 76 additions and 111 deletions
+18 -18
View File
@@ -6,8 +6,8 @@ def dqn_cart_pole():
config = dict()
config['task_fn'] = lambda: CartPole()
config['optimizer_fn'] = lambda params: torch.optim.RMSprop(params, 0.001)
config['network_fn'] = lambda optimizer_fn: FullyConnectedNet([8, 50, 200, 2], optimizer_fn)
# config['network_fn'] = lambda optimizer_fn: DuelingFullyConnectedNet([8, 50, 200, 2], optimizer_fn)
config['network_fn'] = lambda optimizer_fn: FCNet([8, 50, 200, 2], optimizer_fn)
# config['network_fn'] = lambda optimizer_fn: DuelingFCNet([8, 50, 200, 2], optimizer_fn)
config['policy_fn'] = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1)
config['replay_fn'] = lambda: Replay(memory_size=10000, batch_size=10)
config['discount'] = 0.99
@@ -27,7 +27,7 @@ def async_cart_pole():
config = dict()
config['task_fn'] = lambda: CartPole()
config['optimizer_fn'] = lambda params: torch.optim.Adam(params, 0.001)
config['network_fn'] = lambda: FullyConnectedNet([4, 50, 200, 2])
config['network_fn'] = lambda: FCNet([4, 50, 200, 2])
config['policy_fn'] = lambda: GreedyPolicy(epsilon=0.5, final_step=5000, min_epsilon=0.1)
config['bootstrap'] = OneStepQLearning
# config['bootstrap'] = NStepQLearning
@@ -49,7 +49,7 @@ def a3c_cart_pole():
config = dict()
config['task_fn'] = lambda: CartPole()
config['optimizer_fn'] = lambda params: torch.optim.Adam(params, 0.001)
config['network_fn'] = lambda: FCActorCriticNet([4, 200, 2])
config['network_fn'] = lambda: ActorCriticFCNet([4, 200, 2])
config['policy_fn'] = SamplePolicy
config['bootstrap'] = AdvantageActorCritic
config['discount'] = 0.99
@@ -70,8 +70,8 @@ def dqn_pixel_atari(name):
n_actions = 6
config['task_fn'] = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False)
config['optimizer_fn'] = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01)
# config['network_fn'] = lambda optimizer_fn: ConvNet(history_length, n_actions, optimizer_fn)
config['network_fn'] = lambda optimizer_fn: DuelingConvNet(history_length, n_actions, optimizer_fn)
config['network_fn'] = lambda optimizer_fn: NatureConvNet(history_length, n_actions, optimizer_fn)
# config['network_fn'] = lambda optimizer_fn: DuelingNatureConvNet(history_length, n_actions, optimizer_fn)
config['policy_fn'] = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1)
config['replay_fn'] = lambda: Replay(memory_size=1000000, batch_size=32, dtype=np.uint8)
config['discount'] = 0.99
@@ -95,20 +95,19 @@ def async_pixel_atari(name):
config['task_fn'] = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42)
config['optimizer_fn'] = lambda params: torch.optim.Adam(params, lr=0.0001)
config['network_fn'] = lambda: OpenAIConvNet(history_length,
n_actions,
LSTM=False)
config['policy_fn'] = lambda: StochasticGreedyPolicy(epsilons=[1.0, 1.0, 1.0],
final_step=1000000,
min_epsilons=[0.1, 0.01, 0.5],
n_actions)
config['policy_fn'] = lambda: StochasticGreedyPolicy(epsilons=[0.5, 0.5, 0.5],
final_step=2000000,
min_epsilons=[0.1, 0.01, 0.2],
probs=[0.4, 0.3, 0.3])
# config['bootstrap'] = OneStepQLearning
config['bootstrap'] = NStepQLearning
# config['bootstrap'] = OneStepSarsa
# config['bootstrap'] = NStepQLearning
config['bootstrap'] = OneStepSarsa
config['discount'] = 0.99
config['target_network_update_freq'] = 10000
config['step_limit'] = 10000
config['n_workers'] = 16
config['update_interval'] = 32
config['update_interval'] = 20
config['test_interval'] = 50000
config['test_repetitions'] = 1
config['history_length'] = history_length
@@ -122,9 +121,9 @@ def a3c_pixel_atari(name):
n_actions = 6
config['task_fn'] = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42)
config['optimizer_fn'] = lambda params: torch.optim.Adam(params, lr=0.0001)
config['network_fn'] = lambda: OpenAIConvActorCriticNet(history_length,
config['network_fn'] = lambda: OpenAIActorCriticConvNet(history_length,
n_actions,
LSTM=True)
LSTM=False)
config['policy_fn'] = SamplePolicy
config['bootstrap'] = AdvantageActorCritic
config['discount'] = 0.99
@@ -137,6 +136,7 @@ def a3c_pixel_atari(name):
config['history_length'] = history_length
config['logger'] = gym.logger
agent = AsyncAgent(**config)
agent.tag = ''
agent.run()
if __name__ == '__main__':
@@ -148,9 +148,9 @@ if __name__ == '__main__':
# a3c_cart_pole()
# dqn_pixel_atari('PongNoFrameskip-v3')
# async_pixel_atari('PongNoFrameskip-v3')
async_pixel_atari('PongNoFrameskip-v3')
# a3c_pixel_atari('PongNoFrameskip-v3')
# dqn_pixel_atari('BreakoutNoFrameskip-v3')
async_pixel_atari('BreakoutNoFrameskip-v3')
# async_pixel_atari('BreakoutNoFrameskip-v3')
# a3c_pixel_atari('BreakoutNoFrameskip-v3')