This commit is contained in:
Antonin Raffin 2019-10-29 18:43:16 +01:00
parent 9e8f6e0020
commit 0174ec269e
3 changed files with 12 additions and 15 deletions

View file

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

View file

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

View file

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