From aa60e711e107dfdbb4a8597e097a3074d789c78f Mon Sep 17 00:00:00 2001 From: Antonin Raffin Date: Mon, 29 Aug 2022 10:51:45 +0200 Subject: [PATCH] Add policy delay --- stable_baselines3/sac/sac.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/stable_baselines3/sac/sac.py b/stable_baselines3/sac/sac.py index 384bbb0..ed82932 100644 --- a/stable_baselines3/sac/sac.py +++ b/stable_baselines3/sac/sac.py @@ -95,6 +95,7 @@ class SAC(OffPolicyAlgorithm): replay_buffer_class: Optional[ReplayBuffer] = None, replay_buffer_kwargs: Optional[Dict[str, Any]] = None, optimize_memory_usage: bool = False, + policy_delay: int = 1, ent_coef: Union[str, float] = "auto", target_update_interval: int = 1, target_entropy: Union[str, float] = "auto", @@ -145,6 +146,7 @@ class SAC(OffPolicyAlgorithm): self.ent_coef = ent_coef self.target_update_interval = target_update_interval self.ent_coef_optimizer = None + self.policy_delay = policy_delay if _init_setup_model: self._setup_model() @@ -202,11 +204,10 @@ class SAC(OffPolicyAlgorithm): ent_coef_losses, ent_coefs = [], [] actor_losses, critic_losses = [], [] - # TODO: properly handle it when train_freq > 1 - policy_update_delay = gradient_steps for gradient_step in range(gradient_steps): - update_actor = ((gradient_step + 1) % policy_update_delay == 0) or gradient_step == gradient_steps - 1 + self._n_updates += 1 + update_actor = self._n_updates % self.policy_delay == 0 # Sample replay buffer replay_data = self.replay_buffer.sample(batch_size, env=self._vec_normalize_env) @@ -288,8 +289,6 @@ class SAC(OffPolicyAlgorithm): # Copy running stats, see GH issue #996 polyak_update(self.batch_norm_stats, self.batch_norm_stats_target, 1.0) - self._n_updates += gradient_steps - self.logger.record("train/n_updates", self._n_updates, exclude="tensorboard") self.logger.record("train/ent_coef", np.mean(ent_coefs)) self.logger.record("train/actor_loss", np.mean(actor_losses))