From 2c5e41ec474b77bd584c813d8eec11d156762651 Mon Sep 17 00:00:00 2001 From: Antonin RAFFIN Date: Tue, 31 Mar 2020 18:13:03 +0200 Subject: [PATCH] Renaming (everything plural) for consistency --- torchy_baselines/common/distributions.py | 113 +++++++++++------------ 1 file changed, 54 insertions(+), 59 deletions(-) diff --git a/torchy_baselines/common/distributions.py b/torchy_baselines/common/distributions.py index 4c83abb..93b1930 100644 --- a/torchy_baselines/common/distributions.py +++ b/torchy_baselines/common/distributions.py @@ -50,32 +50,30 @@ class Distribution(object): def get_actions(self, deterministic: bool = False) -> th.Tensor: """ - Return an action according to the probabilty distribution. + Return actions according to the probabilty distribution. :param deterministic: (bool) :return: (th.Tensor) """ if deterministic: return self.mode() - else: - return self.sample() + return self.sample() def actions_from_params(self, *args, **kwargs) -> th.Tensor: """ - Returns a sample from the probabilty distribution + Returns samples from the probabilty distribution given its parameters. - :return: (th.Tensor) the action + :return: (th.Tensor) actions """ raise NotImplementedError def log_prob_from_params(self, *args, **kwargs) -> Tuple[th.Tensor, th.Tensor]: """ - Returns a sample and the associated log probabilty - from the probabilty distribution - given its parameters. + Returns samples and the associated log probabilties + from the probabilty distribution given its parameters. - :return: (th.Tuple[th.Tensor, th.Tensor]) action and log prob + :return: (th.Tuple[th.Tensor, th.Tensor]) actions and log prob """ raise NotImplementedError @@ -144,6 +142,7 @@ class DiagGaussianDistribution(Distribution): return self.distribution.mean def sample(self) -> th.Tensor: + # Reparametrization trick to pass gradients return self.distribution.rsample() def entropy(self) -> th.Tensor: @@ -166,20 +165,19 @@ class DiagGaussianDistribution(Distribution): :param log_std: (th.Tensor) :return: (Tuple[th.Tensor, th.Tensor]) """ - action = self.actions_from_params(mean_actions, log_std) - log_prob = self.log_prob(action) - return action, log_prob + actions = self.actions_from_params(mean_actions, log_std) + log_prob = self.log_prob(actions) + return actions, log_prob - def log_prob(self, action: th.Tensor) -> th.Tensor: + def log_prob(self, actions: th.Tensor) -> th.Tensor: """ - Get the log probabilty of an action given a distribution. - Note that you must call ``proba_distribution()`` method - before. + Get the log probabilties of actions according to the distribution. + Note that you must call ``proba_distribution()`` method before. - :param action: (th.Tensor) + :param actions: (th.Tensor) :return: (th.Tensor) """ - log_prob = self.distribution.log_prob(action) + log_prob = self.distribution.log_prob(actions) return sum_independent_dims(log_prob) @@ -196,7 +194,7 @@ class SquashedDiagGaussianDistribution(DiagGaussianDistribution): super(SquashedDiagGaussianDistribution, self).__init__(action_dim) # Avoid NaN (prevents division by zero or log of zero) self.epsilon = epsilon - self.gaussian_action = None + self.gaussian_actions = None def proba_distribution(self, mean_actions: th.Tensor, log_std: th.Tensor) -> 'SquashedDiagGaussianDistribution': @@ -204,9 +202,9 @@ class SquashedDiagGaussianDistribution(DiagGaussianDistribution): return self def mode(self) -> th.Tensor: - self.gaussian_action = self.distribution.mean + self.gaussian_actions = self.distribution.mean # Squash the output - return th.tanh(self.gaussian_action) + return th.tanh(self.gaussian_actions) def entropy(self) -> Optional[th.Tensor]: # No analytical form, @@ -214,29 +212,30 @@ class SquashedDiagGaussianDistribution(DiagGaussianDistribution): return None def sample(self) -> th.Tensor: - self.gaussian_action = self.distribution.rsample() - return th.tanh(self.gaussian_action) + # Reparametrization trick to pass gradients + self.gaussian_actions = self.distribution.rsample() + return th.tanh(self.gaussian_actions) def log_prob_from_params(self, mean_actions: th.Tensor, log_std: th.Tensor) -> Tuple[th.Tensor, th.Tensor]: action = self.actions_from_params(mean_actions, log_std) - log_prob = self.log_prob(action, self.gaussian_action) + log_prob = self.log_prob(action, self.gaussian_actions) return action, log_prob - def log_prob(self, action: th.Tensor, - gaussian_action: Optional[th.Tensor] = None) -> th.Tensor: + def log_prob(self, actions: th.Tensor, + gaussian_actions: Optional[th.Tensor] = None) -> th.Tensor: # Inverse tanh # Naive implementation (not stable): 0.5 * torch.log((1 + x) / (1 - x)) # We use numpy to avoid numerical instability - if gaussian_action is None: + if gaussian_actions is None: # It will be clipped to avoid NaN when inversing tanh - gaussian_action = TanhBijector.inverse(action) + gaussian_actions = TanhBijector.inverse(actions) # Log likelihood for a Gaussian distribution - log_prob = super(SquashedDiagGaussianDistribution, self).log_prob(gaussian_action) + log_prob = super(SquashedDiagGaussianDistribution, self).log_prob(gaussian_actions) # Squash correction (from original SAC implementation) # this comes from the fact that tanh is bijective and differentiable - log_prob -= th.sum(th.log(1 - action ** 2 + self.epsilon), dim=1) + log_prob -= th.sum(th.log(1 - actions ** 2 + self.epsilon), dim=1) return log_prob @@ -258,7 +257,8 @@ class CategoricalDistribution(Distribution): it will be the logits of the Categorical distribution. You can then get probabilties using a softmax. - :param latent_dim: (int) Dimension og the last layer of the policy (before the action layer) + :param latent_dim: (int) Dimension of the last layer + of the policy network (before the action layer) :return: (nn.Linear) """ action_logits = nn.Linear(latent_dim, self.action_dim) @@ -284,13 +284,12 @@ class CategoricalDistribution(Distribution): return self.get_actions(deterministic=deterministic) def log_prob_from_params(self, action_logits: th.Tensor) -> Tuple[th.Tensor, th.Tensor]: - action = self.actions_from_params(action_logits) - log_prob = self.log_prob(action) - return action, log_prob + actions = self.actions_from_params(action_logits) + log_prob = self.log_prob(actions) + return actions, log_prob - def log_prob(self, action: th.Tensor) -> th.Tensor: - log_prob = self.distribution.log_prob(action) - return log_prob + def log_prob(self, actions: th.Tensor) -> th.Tensor: + return self.distribution.log_prob(actions) class StateDependentNoiseDistribution(Distribution): @@ -373,7 +372,9 @@ class StateDependentNoiseDistribution(Distribution): """ std = self.get_std(log_std) self.weights_dist = Normal(th.zeros_like(std), std) + # Reparametrization trick to pass gradients self.exploration_mat = self.weights_dist.rsample() + # Pre-compute matrices in case of parallel exploration self.exploration_matrices = self.weights_dist.rsample((batch_size,)) def proba_distribution_net(self, latent_dim: int, log_std_init: float = -2.0, @@ -419,17 +420,11 @@ class StateDependentNoiseDistribution(Distribution): self.distribution = Normal(mean_actions, th.sqrt(variance + self.epsilon)) return self - def get_actions(self, deterministic: bool = False) -> th.Tensor: - if deterministic: - return self.mode() - else: - return self.sample(self._latent_sde) - def mode(self) -> th.Tensor: - action = self.distribution.mean + actions = self.distribution.mean if self.bijector is not None: - return self.bijector.forward(action) - return action + return self.bijector.forward(actions) + return actions def get_noise(self, latent_sde: th.Tensor) -> th.Tensor: latent_sde = latent_sde if self.learn_features else latent_sde.detach() @@ -443,12 +438,12 @@ class StateDependentNoiseDistribution(Distribution): noise = th.bmm(latent_sde, self.exploration_matrices) return noise.squeeze(1) - def sample(self, latent_sde: th.Tensor) -> th.Tensor: - noise = self.get_noise(latent_sde) - action = self.distribution.mean + noise + def sample(self) -> th.Tensor: + noise = self.get_noise(self._latent_sde) + actions = self.distribution.mean + noise if self.bijector is not None: - return self.bijector.forward(action) - return action + return self.bijector.forward(actions) + return actions def entropy(self) -> Optional[th.Tensor]: # No analytical form, @@ -468,23 +463,23 @@ class StateDependentNoiseDistribution(Distribution): def log_prob_from_params(self, mean_actions: th.Tensor, log_std: th.Tensor, latent_sde: th.Tensor) -> Tuple[th.Tensor, th.Tensor]: - action = self.actions_from_params(mean_actions, log_std, latent_sde) - log_prob = self.log_prob(action) - return action, log_prob + actions = self.actions_from_params(mean_actions, log_std, latent_sde) + log_prob = self.log_prob(actions) + return actions, log_prob - def log_prob(self, action: th.Tensor) -> th.Tensor: + def log_prob(self, actions: th.Tensor) -> th.Tensor: if self.bijector is not None: - gaussian_action = self.bijector.inverse(action) + gaussian_actions = self.bijector.inverse(actions) else: - gaussian_action = action + gaussian_actions = actions # log likelihood for a gaussian - log_prob = self.distribution.log_prob(gaussian_action) + log_prob = self.distribution.log_prob(gaussian_actions) # Sum along action dim log_prob = sum_independent_dims(log_prob) if self.bijector is not None: # Squash correction (from original SAC implementation) - log_prob -= th.sum(self.bijector.log_prob_correction(gaussian_action), dim=1) + log_prob -= th.sum(self.bijector.log_prob_correction(gaussian_actions), dim=1) return log_prob