Making MultivariateNormal Policies work (and porting Normal to

Independent)
This commit is contained in:
2022-07-15 15:03:51 +02:00
parent b1ed9fc2b8
commit ab557a8856
8 changed files with 84 additions and 57 deletions
+2 -2
View File
@@ -1,2 +1,2 @@
from ..trl_pg.policies import CnnPolicy, MlpPolicy, MultiInputPolicy
from ..trl_pg.trl_pg import TRL_PG
from ..ppo.policies import CnnPolicy, MlpPolicy, MultiInputPolicy
from .ppo import PPO
+4 -2
View File
@@ -95,6 +95,7 @@ class ActorCriticPolicy(BasePolicy):
normalize_images: bool = True,
optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam,
optimizer_kwargs: Optional[Dict[str, Any]] = None,
dist_kwargs: Optional[Dict[str, Any]] = None,
):
if optimizer_kwargs is None:
@@ -130,15 +131,16 @@ class ActorCriticPolicy(BasePolicy):
self.normalize_images = normalize_images
self.log_std_init = log_std_init
dist_kwargs = None
# Keyword arguments for gSDE distribution
if use_sde:
dist_kwargs = {
add_dist_kwargs = {
"full_std": full_std,
"squash_output": squash_output,
"use_expln": use_expln,
"learn_features": False,
}
for k in add_dist_kwargs:
dist_kwargs[k] = add_dist_kwargs[k]
if sde_net_arch is not None:
warnings.warn(
+9 -4
View File
@@ -26,7 +26,7 @@ from ..projections.kl_projection_layer import KLProjectionLayer
from ..misc.rollout_buffer import GaussianRolloutCollectorAuxclass
class TRL_PG(GaussianRolloutCollectorAuxclass, OnPolicyAlgorithm):
class PPO(GaussianRolloutCollectorAuxclass, OnPolicyAlgorithm):
"""
Differential Trust Region Layer (TRL) for Policy Gradient (PG)
@@ -248,7 +248,12 @@ class TRL_PG(GaussianRolloutCollectorAuxclass, OnPolicyAlgorithm):
q_dist = new_dist_like(
p_dist, rollout_data.means, rollout_data.stds)
proj_p = self.projection(p_dist, q_dist, self._global_steps)
log_prob = proj_p.log_prob(actions).sum(dim=1)
if isinstance(p_dist, th.distributions.Normal):
# Normal uses a weird mapping from dimensions into batch_shape
log_prob = proj_p.log_prob(actions).sum(dim=1)
else:
# UniversalGaussianDistribution instead uses Independent (or MultivariateNormal), which has a more rational dim mapping
log_prob = proj_p.log_prob(actions)
values = self.policy.value_net(latent_vf)
entropy = proj_p.entropy()
@@ -373,10 +378,10 @@ class TRL_PG(GaussianRolloutCollectorAuxclass, OnPolicyAlgorithm):
eval_env: Optional[GymEnv] = None,
eval_freq: int = -1,
n_eval_episodes: int = 5,
tb_log_name: str = "TRL_PG",
tb_log_name: str = "PPO",
eval_log_path: Optional[str] = None,
reset_num_timesteps: bool = True,
) -> "TRL_PG":
) -> "PPO":
return super().learn(
total_timesteps=total_timesteps,