Fixing bugs with w2 and sqrt_induced_gaussian
This commit is contained in:
@@ -99,6 +99,7 @@ class ActorCriticPolicy(BasePolicy):
|
||||
optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam,
|
||||
optimizer_kwargs: Optional[Dict[str, Any]] = None,
|
||||
dist_kwargs: Optional[Dict[str, Any]] = None,
|
||||
sqrt_induced_gaussian=False,
|
||||
):
|
||||
|
||||
if optimizer_kwargs is None:
|
||||
@@ -152,6 +153,8 @@ class ActorCriticPolicy(BasePolicy):
|
||||
self.use_sde = use_sde
|
||||
self.dist_kwargs = dist_kwargs
|
||||
|
||||
self.sqrt_induced_gaussian = sqrt_induced_gaussian
|
||||
|
||||
# Action distribution
|
||||
self.action_dist = make_proba_distribution(
|
||||
action_space, use_sde=use_sde, dist_kwargs=dist_kwargs)
|
||||
@@ -289,18 +292,6 @@ class ActorCriticPolicy(BasePolicy):
|
||||
"""
|
||||
mean_actions = self.action_net(latent_pi)
|
||||
|
||||
if isinstance(self.projection, WassersteinProjectionLayer):
|
||||
if isinstance(self.action_dist, UniversalGaussianDistribution):
|
||||
cov_sqrt = self.chol_net(latent_pi)
|
||||
dist = self.action_dist.proba_distribution_from_sqrt(
|
||||
mean_actions, cov_sqrt, latent_pi)
|
||||
mean, chol = get_mean_and_chol(dist, expand=False)
|
||||
self.chol = chol
|
||||
return dist
|
||||
else:
|
||||
raise Exception(
|
||||
'Need to use UniversalGaussianDistribution to use WassersteinProjection (uses sqrt-induced-cov)')
|
||||
|
||||
if isinstance(self.action_dist, DiagGaussianDistribution):
|
||||
return self.action_dist.proba_distribution(mean_actions, self.log_std)
|
||||
elif isinstance(self.action_dist, CategoricalDistribution):
|
||||
@@ -315,9 +306,17 @@ class ActorCriticPolicy(BasePolicy):
|
||||
elif isinstance(self.action_dist, StateDependentNoiseDistribution):
|
||||
return self.action_dist.proba_distribution(mean_actions, self.log_std, latent_pi)
|
||||
elif isinstance(self.action_dist, UniversalGaussianDistribution):
|
||||
chol = self.chol_net(latent_pi)
|
||||
self.chol = chol
|
||||
return self.action_dist.proba_distribution(mean_actions, chol, latent_pi)
|
||||
if self.sqrt_induced_gaussian:
|
||||
cov_sqrt = self.chol_net(latent_pi)
|
||||
dist = self.action_dist.proba_distribution_from_sqrt(
|
||||
mean_actions, cov_sqrt, latent_pi)
|
||||
mean, chol = get_mean_and_chol(dist, expand=False)
|
||||
self.chol = chol
|
||||
return dist
|
||||
else:
|
||||
chol = self.chol_net(latent_pi)
|
||||
self.chol = chol
|
||||
return self.action_dist.proba_distribution(mean_actions, chol, latent_pi)
|
||||
else:
|
||||
raise ValueError("Invalid action distribution")
|
||||
|
||||
|
||||
@@ -16,7 +16,7 @@ from stable_baselines3.common.callbacks import BaseCallback
|
||||
from stable_baselines3.common.utils import obs_as_tensor
|
||||
from stable_baselines3.common.vec_env import VecNormalize
|
||||
|
||||
from ..misc.distTools import new_dist_like
|
||||
from ..misc.distTools import new_dist_like, new_dist_like_from_sqrt
|
||||
|
||||
from ..projections.base_projection_layer import BaseProjectionLayer
|
||||
from ..projections.frob_projection_layer import FrobeniusProjectionLayer
|
||||
@@ -133,7 +133,9 @@ class PPO(GaussianRolloutCollectorAuxclass, OnPolicyAlgorithm):
|
||||
use_sde=use_sde,
|
||||
sde_sample_freq=sde_sample_freq,
|
||||
tensorboard_log=tensorboard_log,
|
||||
policy_kwargs=policy_kwargs,
|
||||
policy_kwargs=policy_kwargs |
|
||||
{'sqrt_induced_gaussian': isinstance(
|
||||
projection, WassersteinProjectionLayer)},
|
||||
verbose=verbose,
|
||||
device=device,
|
||||
create_eval_env=create_eval_env,
|
||||
@@ -245,8 +247,12 @@ class PPO(GaussianRolloutCollectorAuxclass, OnPolicyAlgorithm):
|
||||
latent_pi, latent_vf = pol.mlp_extractor(features)
|
||||
p = pol._get_action_dist_from_latent(latent_pi)
|
||||
p_dist = p.distribution
|
||||
q_dist = new_dist_like(
|
||||
p_dist, rollout_data.means, rollout_data.chols)
|
||||
if isinstance(self.projection, WassersteinProjectionLayer):
|
||||
q_dist = new_dist_like_from_sqrt(
|
||||
p_dist, rollout_data.means, rollout_data.chols)
|
||||
else:
|
||||
q_dist = new_dist_like(
|
||||
p_dist, rollout_data.means, rollout_data.chols)
|
||||
proj_p = self.projection(p_dist, q_dist, self._global_steps)
|
||||
if isinstance(p_dist, th.distributions.Normal):
|
||||
# Normal uses a weird mapping from dimensions into batch_shape
|
||||
|
||||
Reference in New Issue
Block a user