Fix a bug of QR DQN

This commit is contained in:
Shangtong Zhang
2018-03-21 23:49:41 -06:00
parent 6863fc9c14
commit 0d6e7fa54e
2 changed files with 5 additions and 3 deletions
+4 -2
View File
@@ -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()