diff --git a/setup.py b/setup.py index 044d839..21fedaa 100644 --- a/setup.py +++ b/setup.py @@ -10,6 +10,7 @@ setup(name='torchy_baselines', 'gym[classic_control]>=0.11', 'numpy', 'torch>=1.4.0', + # For saving models 'cloudpickle', # For reading logs 'pandas', @@ -31,7 +32,7 @@ setup(name='torchy_baselines', # For spelling 'sphinxcontrib.spelling', # Type hints support - 'sphinx-autodoc-typehints' + # 'sphinx-autodoc-typehints' ], 'extra': [ # For render @@ -47,7 +48,7 @@ setup(name='torchy_baselines', license="MIT", long_description="", long_description_content_type='text/markdown', - version="0.2.3", + version="0.2.4", ) # python setup.py sdist diff --git a/torchy_baselines/__init__.py b/torchy_baselines/__init__.py index 8c66376..b228acb 100644 --- a/torchy_baselines/__init__.py +++ b/torchy_baselines/__init__.py @@ -4,4 +4,4 @@ from torchy_baselines.ppo import PPO from torchy_baselines.sac import SAC from torchy_baselines.td3 import TD3 -__version__ = "0.2.3" +__version__ = "0.2.4" diff --git a/torchy_baselines/common/base_class.py b/torchy_baselines/common/base_class.py index 0462002..eabe92e 100644 --- a/torchy_baselines/common/base_class.py +++ b/torchy_baselines/common/base_class.py @@ -176,7 +176,7 @@ class BaseRLModel(ABC): low, high = self.action_space.low, self.action_space.high return low + (0.5 * (scaled_action + 1.0) * (high - low)) - def _setup_learning_rate(self) -> None: + def _setup_lr_schedule(self) -> None: """Transform to callable if needed.""" self.lr_schedule = get_schedule_fn(self.learning_rate) diff --git a/torchy_baselines/ppo/policies.py b/torchy_baselines/ppo/policies.py index 8487f6d..d54e7dd 100644 --- a/torchy_baselines/ppo/policies.py +++ b/torchy_baselines/ppo/policies.py @@ -19,7 +19,7 @@ class PPOPolicy(BasePolicy): :param observation_space: (gym.spaces.Space) Observation space :param action_space: (gym.spaces.Space) Action space - :param lr_schedule: (callable) Learning rate schedule (could be constant) + :param lr_schedule: (Callable) Learning rate schedule (could be constant) :param net_arch: ([int or dict]) The specification of the policy and value networks. :param device: (str or th.device) Device on which the code should run. :param activation_fn: (nn.Module) Activation function diff --git a/torchy_baselines/ppo/ppo.py b/torchy_baselines/ppo/ppo.py index 5aa8f56..098aa3c 100644 --- a/torchy_baselines/ppo/ppo.py +++ b/torchy_baselines/ppo/ppo.py @@ -122,7 +122,7 @@ class PPO(BaseRLModel): self._setup_model() def _setup_model(self) -> None: - self._setup_learning_rate() + self._setup_lr_schedule() # TODO: preprocessing: one hot vector for obs discrete state_dim = self.observation_space.shape[0] if isinstance(self.action_space, spaces.Box): @@ -137,7 +137,7 @@ class PPO(BaseRLModel): self.rollout_buffer = RolloutBuffer(self.n_steps, state_dim, action_dim, self.device, gamma=self.gamma, gae_lambda=self.gae_lambda, n_envs=self.n_envs) self.policy = self.policy_class(self.observation_space, self.action_space, - self.learning_rate, use_sde=self.use_sde, device=self.device, + self.lr_schedule, use_sde=self.use_sde, device=self.device, **self.policy_kwargs) self.policy = self.policy.to(self.device) diff --git a/torchy_baselines/sac/sac.py b/torchy_baselines/sac/sac.py index 907aa05..8878c3f 100644 --- a/torchy_baselines/sac/sac.py +++ b/torchy_baselines/sac/sac.py @@ -118,7 +118,7 @@ class SAC(OffPolicyRLModel): self._setup_model() def _setup_model(self) -> None: - self._setup_learning_rate() + self._setup_lr_schedule() obs_dim, action_dim = self.observation_space.shape[0], self.action_space.shape[0] if self.seed is not None: self.set_random_seed(self.seed) diff --git a/torchy_baselines/td3/td3.py b/torchy_baselines/td3/td3.py index 331a6e0..70f3a5f 100644 --- a/torchy_baselines/td3/td3.py +++ b/torchy_baselines/td3/td3.py @@ -118,7 +118,7 @@ class TD3(OffPolicyRLModel): self._setup_model() def _setup_model(self) -> None: - self._setup_learning_rate() + self._setup_lr_schedule() obs_dim, action_dim = self.observation_space.shape[0], self.action_space.shape[0] self.set_random_seed(self.seed) self.replay_buffer = ReplayBuffer(self.buffer_size, obs_dim, action_dim, self.device)