QR DQN Atari

This commit is contained in:
Shangtong Zhang
2018-03-21 10:55:59 -06:00
parent 4fbf811188
commit 6863fc9c14
2 changed files with 45 additions and 1 deletions
+23 -1
View File
@@ -429,6 +429,27 @@ def quantile_regression_dqn_cart_pole():
config.num_quantiles = 20
run_episodes(QuantileRegressionDQNAgent(config))
def quantile_regression_dqn_pixel_atari(name):
config = Config()
config.history_length = 4
config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False,
history_length=config.history_length)
action_dim = config.task_fn().action_dim
config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.00005, eps=0.01 / 32)
config.network_fn = lambda: QuantileConvNet(config.history_length, action_dim, config.num_quantiles, gpu=0)
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.reward_shift_fn = lambda r: np.sign(r)
config.discount = 0.99
config.target_network_update_freq = 10000
config.exploration_steps= 50000
config.logger = Logger('./log', logger)
config.test_interval = 10
config.test_repetitions = 1
config.double_q = False
config.num_quantiles = 200
run_episodes(QuantileRegressionDQNAgent(config))
if __name__ == '__main__':
mkdir('data')
mkdir('data/video')
@@ -439,7 +460,7 @@ if __name__ == '__main__':
# dqn_cart_pole()
# categorical_dqn_cart_pole()
quantile_regression_dqn_cart_pole()
# quantile_regression_dqn_cart_pole()
# async_cart_pole()
# a3c_cart_pole()
# a2c_cart_pole()
@@ -451,6 +472,7 @@ if __name__ == '__main__':
# dqn_pixel_atari('PongNoFrameskip-v4')
# categorical_dqn_pixel_atari('PongNoFrameskip-v4')
quantile_regression_dqn_pixel_atari('PongNoFrameskip-v4')
# n_step_dqn_pixel_atari('PongNoFrameskip-v4')
# async_pixel_atari('PongNoFrameskip-v4')
# a3c_pixel_atari('PongNoFrameskip-v4')
+22
View File
@@ -161,4 +161,26 @@ class CategoricalConvNet(nn.Module, CategoricalNet):
y = F.relu(self.conv3(y))
y = y.view(y.size(0), -1)
y = F.relu(self.fc4(y))
return y
class QuantileConvNet(nn.Module, QuantileNet):
def __init__(self, in_channels, n_actions, n_quantiles, gpu=0):
super(QuantileConvNet, self).__init__()
self.conv1 = nn.Conv2d(in_channels, 32, kernel_size=8, stride=4)
self.conv2 = nn.Conv2d(32, 64, kernel_size=4, stride=2)
self.conv3 = nn.Conv2d(64, 64, kernel_size=3, stride=1)
self.fc4 = nn.Linear(7 * 7 * 64, 512)
self.fc5 = nn.Linear(512, n_actions * n_quantiles)
self.n_actions = n_actions
self.n_quantiles = n_quantiles
BasicNet.__init__(self, gpu)
def forward(self, x):
x = self.variable(x)
y = F.relu(self.conv1(x))
y = F.relu(self.conv2(y))
y = F.relu(self.conv3(y))
y = y.view(y.size(0), -1)
y = F.relu(self.fc4(y))
y = self.fc5(y)
return y