Optimize computation of QR DQN loss

This commit is contained in:
Shangtong Zhang
2018-03-22 00:03:38 -06:00
parent 0d6e7fa54e
commit 14d9059b5c
+4 -5
View File
@@ -78,13 +78,12 @@ class QuantileRegressionDQNAgent:
actions = actions.view(-1, 1, 1).expand(-1, -1, quantiles.size(2))
quantiles = quantiles.gather(1, Variable(actions)).squeeze(1)
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()
quantiles_next = quantiles_next.t().unsqueeze(-1)
diff = Variable(quantiles_next) - 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()
loss.mean(1).sum().backward()
self.optimizer.step()
if not deterministic and self.total_steps % self.config.target_network_update_freq == 0:
self.target_network.load_state_dict(self.learning_network.state_dict())