mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-30 20:18:15 +00:00
Add target update param for pretrain + update expectation computation
This commit is contained in:
parent
73a6d5117e
commit
5696a4688d
2 changed files with 52 additions and 41 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()}")
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue