mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Fix a bug of QR DQN
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user