From 2eefb04b9dd2f7a7648c381aa10aa2f68c46e61e Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Wed, 25 Apr 2018 23:37:23 -0600 Subject: [PATCH] Refactor DDPG --- agent/DDPG_agent.py | 9 ++------- network/network_utils.py | 2 ++ 2 files changed, 4 insertions(+), 7 deletions(-) diff --git a/agent/DDPG_agent.py b/agent/DDPG_agent.py index 0b62391..b5590e0 100644 --- a/agent/DDPG_agent.py +++ b/agent/DDPG_agent.py @@ -91,19 +91,14 @@ class DDPGAgent(BaseAgent): q = critic.predict(states, actions) critic_loss = self.criterion(q, q_next) - critic.zero_grad() self.critic_opt.zero_grad() critic_loss.backward() self.critic_opt.step() - actions = actor.predict(states, False) - var_actions = actions.detach().requires_grad_() - q = critic.predict(states, var_actions) - q.backward(critic.tensor(np.ones(q.size()))) + policy_loss = -critic.predict(states, actor.predict(states)).mean() - actor.zero_grad() self.actor_opt.zero_grad() - actions.backward(-var_actions.grad) + policy_loss.backward() torch.nn.utils.clip_grad_value_(actor.parameters(), config.gradient_clip) self.actor_opt.step() diff --git a/network/network_utils.py b/network/network_utils.py index d390c73..3a54b7b 100644 --- a/network/network_utils.py +++ b/network/network_utils.py @@ -18,6 +18,8 @@ class BaseNet: self.to(self.device) def tensor(self, x): + if isinstance(x, torch.Tensor): + return x x = torch.tensor(x, device=self.device, dtype=torch.float32) return x