From a9d4a0fa72007c720b9bb2b449217dd323b351e4 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Sun, 4 Jun 2017 19:21:04 -0600 Subject: [PATCH] New network for value based async methods --- README.md | 2 +- main.py | 14 +++++++++----- network.py | 18 ++++++++++++++++++ policy.py | 14 ++++++++++++++ 4 files changed, 42 insertions(+), 6 deletions(-) diff --git a/README.md b/README.md index be09257..5781eda 100644 --- a/README.md +++ b/README.md @@ -13,7 +13,7 @@ Implemented algorithms: * Async N-Step Q-Learning # Curves -> Curves for CartPole is trivial so I didn't place it here. +> Curves for CartPole are trivial so I didn't place it here. ## DQN, Double DQN, Dueling DQN ![Loading...](https://raw.githubusercontent.com/ShangtongZhang/DeepRL/master/images/DQN-breakout.png) ![Loading...](https://raw.githubusercontent.com/ShangtongZhang/DeepRL/master/images/DQN-Pong.png) diff --git a/main.py b/main.py index 83a5ca2..ef26e62 100644 --- a/main.py +++ b/main.py @@ -93,13 +93,17 @@ def async_pixel_atari(name): n_actions = 6 config['task_fn'] = lambda: PixelAtari(name, no_op=30, frame_skip=4) config['optimizer_fn'] = lambda params: torch.optim.Adam(params, lr=0.0001) - config['network_fn'] = lambda: ConvNet(history_length, n_actions, gpu=False) - config['policy_fn'] = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) + # config['optimizer_fn'] = lambda params: torch.optim.RMSprop(params, lr=0.0001) + config['network_fn'] = lambda: NipsConvNet(history_length, n_actions, gpu=False) + config['policy_fn'] = lambda: StochasticGreedyPolicy(epsilons=[1.0, 1.0, 1.0], + final_step=int(4000000/16), + min_epsilons=[0.1, 0.01, 0.5], + probs=[0.4, 0.3, 0.3]) # config['bootstrap_fn'] = OneStepQLearning # config['bootstrap_fn'] = NStepQLearning config['bootstrap_fn'] = OneStepSarsa config['discount'] = 0.99 - config['target_network_update_freq'] = 10000 + config['target_network_update_freq'] = 40000 config['step_limit'] = 10000 config['n_workers'] = 16 config['batch_size'] = 20 @@ -136,11 +140,11 @@ if __name__ == '__main__': gym.logger.setLevel(logging.INFO) # async_cart_pole() - dqn_cart_pole() + # dqn_cart_pole() # dqn_pixel_atari('BreakoutNoFrameskip-v3') # async_pixel_atari('BreakoutNoFrameskip-v3') # a3c_pixel_atari('BreakoutNoFrameskip-v3') # a3c_cart_pole() # dqn_pixel_atari('PongNoFrameskip-v3') - # async_pixel_atari('PongNoFrameskip-v3') + async_pixel_atari('PongNoFrameskip-v3') # a3c_pixel_atari('PongNoFrameskip-v3') diff --git a/network.py b/network.py index e13467d..631d3ce 100644 --- a/network.py +++ b/network.py @@ -134,6 +134,24 @@ class ConvNet(nn.Module, VanillaNet): y = F.relu(self.fc4(y)) return self.fc5(y) +class NipsConvNet(nn.Module, VanillaNet): + def __init__(self, in_channels, n_actions, optimizer_fn=None, gpu=True): + super(NipsConvNet, self).__init__() + self.conv1 = nn.Conv2d(in_channels, 16, kernel_size=8, stride=4) + self.conv2 = nn.Conv2d(16, 32, kernel_size=4, stride=2) + self.fc3 = nn.Linear(9 * 9 * 32, 256) + self.fc4 = nn.Linear(256, n_actions) + self.criterion = nn.MSELoss() + BasicNet.__init__(self, optimizer_fn, gpu) + + def forward(self, x): + x = self.to_torch_variable(x) + y = F.relu(self.conv1(x)) + y = F.relu(self.conv2(y)) + y = y.view(y.size(0), -1) + y = F.relu(self.fc3(y)) + return self.fc4(y) + # Network for pixel Atari game with dueling architecture class DuelingConvNet(nn.Module, DuelingNet): def __init__(self, in_channels, n_actions, optimizer_fn=None, gpu=True): diff --git a/policy.py b/policy.py index 71cd9d0..3e4c13e 100644 --- a/policy.py +++ b/policy.py @@ -23,6 +23,20 @@ class GreedyPolicy: self.epsilon = max(self.epsilon, self.min_epsilon) self.current_steps += 1 +class StochasticGreedyPolicy: + def __init__(self, epsilons, final_step, min_epsilons, probs): + self.policies = [] + self.probs = probs + for epsilon, min_epsilon in zip(epsilons, min_epsilons): + self.policies.append(GreedyPolicy(epsilon, final_step, min_epsilon)) + + def sample(self, action_value): + return np.random.choice(self.policies, p=self.probs).sample(action_value) + + def update_epsilon(self): + for policy in self.policies: + policy.update_epsilon() + class SamplePolicy: def sample(self, action_value): return np.random.choice(np.arange(len(action_value)), p=action_value)