From 1f0443f332821323403548f5994b2700eb943f23 Mon Sep 17 00:00:00 2001 From: Adam Gleave Date: Thu, 2 Jul 2020 18:49:59 -0700 Subject: [PATCH] Review base_class --- stable_baselines3/common/base_class.py | 100 +++++++++++++------------ 1 file changed, 53 insertions(+), 47 deletions(-) diff --git a/stable_baselines3/common/base_class.py b/stable_baselines3/common/base_class.py index 4453231..f66220e 100644 --- a/stable_baselines3/common/base_class.py +++ b/stable_baselines3/common/base_class.py @@ -1,5 +1,7 @@ +"""Abstract base classes for RL algorithms.""" + import time -from typing import Union, Type, Optional, Dict, Any, List, Tuple, Callable +from typing import Union, Type, Optional, Dict, Any, Iterable, List, Tuple, Callable from abc import ABC, abstractmethod from collections import deque import pathlib @@ -23,12 +25,30 @@ from stable_baselines3.common.monitor import Monitor from stable_baselines3.common.noise import ActionNoise +def maybe_make_env(env: Union[GymEnv, str, None], monitor_wrapper: bool, verbose: int) -> Optional[GymEnv]: + """If env is a string, make the environment; otherwise, return env. + + :param env: (Union[GymEnv, str, None]) The environment to learn from. + :param monitor_wrapper: (bool) Whether to wrap env in a Monitor when creating env. + :param verbose: (int) logging verbosity + :return A Gym (vector) environment. + """ + if isinstance(env, str): + if verbose >= 1: + print(f"Creating environment from the given name '{env}'") + env = gym.make(env) + if monitor_wrapper: + env = Monitor(env, filename=None) + + return env + + class BaseAlgorithm(ABC): """ The base of RL algorithms :param policy: (Type[BasePolicy]) Policy object - :param env: (Union[GymEnv, str]) The environment to learn from + :param env: (Union[GymEnv, str, None]) The environment to learn from (if registered in Gym, can be str. Can be None for loading trained models) :param policy_base: (Type[BasePolicy]) The base policy used by this method :param learning_rate: (float or callable) learning rate for the optimizer, @@ -54,7 +74,7 @@ class BaseAlgorithm(ABC): def __init__(self, policy: Type[BasePolicy], - env: Union[GymEnv, str], + env: Union[GymEnv, str, None], policy_base: Type[BasePolicy], learning_rate: Union[float, Callable], policy_kwargs: Dict[str, Any] = None, @@ -116,18 +136,9 @@ class BaseAlgorithm(ABC): if env is not None: if isinstance(env, str): if create_eval_env: - eval_env = gym.make(env) - if monitor_wrapper: - eval_env = Monitor(eval_env, filename=None) - self.eval_env = DummyVecEnv([lambda: eval_env]) - if self.verbose >= 1: - print("Creating environment from the given name, wrapped in a DummyVecEnv.") - - env = gym.make(env) - if monitor_wrapper: - env = Monitor(env, filename=None) - env = DummyVecEnv([lambda: env]) + self.eval_env = maybe_make_env(env, monitor_wrapper, self.verbose) + env = maybe_make_env(env, monitor_wrapper, self.verbose) env = self._wrap_env(env) self.observation_space = env.observation_space @@ -136,8 +147,8 @@ class BaseAlgorithm(ABC): self.env = env if not support_multi_env and self.n_envs > 1: - raise ValueError("Error: the model does not support multiple envs requires a single vectorized" - " environment.") + raise ValueError("Error: the model does not support multiple envs; it requires " + "a single vectorized environment.") def _wrap_env(self, env: GymEnv) -> VecEnv: if not isinstance(env, VecEnv): @@ -153,10 +164,7 @@ class BaseAlgorithm(ABC): @abstractmethod def _setup_model(self) -> None: - """ - Create networks, buffer and optimizers - """ - raise NotImplementedError() + """Create networks, buffer and optimizers.""" def _get_eval_env(self, eval_env: Optional[GymEnv]) -> Optional[GymEnv]: """ @@ -238,7 +246,7 @@ class BaseAlgorithm(ABC): def get_torch_variables(self) -> Tuple[List[str], List[str]]: """ - Get the name of the torch variable that will be saved. + Get the name of the torch variables that will be saved. ``th.save`` and ``th.load`` will be used with the right device instead of the default pickling strategy. @@ -263,10 +271,9 @@ class BaseAlgorithm(ABC): Return a trained model. :param total_timesteps: (int) The total number of samples (env steps) to train on - :param callback: (function (dict, dict)) -> boolean function called at every steps with state of the algorithm. - It takes the local and global variables. If it returns False, training is aborted. + :param callback: (MaybeCallback) callback(s) called at every step with state of the algorithm. :param log_interval: (int) The number of timesteps before logging. - :param tb_log_name: (str) the name of the run for tensorboard log + :param tb_log_name: (str) the name of the run for TensorBoard logging :param eval_env: (gym.Env) Environment that will be used to evaluate the agent :param eval_freq: (int) Evaluate the agent every ``eval_freq`` timesteps (this may vary a little) :param n_eval_episodes: (int) Number of episode to evaluate the agent @@ -274,7 +281,6 @@ class BaseAlgorithm(ABC): :param reset_num_timesteps: (bool) whether or not to reset the current timestep number (used in logging) :return: (BaseAlgorithm) the trained model """ - raise NotImplementedError() def predict(self, observation: np.ndarray, state: Optional[np.ndarray] = None, @@ -329,8 +335,6 @@ class BaseAlgorithm(ABC): # load parameters model.__dict__.update(data) model.__dict__.update(kwargs) - if not hasattr(model, "_setup_model") and len(params) > 0: - raise NotImplementedError(f"{cls} has no ``_setup_model()`` method") model._setup_model() # put state_dicts back in place @@ -366,14 +370,18 @@ class BaseAlgorithm(ABC): self.eval_env.seed(seed) def _init_callback(self, - callback: Union[None, Callable, List[BaseCallback], BaseCallback], + callback: MaybeCallback, eval_env: Optional[VecEnv] = None, eval_freq: int = 10000, n_eval_episodes: int = 5, log_path: Optional[str] = None) -> BaseCallback: """ - :param callback: (Union[callable, [BaseCallback], BaseCallback, None]) - :return: (BaseCallback) + :param callback: (MaybeCallback) Callback(s) called at every step with state of the algorithm. + :param eval_freq: (Optional[int]) How many steps between evaluations; if None, do not evaluate. + :param n_eval_episodes: (int) How many episodes to play per evaluation + :param n_eval_episodes: (int) Number of episodes to rollout during evaluation. + :param log_path: (Optional[str]) Path to a folder where the evaluations will be saved + :return: (BaseCallback) A hybrid callback calling `callback` and performing evaluation. """ # Convert a list of callbacks into a callback if isinstance(callback, list): @@ -396,7 +404,7 @@ class BaseAlgorithm(ABC): def _setup_learn(self, total_timesteps: int, eval_env: Optional[GymEnv], - callback: Union[None, Callable, List[BaseCallback], BaseCallback] = None, + callback: MaybeCallback = None, eval_freq: int = 10000, n_eval_episodes: int = 5, log_path: Optional[str] = None, @@ -407,11 +415,11 @@ class BaseAlgorithm(ABC): Initialize different variables needed for training. :param total_timesteps: (int) The total number of samples (env steps) to train on - :param eval_env: (Optional[GymEnv]) - :param callback: (Union[None, BaseCallback, List[BaseCallback, Callable]]) + :param eval_env: (Optional[VecEnv]) Environment to use for evaluation. + :param callback: (MaybeCallback) Callback(s) called at every step with state of the algorithm. :param eval_freq: (int) How many steps between evaluations :param n_eval_episodes: (int) How many episodes to play per evaluation - :param log_path (Optional[str]): Path to a log folder + :param log_path: (Optional[str]) Path to a folder where the evaluations will be saved :param reset_num_timesteps: (bool) Whether to reset or not the ``num_timesteps`` attribute :param tb_log_name: (str) the name of the run for tensorboard log :return: (Tuple[int, BaseCallback]) @@ -480,8 +488,8 @@ class BaseAlgorithm(ABC): def save( self, path: Union[str, pathlib.Path, io.BufferedIOBase], - exclude: Optional[List[str]] = None, - include: Optional[List[str]] = None, + exclude: Optional[Iterable[str]] = None, + include: Optional[Iterable[str]] = None, ) -> None: """ Save all the attributes of the object and the model parameters in a zip-file. @@ -492,16 +500,15 @@ class BaseAlgorithm(ABC): """ # copy parameter list so we don't mutate the original dict data = self.__dict__.copy() - # use standard list of excluded parameters if none given - if exclude is None: - exclude = self.excluded_save_params() - else: - # append standard exclude params to the given params - exclude.extend([param for param in self.excluded_save_params() if param not in exclude]) - # do not exclude params if they are specifically included + # Exclude is union of specified parameters (if any) and standard exclusions + if exclude is None: + exclude = [] + exclude = set(exclude).union(self.excluded_save_params()) + + # Do not exclude params if they are specifically included if include is not None: - exclude = [param_name for param_name in exclude if param_name not in include] + exclude = exclude.difference(include) state_dicts_names, tensors_names = self.get_torch_variables() # any params that are in the save vars must not be saved by data @@ -509,12 +516,11 @@ class BaseAlgorithm(ABC): for torch_var in torch_variables: # we need to get only the name of the top most module as we'll remove that var_name = torch_var.split('.')[0] - exclude.append(var_name) + exclude.add(var_name) # Remove parameter entries of parameters which are to be excluded for param_name in exclude: - if param_name in data: - data.pop(param_name, None) + data.pop(param_name, None) # Build dict of tensor variables tensors = None