From 4112e27a1e723d04a93ff1e66d66def442ec4ab8 Mon Sep 17 00:00:00 2001 From: Antonin Raffin Date: Thu, 24 Sep 2020 14:36:33 +0200 Subject: [PATCH] Reformat --- stable_baselines3/sac/sac.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/stable_baselines3/sac/sac.py b/stable_baselines3/sac/sac.py index b12aaf0..9c1e40f 100644 --- a/stable_baselines3/sac/sac.py +++ b/stable_baselines3/sac/sac.py @@ -202,7 +202,6 @@ class SAC(OffPolicyAlgorithm): self.alpha_coef_tensor = th.tensor(float(self.alpha_coef)).to(self.device) self.log_alpha = th.log(self.alpha_coef_tensor) - def _create_aliases(self) -> None: self.actor = self.policy.actor self.critic = self.policy.critic @@ -230,7 +229,9 @@ class SAC(OffPolicyAlgorithm): policy_actions.append(actions_pi) n_log_probs.append(log_prob) # (batch, n, action_dim) - policy_actions = th.cat(policy_actions, dim=1).view(len(replay_data.observations), self.n_action_samples, action_dim) + policy_actions = th.cat(policy_actions, dim=1).view( + len(replay_data.observations), self.n_action_samples, action_dim + ) # assert policy_actions.shape == (len(replay_data.observations), self.n_action_samples, action_dim) # (batch, n, 1) n_log_probs = th.cat(n_log_probs).view(len(replay_data.observations), self.n_action_samples, 1)