mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-30 20:18:15 +00:00
Clean up
This commit is contained in:
parent
9e8f6e0020
commit
0174ec269e
3 changed files with 12 additions and 15 deletions
|
|
@ -112,6 +112,7 @@ class A2C(PPO):
|
|||
# Optimization step
|
||||
self.policy.optimizer.zero_grad()
|
||||
loss.backward()
|
||||
|
||||
# Clip grad norm
|
||||
th.nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
|
||||
self.policy.optimizer.step()
|
||||
|
|
@ -123,9 +124,10 @@ class A2C(PPO):
|
|||
logger.logkv("entropy", entropy.mean().item())
|
||||
logger.logkv("policy_loss", policy_loss.item())
|
||||
logger.logkv("value_loss", value_loss.item())
|
||||
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
||||
|
||||
if self.use_sde:
|
||||
logger.logkv("noise net std", th.exp(self.policy.log_std).mean().item())
|
||||
pass
|
||||
# print(th.exp(self.policy.log_std).detach())
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -180,7 +180,7 @@ class StateDependentNoiseDistribution(Distribution):
|
|||
def get_std(self, log_std):
|
||||
if self.use_expln:
|
||||
# From SDE paper, it allows to keep variance
|
||||
# above zero and prevent it from growing too fast
|
||||
# above zero and prevent it from growing too fast
|
||||
if log_std <= 0:
|
||||
return th.exp(log_std)
|
||||
else:
|
||||
|
|
@ -192,8 +192,7 @@ class StateDependentNoiseDistribution(Distribution):
|
|||
self.weights_dist = Normal(th.zeros_like(log_std), self.get_std(log_std))
|
||||
self.noise_weights = self.weights_dist.rsample()
|
||||
|
||||
def proba_distribution_net(self, latent_dim, log_std_init=-3):
|
||||
print("Log std init:", log_std_init)
|
||||
def proba_distribution_net(self, latent_dim, log_std_init=-1):
|
||||
mean_actions = nn.Linear(latent_dim, self.action_dim)
|
||||
# TODO: log_std_init depending on the number of layers?
|
||||
log_std = nn.Parameter(th.ones(self.features_dim, self.action_dim) * log_std_init)
|
||||
|
|
|
|||
|
|
@ -103,7 +103,7 @@ class PPOPolicy(BasePolicy):
|
|||
def __init__(self, observation_space, action_space,
|
||||
learning_rate, net_arch=None, device='cpu',
|
||||
activation_fn=nn.Tanh, adam_epsilon=1e-5,
|
||||
ortho_init=True, use_sde=False):
|
||||
ortho_init=True, use_sde=False, log_std_init=0.0):
|
||||
super(PPOPolicy, self).__init__(observation_space, action_space, device)
|
||||
self.obs_dim = self.observation_space.shape[0]
|
||||
if net_arch is None:
|
||||
|
|
@ -123,6 +123,7 @@ class PPOPolicy(BasePolicy):
|
|||
# In the future, feature_extractor will be replaced with a CNN
|
||||
self.features_extractor = nn.Flatten()
|
||||
self.features_dim = self.obs_dim
|
||||
self.log_std_init = log_std_init
|
||||
# Action distribution
|
||||
self.action_dist = make_proba_distribution(action_space, self.features_dim, use_sde=use_sde)
|
||||
|
||||
|
|
@ -130,22 +131,14 @@ class PPOPolicy(BasePolicy):
|
|||
|
||||
def reset_noise_net(self):
|
||||
self.action_dist.sample_weights(self.log_std)
|
||||
# weights_dist = Normal(th.zeros_like(self.noise_log_sigma), th.exp(self.noise_log_sigma))
|
||||
# self.noise_net = weights_dist.rsample()
|
||||
# noise = th.mm(state, weights)
|
||||
# variance = th.mm(state ** 2, sigma ** 2)
|
||||
# action_dist = Normal(mu, th.sqrt(variance))
|
||||
# # action_dist.log_prob((mu + noise).detach())
|
||||
# action_dist.log_prob(action)
|
||||
# # action_dist = Normal(mu_j + noise_j, sum of s_i * sigma_ij)
|
||||
# # log_prob = distribution.log_prob(self.noise_net)
|
||||
|
||||
def _build(self, learning_rate):
|
||||
self.mlp_extractor = MlpExtractor(self.features_dim, net_arch=self.net_arch,
|
||||
activation_fn=self.activation_fn, device=self.device)
|
||||
|
||||
if isinstance(self.action_dist, (DiagGaussianDistribution, StateDependentNoiseDistribution)):
|
||||
self.action_net, self.log_std = self.action_dist.proba_distribution_net(latent_dim=self.mlp_extractor.latent_dim_pi)
|
||||
self.action_net, self.log_std = self.action_dist.proba_distribution_net(latent_dim=self.mlp_extractor.latent_dim_pi,
|
||||
log_std_init=self.log_std_init)
|
||||
elif isinstance(self.action_dist, CategoricalDistribution):
|
||||
self.action_net = self.action_dist.proba_distribution_net(latent_dim=self.mlp_extractor.latent_dim_pi)
|
||||
|
||||
|
|
@ -177,10 +170,13 @@ class PPOPolicy(BasePolicy):
|
|||
|
||||
def _get_action_dist_from_latent(self, latent, obs, deterministic=False):
|
||||
mean_actions = self.action_net(latent)
|
||||
|
||||
if isinstance(self.action_dist, DiagGaussianDistribution):
|
||||
return self.action_dist.proba_distribution(mean_actions, self.log_std, deterministic=deterministic)
|
||||
|
||||
elif isinstance(self.action_dist, CategoricalDistribution):
|
||||
return self.action_dist.proba_distribution(mean_actions, deterministic=deterministic)
|
||||
|
||||
elif isinstance(self.action_dist, StateDependentNoiseDistribution):
|
||||
return self.action_dist.proba_distribution(mean_actions, self.log_std, obs, deterministic=deterministic)
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue