Implemented SDE

This commit is contained in:
2022-08-10 11:54:52 +02:00
parent 12e422aec7
commit 520dc98eb5
4 changed files with 61 additions and 13 deletions
+10 -1
View File
@@ -35,6 +35,8 @@ from stable_baselines3.common.torch_layers import (
NatureCNN,
)
from stable_baselines3.common.preprocessing import get_action_dim
from metastable_baselines.projections.w2_projection_layer import WassersteinProjectionLayer
from ..distributions import UniversalGaussianDistribution, make_proba_distribution
@@ -196,7 +198,14 @@ class ActorCriticPolicy(BasePolicy):
assert isinstance(
self.action_dist, StateDependentNoiseDistribution) or isinstance(
self.action_dist, UniversalGaussianDistribution), "reset_noise() is only available when using gSDE"
self.action_dist.sample_weights(self.log_std, batch_size=n_envs)
if isinstance(
self.action_dist, StateDependentNoiseDistribution):
self.action_dist.sample_weights(self.log_std, batch_size=n_envs)
if isinstance(
self.action_dist, UniversalGaussianDistribution):
self.action_dist.sample_weights(
get_action_dim(self.action_space), batch_size=n_envs)
def _build_mlp_extractor(self) -> None:
"""
+1 -1
View File
@@ -105,7 +105,7 @@ class PPO(GaussianRolloutCollectorAuxclass, OnPolicyAlgorithm):
target_kl: Optional[float] = None,
tensorboard_log: Optional[str] = None,
create_eval_env: bool = False,
policy_kwargs: Optional[Dict[str, Any]] = None,
policy_kwargs: Optional[Dict[str, Any]] = {},
verbose: int = 0,
seed: Optional[int] = None,
device: Union[th.device, str] = "auto",