From 5696a4688d219810fd23ed458c374af24e028d81 Mon Sep 17 00:00:00 2001 From: Antonin RAFFIN Date: Mon, 31 Aug 2020 09:41:14 +0200 Subject: [PATCH] Add target update param for pretrain + update expectation computation --- stable_baselines3/sac/sac.py | 45 ++++++++++++++++++--------------- stable_baselines3/tqc/tqc.py | 48 ++++++++++++++++++++---------------- 2 files changed, 52 insertions(+), 41 deletions(-) diff --git a/stable_baselines3/sac/sac.py b/stable_baselines3/sac/sac.py index c4e16c0..f605c3c 100644 --- a/stable_baselines3/sac/sac.py +++ b/stable_baselines3/sac/sac.py @@ -382,9 +382,10 @@ class SAC(OffPolicyAlgorithm): self, gradient_steps: int, batch_size: int = 64, - n_action_samples: int = 4, - target_update_interval: int = 100, - strategy: str = "binary", + n_action_samples: int = -1, + target_update_interval: int = 1, + tau: float = 0.005, + strategy: str = "exp", reduce: str = "mean", exp_temperature: float = 1.0, off_policy_update_freq: int = -1, @@ -462,23 +463,27 @@ class SAC(OffPolicyAlgorithm): with th.no_grad(): qf_buffer = th.min(*self.critic(replay_data.observations, replay_data.actions)) - qf_agg = None - for _ in range(n_action_samples): - if self.use_sde: - self.actor.reset_noise() - actions_pi, _ = self.actor.action_log_prob(replay_data.observations) + if n_action_samples == -1: + actions_pi = self.actor.forward(replay_data.observations, deterministic=True) + qf_agg = th.min(*self.critic(replay_data.observations, actions_pi)) + else: + qf_agg = None + for _ in range(n_action_samples): + if self.use_sde: + self.actor.reset_noise() + actions_pi, _ = self.actor.action_log_prob(replay_data.observations) - qf_pi = th.min(*self.critic(replay_data.observations, actions_pi.detach())) - if qf_agg is None: - if reduce == "max": - qf_agg = qf_pi + qf_pi = th.min(*self.critic(replay_data.observations, actions_pi.detach())) + if qf_agg is None: + if reduce == "max": + qf_agg = qf_pi + else: + qf_agg = qf_pi / n_action_samples else: - qf_agg = qf_pi / n_action_samples - else: - if reduce == "max": - qf_agg = th.max(qf_pi, qf_agg) - else: - qf_agg += qf_pi / n_action_samples + if reduce == "max": + qf_agg = th.max(qf_pi, qf_agg) + else: + qf_agg += qf_pi / n_action_samples advantage = qf_buffer - qf_agg if strategy == "binary": @@ -503,9 +508,9 @@ class SAC(OffPolicyAlgorithm): actor_loss.backward() self.actor.optimizer.step() - # Hard copy + # Update target networks if gradient_step % target_update_interval == 0: - self.critic_target.load_state_dict(self.critic.state_dict()) + polyak_update(self.critic.parameters(), self.critic_target.parameters(), tau) def learn( self, diff --git a/stable_baselines3/tqc/tqc.py b/stable_baselines3/tqc/tqc.py index ee4bbc3..34f370b 100644 --- a/stable_baselines3/tqc/tqc.py +++ b/stable_baselines3/tqc/tqc.py @@ -282,9 +282,10 @@ class TQC(OffPolicyAlgorithm): self, gradient_steps: int, batch_size: int = 64, - n_action_samples: int = 4, - target_update_interval: int = 100, - strategy: str = "binary", + n_action_samples: int = -1, + target_update_interval: int = 1, + tau: float = 0.005, + strategy: str = "exp", reduce: str = "mean", exp_temperature: float = 1.0, off_policy_update_freq: int = -1, @@ -387,26 +388,30 @@ class TQC(OffPolicyAlgorithm): # else: # qf_agg = qf_pi.reshape(n_action_samples, batch_size, 1).mean(axis=0) with th.no_grad(): - qf_buffer = self.critic(replay_data.observations, replay_data.actions).mean(2).mean(1, keepdim=True) - qf_agg = None - for _ in range(n_action_samples): - if self.use_sde: - self.actor.reset_noise() - actions_pi, _ = self.actor.action_log_prob(replay_data.observations) + # Use the mean (as done in AWAC, cf rlkit) + if n_action_samples == -1: + actions_pi = self.actor.forward(replay_data.observations, deterministic=True) + qf_agg = self.critic(replay_data.observations, actions_pi).mean(2).mean(1, keepdim=True) + else: + qf_agg = None + for _ in range(n_action_samples): + if self.use_sde: + self.actor.reset_noise() + actions_pi, _ = self.actor.action_log_prob(replay_data.observations) - qf_pi = self.critic(replay_data.observations, actions_pi.detach()).mean(2).mean(1, keepdim=True) - if qf_agg is None: - if reduce == "max": - qf_agg = qf_pi + qf_pi = self.critic(replay_data.observations, actions_pi.detach()).mean(2).mean(1, keepdim=True) + if qf_agg is None: + if reduce == "max": + qf_agg = qf_pi + else: + qf_agg = qf_pi / n_action_samples else: - qf_agg = qf_pi / n_action_samples - else: - if reduce == "max": - qf_agg = th.max(qf_pi, qf_agg) - else: - qf_agg += qf_pi / n_action_samples + if reduce == "max": + qf_agg = th.max(qf_pi, qf_agg) + else: + qf_agg += qf_pi / n_action_samples advantage = qf_buffer - qf_agg if strategy == "binary": @@ -431,9 +436,10 @@ class TQC(OffPolicyAlgorithm): actor_loss.backward() self.actor.optimizer.step() - # Hard copy + # Update target networks if gradient_step % target_update_interval == 0: - self.critic_target.load_state_dict(self.critic.state_dict()) + polyak_update(self.critic.parameters(), self.critic_target.parameters(), tau) + if self.use_sde: print(f"std={(self.actor.get_std()).mean().item()}")