From 14d9059b5cf0a8e6af6449a8a194039e99e50e49 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Thu, 22 Mar 2018 00:03:38 -0600 Subject: [PATCH] Optimize computation of QR DQN loss --- agent/QuantileRegressionDQN_agent.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/agent/QuantileRegressionDQN_agent.py b/agent/QuantileRegressionDQN_agent.py index fe1a281..b9cd72c 100644 --- a/agent/QuantileRegressionDQN_agent.py +++ b/agent/QuantileRegressionDQN_agent.py @@ -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())