Implemented binding to PCA

This commit is contained in:
2023-08-21 16:43:41 +02:00
parent 497ee7e5fb
commit 9a4c43e233
6 changed files with 1959 additions and 10 deletions
+8 -7
View File
@@ -6,15 +6,17 @@ import torch as th
from gym import spaces
from torch.nn import functional as F
from stable_baselines3.common.on_policy_algorithm import OnPolicyAlgorithm
from stable_baselines3.common.policies import ActorCriticCnnPolicy, ActorCriticPolicy, BasePolicy, MultiInputActorCriticPolicy
# from stable_baselines3.common.on_policy_algorithm import OnPolicyAlgorithm
from ..common.on_policy_algorithm import BetterOnPolicyAlgorithm
# from stable_baselines3.common.policies import ActorCriticCnnPolicy, ActorCriticPolicy, BasePolicy, MultiInputActorCriticPolicy
from ..common.policies import ActorCriticPolicy, BasePolicy
from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, Schedule
from stable_baselines3.common.utils import explained_variance, get_schedule_fn
SelfPPO = TypeVar("SelfPPO", bound="PPO")
class PPO(OnPolicyAlgorithm):
class PPO(BetterOnPolicyAlgorithm):
"""
Proximal Policy Optimization algorithm (PPO) (clip version)
@@ -69,9 +71,7 @@ class PPO(OnPolicyAlgorithm):
"""
policy_aliases: Dict[str, Type[BasePolicy]] = {
"MlpPolicy": ActorCriticPolicy,
"CnnPolicy": ActorCriticCnnPolicy,
"MultiInputPolicy": MultiInputActorCriticPolicy,
"MlpPolicy": ActorCriticPolicy
}
def __init__(
@@ -92,6 +92,7 @@ class PPO(OnPolicyAlgorithm):
max_grad_norm: float = 0.5,
use_sde: bool = False,
sde_sample_freq: int = -1,
use_pca: bool = False,
target_kl: Optional[float] = None,
stats_window_size: int = 100,
tensorboard_log: Optional[str] = None,
@@ -113,6 +114,7 @@ class PPO(OnPolicyAlgorithm):
max_grad_norm=max_grad_norm,
use_sde=use_sde,
sde_sample_freq=sde_sample_freq,
use_pca=use_pca,
stats_window_size=stats_window_size,
tensorboard_log=tensorboard_log,
policy_kwargs=policy_kwargs,
@@ -315,4 +317,3 @@ class PPO(OnPolicyAlgorithm):
reset_num_timesteps=reset_num_timesteps,
progress_bar=progress_bar,
)