Add target update param for pretrain + update expectation computation

This commit is contained in:
Antonin RAFFIN 2020-08-31 09:41:14 +02:00
parent 73a6d5117e
commit 5696a4688d
2 changed files with 52 additions and 41 deletions

View file

@ -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,

View file

@ -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()}")