From 0d6e7fa54ec50f96c7a3d2cf16796aec81c1cbde Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Wed, 21 Mar 2018 23:49:41 -0600 Subject: [PATCH] Fix a bug of QR DQN --- agent/QuantileRegressionDQN_agent.py | 6 ++++-- main.py | 2 +- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/agent/QuantileRegressionDQN_agent.py b/agent/QuantileRegressionDQN_agent.py index 56ba90f..fe1a281 100644 --- a/agent/QuantileRegressionDQN_agent.py +++ b/agent/QuantileRegressionDQN_agent.py @@ -78,8 +78,10 @@ class QuantileRegressionDQNAgent: actions = actions.view(-1, 1, 1).expand(-1, -1, quantiles.size(2)) quantiles = quantiles.gather(1, Variable(actions)).squeeze(1) - diff = Variable(quantiles_next) - quantiles - loss = self.huber(diff) * Variable(self.cumulative_density.view(1, -1) - (diff.data < 0).float()).abs() + loss = 0.0 + for i in range(self.config.num_quantiles): + diff = Variable(quantiles_next[:, i].contiguous().view(-1, 1)) - quantiles + loss += self.huber(diff) * Variable(self.cumulative_density.view(1, -1) - (diff.data < 0).float()).abs() self.optimizer.zero_grad() loss.sum(-1).mean().backward() diff --git a/main.py b/main.py index bc5d15d..0caff62 100644 --- a/main.py +++ b/main.py @@ -437,7 +437,7 @@ def quantile_regression_dqn_pixel_atari(name): 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.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.01) 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