From d22b66fc1035b91f6f8504ab4f6938465f04ccdd Mon Sep 17 00:00:00 2001 From: Antonin Raffin Date: Mon, 9 Sep 2019 16:45:55 +0200 Subject: [PATCH] Fixes for CUDA support --- torchy_baselines/cem_rl/cem_rl.py | 4 +--- torchy_baselines/td3/policies.py | 28 ++++++++++++++++------------ 2 files changed, 17 insertions(+), 15 deletions(-) diff --git a/torchy_baselines/cem_rl/cem_rl.py b/torchy_baselines/cem_rl/cem_rl.py index fccd9b3..6261f35 100644 --- a/torchy_baselines/cem_rl/cem_rl.py +++ b/torchy_baselines/cem_rl/cem_rl.py @@ -16,11 +16,10 @@ class CEMRL(TD3): Paper: https://arxiv.org/abs/1810.01222 Code: https://github.com/apourchot/CEM-RL """ - def __init__(self, policy, env, policy_kwargs=None, verbose=0, sigma_init=1e-3, pop_size=10, damp=1e-3, damp_limit=1e-5, elitism=False, n_grad=5, policy_freq=2, batch_size=100, - buffer_size=int(1e6), learning_rate=1e-3, seed=0, device='cpu', + buffer_size=int(1e6), learning_rate=1e-3, seed=0, device='auto', action_noise_std=0.0, start_timesteps=100, _init_setup_model=True): super(CEMRL, self).__init__(policy, env, policy_kwargs, verbose, @@ -149,7 +148,6 @@ class CEMRL(TD3): episode_reward += reward # Store data in replay buffer - # self.replay_buffer.add(state, next_state, action, reward, done) self.replay_buffer.add(obs, new_obs, action, reward, done_bool) obs = new_obs diff --git a/torchy_baselines/td3/policies.py b/torchy_baselines/td3/policies.py index e3b05c5..00d05f6 100644 --- a/torchy_baselines/td3/policies.py +++ b/torchy_baselines/td3/policies.py @@ -25,11 +25,11 @@ class BaseNetwork(nn.Module): :return: (np.ndarray) """ - return th.nn.utils.parameters_to_vector(self.parameters()).detach().numpy() + return th.nn.utils.parameters_to_vector(self.parameters()).detach().cpu().numpy() class Actor(BaseNetwork): - def __init__(self, state_dim, action_dim, net_arch=None): + def __init__(self, state_dim, action_dim, net_arch=None, activation_fn=nn.ReLU): super(Actor, self).__init__() if net_arch is None: @@ -39,9 +39,9 @@ class Actor(BaseNetwork): self.actor_net = nn.Sequential( nn.Linear(state_dim, net_arch[0]), - nn.ReLU(), + activation_fn(), nn.Linear(net_arch[0], net_arch[1]), - nn.ReLU(), + activation_fn(), nn.Linear(net_arch[1], action_dim), nn.Tanh(), ) @@ -51,7 +51,7 @@ class Actor(BaseNetwork): class Critic(BaseNetwork): - def __init__(self, state_dim, action_dim, net_arch=None): + def __init__(self, state_dim, action_dim, net_arch=None, activation_fn=nn.ReLU): super(Critic, self).__init__() if net_arch is None: @@ -59,17 +59,17 @@ class Critic(BaseNetwork): self.q1_net = nn.Sequential( nn.Linear(state_dim + action_dim, net_arch[0]), - nn.ReLU(), + activation_fn(), nn.Linear(net_arch[0], net_arch[1]), - nn.ReLU(), + activation_fn(), nn.Linear(net_arch[1], 1), ) self.q2_net = nn.Sequential( nn.Linear(state_dim + action_dim, net_arch[0]), - nn.ReLU(), + activation_fn(), nn.Linear(net_arch[0], net_arch[1]), - nn.ReLU(), + activation_fn(), nn.Linear(net_arch[1], 1), ) @@ -83,11 +83,15 @@ class Critic(BaseNetwork): class TD3Policy(BasePolicy): def __init__(self, observation_space, action_space, - learning_rate=1e-3, net_arch=None, device='cpu'): + learning_rate=1e-3, net_arch=None, device='cpu', + activation_fn=nn.ReLU): super(TD3Policy, self).__init__(observation_space, action_space, device) self.state_dim = self.observation_space.shape[0] self.action_dim = self.action_space.shape[0] self.net_arch = net_arch + self.activation_fn = activation_fn + self.actor, self.actor_target = None, None + self.critic, self.critic_target = None, None self._build(learning_rate) def _build(self, learning_rate): @@ -102,10 +106,10 @@ class TD3Policy(BasePolicy): self.critic.optimizer = th.optim.Adam(self.critic.parameters(), lr=learning_rate) def make_actor(self): - return Actor(self.state_dim, self.action_dim, self.net_arch).to(self.device) + return Actor(self.state_dim, self.action_dim, self.net_arch, self.activation_fn).to(self.device) def make_critic(self): - return Critic(self.state_dim, self.action_dim, self.net_arch).to(self.device) + return Critic(self.state_dim, self.action_dim, self.net_arch, self.activation_fn).to(self.device) MlpPolicy = TD3Policy