Making MultivariateNormal Policies work (and porting Normal to
Independent)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user