diff --git a/torchy_baselines/a2c/a2c.py b/torchy_baselines/a2c/a2c.py index 4c24c7f..6e0a1f2 100644 --- a/torchy_baselines/a2c/a2c.py +++ b/torchy_baselines/a2c/a2c.py @@ -7,11 +7,11 @@ import torch.nn.functional as F from torchy_baselines.common.utils import explained_variance from torchy_baselines.common import logger from torchy_baselines.common.type_aliases import GymEnv +from torchy_baselines.common.callbacks import BaseCallback from torchy_baselines.ppo.ppo import PPO from torchy_baselines.ppo.policies import PPOPolicy - class A2C(PPO): """ Advantage Actor Critic (A2C) @@ -89,14 +89,14 @@ class A2C(PPO): if _init_setup_model: self._setup_model() - def _setup_model(self): + def _setup_model(self) -> None: super(A2C, self)._setup_model() if self.use_rms_prop: self.policy.optimizer = th.optim.RMSprop(self.policy.parameters(), lr=self.learning_rate(1), alpha=0.99, eps=self.rms_prop_eps, weight_decay=0) - def train(self, gradient_steps: int, batch_size=None): + def train(self, gradient_steps: int, batch_size: Optional[int] = None) -> None: # Update optimizer learning rate self._update_learning_rate(self.policy.optimizer) # A2C with gradient_steps > 1 does not make sense @@ -153,9 +153,16 @@ class A2C(PPO): if hasattr(self.policy, 'log_std'): logger.logkv("std", th.exp(self.policy.log_std).mean().item()) - def learn(self, total_timesteps, callback=None, log_interval=100, - eval_env=None, eval_freq=-1, n_eval_episodes=5, - tb_log_name="A2C", eval_log_path=None, reset_num_timesteps=True): + def learn(self, + total_timesteps: int, + callback: Optional[BaseCallback] = None, + log_interval: int = 100, + eval_env: Optional[GymEnv] = None, + eval_freq: int = -1, + n_eval_episodes: int = 5, + tb_log_name: str = "A2C", + eval_log_path: Optional[str] = None, + reset_num_timesteps: bool = True) -> 'A2C': return super(A2C, self).learn(total_timesteps=total_timesteps, callback=callback, log_interval=log_interval, eval_env=eval_env, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes, diff --git a/torchy_baselines/cem_rl/cem.py b/torchy_baselines/cem_rl/cem.py index ee4f484..d1d221c 100644 --- a/torchy_baselines/cem_rl/cem.py +++ b/torchy_baselines/cem_rl/cem.py @@ -1,3 +1,5 @@ +from typing import Type, Tuple, Optional, List + import numpy as np @@ -21,9 +23,16 @@ class CEM(object): :param antithetic: (bool) Use a finite difference like method for sampling (mu + epsilon, mu - epsilon) """ - def __init__(self, num_params, mu_init=None, sigma_init=1e-3, - pop_size=256, damping_init=1e-3, damping_final=1e-5, - parents=None, elitism=False, antithetic=False): + def __init__(self, + num_params: int, + mu_init: Optional[np.ndarray] = None, + sigma_init: float = 1e-3, + pop_size: int = 256, + damping_init: float = 1e-3, + damping_final: float = 1e-5, + parents: Optional[int] = None, + elitism: bool = False, + antithetic: bool = False): super(CEM, self).__init__() self.num_params = num_params @@ -66,7 +75,7 @@ class CEM(object): for i in range(1, self.parents + 1)]) self.weights /= self.weights.sum() - def ask(self, pop_size): + def ask(self, pop_size: int) -> List[np.ndarray]: """ Returns a list of candidates parameters @@ -87,7 +96,7 @@ class CEM(object): return individuals - def tell(self, solutions, scores): + def tell(self, solutions: List[np.ndarray], scores: List[float]) -> None: """ Updates the distribution @@ -114,7 +123,7 @@ class CEM(object): self.elite = solutions[idx_sorted[0]] self.elite_score = scores[idx_sorted[0]] - def get_distrib_params(self): + def get_distrib_params(self) -> Tuple[np.ndarray, np.ndarray]: """ Returns the parameters of the distribution: the mean and standard deviation. diff --git a/torchy_baselines/ppo/ppo.py b/torchy_baselines/ppo/ppo.py index 6f9a483..a7dcf0b 100644 --- a/torchy_baselines/ppo/ppo.py +++ b/torchy_baselines/ppo/ppo.py @@ -121,7 +121,7 @@ class PPO(BaseRLModel): if _init_setup_model: self._setup_model() - def _setup_model(self): + def _setup_model(self) -> None: self._setup_learning_rate() # TODO: preprocessing: one hot vector for obs discrete state_dim = self.observation_space.shape[0] @@ -284,9 +284,16 @@ class PPO(BaseRLModel): if hasattr(self.policy, 'log_std'): logger.logkv("std", th.exp(self.policy.log_std).mean().item()) - def learn(self, total_timesteps, callback=None, log_interval=1, - eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="PPO", - eval_log_path=None, reset_num_timesteps=True): + def learn(self, + total_timesteps: int, + callback: Optional[BaseCallback] = None, + log_interval: int = 1, + eval_env: Optional[GymEnv] = None, + eval_freq: int = -1, + n_eval_episodes: int = 5, + tb_log_name: str = "PPO", + eval_log_path: Optional[str] = None, + reset_num_timesteps: bool = True) -> 'PPO': episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq, n_eval_episodes, eval_log_path, reset_num_timesteps)