mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-17 22:30:59 +00:00
Review base_class
This commit is contained in:
parent
2affbd6856
commit
1f0443f332
1 changed files with 53 additions and 47 deletions
|
|
@ -1,5 +1,7 @@
|
||||||
|
"""Abstract base classes for RL algorithms."""
|
||||||
|
|
||||||
import time
|
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 abc import ABC, abstractmethod
|
||||||
from collections import deque
|
from collections import deque
|
||||||
import pathlib
|
import pathlib
|
||||||
|
|
@ -23,12 +25,30 @@ from stable_baselines3.common.monitor import Monitor
|
||||||
from stable_baselines3.common.noise import ActionNoise
|
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):
|
class BaseAlgorithm(ABC):
|
||||||
"""
|
"""
|
||||||
The base of RL algorithms
|
The base of RL algorithms
|
||||||
|
|
||||||
:param policy: (Type[BasePolicy]) Policy object
|
: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)
|
(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 policy_base: (Type[BasePolicy]) The base policy used by this method
|
||||||
:param learning_rate: (float or callable) learning rate for the optimizer,
|
:param learning_rate: (float or callable) learning rate for the optimizer,
|
||||||
|
|
@ -54,7 +74,7 @@ class BaseAlgorithm(ABC):
|
||||||
|
|
||||||
def __init__(self,
|
def __init__(self,
|
||||||
policy: Type[BasePolicy],
|
policy: Type[BasePolicy],
|
||||||
env: Union[GymEnv, str],
|
env: Union[GymEnv, str, None],
|
||||||
policy_base: Type[BasePolicy],
|
policy_base: Type[BasePolicy],
|
||||||
learning_rate: Union[float, Callable],
|
learning_rate: Union[float, Callable],
|
||||||
policy_kwargs: Dict[str, Any] = None,
|
policy_kwargs: Dict[str, Any] = None,
|
||||||
|
|
@ -116,18 +136,9 @@ class BaseAlgorithm(ABC):
|
||||||
if env is not None:
|
if env is not None:
|
||||||
if isinstance(env, str):
|
if isinstance(env, str):
|
||||||
if create_eval_env:
|
if create_eval_env:
|
||||||
eval_env = gym.make(env)
|
self.eval_env = maybe_make_env(env, monitor_wrapper, self.verbose)
|
||||||
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])
|
|
||||||
|
|
||||||
|
env = maybe_make_env(env, monitor_wrapper, self.verbose)
|
||||||
env = self._wrap_env(env)
|
env = self._wrap_env(env)
|
||||||
|
|
||||||
self.observation_space = env.observation_space
|
self.observation_space = env.observation_space
|
||||||
|
|
@ -136,8 +147,8 @@ class BaseAlgorithm(ABC):
|
||||||
self.env = env
|
self.env = env
|
||||||
|
|
||||||
if not support_multi_env and self.n_envs > 1:
|
if not support_multi_env and self.n_envs > 1:
|
||||||
raise ValueError("Error: the model does not support multiple envs requires a single vectorized"
|
raise ValueError("Error: the model does not support multiple envs; it requires "
|
||||||
" environment.")
|
"a single vectorized environment.")
|
||||||
|
|
||||||
def _wrap_env(self, env: GymEnv) -> VecEnv:
|
def _wrap_env(self, env: GymEnv) -> VecEnv:
|
||||||
if not isinstance(env, VecEnv):
|
if not isinstance(env, VecEnv):
|
||||||
|
|
@ -153,10 +164,7 @@ class BaseAlgorithm(ABC):
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def _setup_model(self) -> None:
|
def _setup_model(self) -> None:
|
||||||
"""
|
"""Create networks, buffer and optimizers."""
|
||||||
Create networks, buffer and optimizers
|
|
||||||
"""
|
|
||||||
raise NotImplementedError()
|
|
||||||
|
|
||||||
def _get_eval_env(self, eval_env: Optional[GymEnv]) -> Optional[GymEnv]:
|
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]]:
|
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
|
``th.save`` and ``th.load`` will be used with the right device
|
||||||
instead of the default pickling strategy.
|
instead of the default pickling strategy.
|
||||||
|
|
||||||
|
|
@ -263,10 +271,9 @@ class BaseAlgorithm(ABC):
|
||||||
Return a trained model.
|
Return a trained model.
|
||||||
|
|
||||||
:param total_timesteps: (int) The total number of samples (env steps) to train on
|
: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.
|
:param callback: (MaybeCallback) callback(s) called at every step with state of the algorithm.
|
||||||
It takes the local and global variables. If it returns False, training is aborted.
|
|
||||||
:param log_interval: (int) The number of timesteps before logging.
|
: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_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 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
|
: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)
|
:param reset_num_timesteps: (bool) whether or not to reset the current timestep number (used in logging)
|
||||||
:return: (BaseAlgorithm) the trained model
|
:return: (BaseAlgorithm) the trained model
|
||||||
"""
|
"""
|
||||||
raise NotImplementedError()
|
|
||||||
|
|
||||||
def predict(self, observation: np.ndarray,
|
def predict(self, observation: np.ndarray,
|
||||||
state: Optional[np.ndarray] = None,
|
state: Optional[np.ndarray] = None,
|
||||||
|
|
@ -329,8 +335,6 @@ class BaseAlgorithm(ABC):
|
||||||
# load parameters
|
# load parameters
|
||||||
model.__dict__.update(data)
|
model.__dict__.update(data)
|
||||||
model.__dict__.update(kwargs)
|
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()
|
model._setup_model()
|
||||||
|
|
||||||
# put state_dicts back in place
|
# put state_dicts back in place
|
||||||
|
|
@ -366,14 +370,18 @@ class BaseAlgorithm(ABC):
|
||||||
self.eval_env.seed(seed)
|
self.eval_env.seed(seed)
|
||||||
|
|
||||||
def _init_callback(self,
|
def _init_callback(self,
|
||||||
callback: Union[None, Callable, List[BaseCallback], BaseCallback],
|
callback: MaybeCallback,
|
||||||
eval_env: Optional[VecEnv] = None,
|
eval_env: Optional[VecEnv] = None,
|
||||||
eval_freq: int = 10000,
|
eval_freq: int = 10000,
|
||||||
n_eval_episodes: int = 5,
|
n_eval_episodes: int = 5,
|
||||||
log_path: Optional[str] = None) -> BaseCallback:
|
log_path: Optional[str] = None) -> BaseCallback:
|
||||||
"""
|
"""
|
||||||
:param callback: (Union[callable, [BaseCallback], BaseCallback, None])
|
:param callback: (MaybeCallback) Callback(s) called at every step with state of the algorithm.
|
||||||
:return: (BaseCallback)
|
: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
|
# Convert a list of callbacks into a callback
|
||||||
if isinstance(callback, list):
|
if isinstance(callback, list):
|
||||||
|
|
@ -396,7 +404,7 @@ class BaseAlgorithm(ABC):
|
||||||
def _setup_learn(self,
|
def _setup_learn(self,
|
||||||
total_timesteps: int,
|
total_timesteps: int,
|
||||||
eval_env: Optional[GymEnv],
|
eval_env: Optional[GymEnv],
|
||||||
callback: Union[None, Callable, List[BaseCallback], BaseCallback] = None,
|
callback: MaybeCallback = None,
|
||||||
eval_freq: int = 10000,
|
eval_freq: int = 10000,
|
||||||
n_eval_episodes: int = 5,
|
n_eval_episodes: int = 5,
|
||||||
log_path: Optional[str] = None,
|
log_path: Optional[str] = None,
|
||||||
|
|
@ -407,11 +415,11 @@ class BaseAlgorithm(ABC):
|
||||||
Initialize different variables needed for training.
|
Initialize different variables needed for training.
|
||||||
|
|
||||||
:param total_timesteps: (int) The total number of samples (env steps) to train on
|
:param total_timesteps: (int) The total number of samples (env steps) to train on
|
||||||
:param eval_env: (Optional[GymEnv])
|
:param eval_env: (Optional[VecEnv]) Environment to use for evaluation.
|
||||||
:param callback: (Union[None, BaseCallback, List[BaseCallback, Callable]])
|
:param callback: (MaybeCallback) Callback(s) called at every step with state of the algorithm.
|
||||||
:param eval_freq: (int) How many steps between evaluations
|
:param eval_freq: (int) How many steps between evaluations
|
||||||
:param n_eval_episodes: (int) How many episodes to play per evaluation
|
: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 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
|
:param tb_log_name: (str) the name of the run for tensorboard log
|
||||||
:return: (Tuple[int, BaseCallback])
|
:return: (Tuple[int, BaseCallback])
|
||||||
|
|
@ -480,8 +488,8 @@ class BaseAlgorithm(ABC):
|
||||||
def save(
|
def save(
|
||||||
self,
|
self,
|
||||||
path: Union[str, pathlib.Path, io.BufferedIOBase],
|
path: Union[str, pathlib.Path, io.BufferedIOBase],
|
||||||
exclude: Optional[List[str]] = None,
|
exclude: Optional[Iterable[str]] = None,
|
||||||
include: Optional[List[str]] = None,
|
include: Optional[Iterable[str]] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
Save all the attributes of the object and the model parameters in a zip-file.
|
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
|
# copy parameter list so we don't mutate the original dict
|
||||||
data = self.__dict__.copy()
|
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:
|
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()
|
state_dicts_names, tensors_names = self.get_torch_variables()
|
||||||
# any params that are in the save vars must not be saved by data
|
# any params that are in the save vars must not be saved by data
|
||||||
|
|
@ -509,11 +516,10 @@ class BaseAlgorithm(ABC):
|
||||||
for torch_var in torch_variables:
|
for torch_var in torch_variables:
|
||||||
# we need to get only the name of the top most module as we'll remove that
|
# we need to get only the name of the top most module as we'll remove that
|
||||||
var_name = torch_var.split('.')[0]
|
var_name = torch_var.split('.')[0]
|
||||||
exclude.append(var_name)
|
exclude.add(var_name)
|
||||||
|
|
||||||
# Remove parameter entries of parameters which are to be excluded
|
# Remove parameter entries of parameters which are to be excluded
|
||||||
for param_name in exclude:
|
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
|
# Build dict of tensor variables
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue