diff --git a/torchy_baselines/a2c/a2c.py b/torchy_baselines/a2c/a2c.py index 6b51227..51eb846 100644 --- a/torchy_baselines/a2c/a2c.py +++ b/torchy_baselines/a2c/a2c.py @@ -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()) diff --git a/torchy_baselines/common/distributions.py b/torchy_baselines/common/distributions.py index 2944d8c..5c9cdac 100644 --- a/torchy_baselines/common/distributions.py +++ b/torchy_baselines/common/distributions.py @@ -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) diff --git a/torchy_baselines/ppo/policies.py b/torchy_baselines/ppo/policies.py index 08dd7dd..25d0bf3 100644 --- a/torchy_baselines/ppo/policies.py +++ b/torchy_baselines/ppo/policies.py @@ -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)