mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
New network for value based async methods
This commit is contained in:
@@ -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
|
||||

|
||||

|
||||
|
||||
@@ -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')
|
||||
|
||||
+18
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user