Refactored doc-strings

This commit is contained in:
Noah Dormann 2019-11-28 16:30:13 +01:00
parent 7ce610fade
commit 6928879f5a
4 changed files with 18 additions and 35 deletions

View file

@ -56,12 +56,10 @@ class BaseRLModel(object):
self.action_space = None self.action_space = None
self.n_envs = None self.n_envs = None
self.num_timesteps = 0 self.num_timesteps = 0
self.params = None
self.eval_env = None self.eval_env = None
self.replay_buffer = None self.replay_buffer = None
self.seed = seed self.seed = seed
self.action_noise = None self.action_noise = None
self.params = None
# Track the training progress (from 1 to 0) # Track the training progress (from 1 to 0)
# this is used to update the learning rate # this is used to update the learning rate
self._current_progress = 1 self._current_progress = 1
@ -176,16 +174,6 @@ class BaseRLModel(object):
""" """
pass pass
def get_parameter_list(self):
"""
Get pytorch Variables of model's parameters
This includes all variables necessary for continuing training (saving / loading).
:return: (list) List of pytorch Variables
"""
return self.params
def get_policy_parameters(self): def get_policy_parameters(self):
""" """
Get current model policy parameters as dictionary of variable name -> tensors. Get current model policy parameters as dictionary of variable name -> tensors.
@ -253,14 +241,8 @@ class BaseRLModel(object):
def load_parameters(self, load_dict, opt_params=None): def load_parameters(self, load_dict, opt_params=None):
""" """
Load model parameters from a dictionary Load model parameters from a dictionary
load_dict should contain all keys from torch.model.state_dict()
Dictionary should contain all entries of torch model.state_dict() If opt_params are given this does also load agent's optimizer-parameters, but can only be handled in child classes.
This does not load agent's hyper-parameters.
.. warning::
This function does not update trainer/optimizer variables (e.g. momentum).
As such training after using this function may lead to less-than-optimal results.
:param load_dict: (dict) dict of parameters from model.state_dict() :param load_dict: (dict) dict of parameters from model.state_dict()
@ -273,11 +255,11 @@ class BaseRLModel(object):
@classmethod @classmethod
def load(cls, load_path, env=None, **kwargs): def load(cls, load_path, env=None, **kwargs):
""" """
Load the model from file Load the model from a zip-file
:param load_path: (str) the saved parameter location :param load_path: (str) the location of the saved data
:param env: (Gym Envrionment) the new environment to run the loaded model on :param env: (Gym Envrionment) the new environment to run the loaded model on
(can be None if you only need prediction from a trained model) (can be None if you only need prediction from a trained model) has priority over any saved environment
:param kwargs: extra arguments to change the model when loading :param kwargs: extra arguments to change the model when loading
""" """
data, params, opt_params = cls._load_from_file(load_path) data, params, opt_params = cls._load_from_file(load_path)
@ -287,7 +269,9 @@ class BaseRLModel(object):
"Stored kwargs: {}, specified kwargs: {}".format(data['policy_kwargs'], "Stored kwargs: {}, specified kwargs: {}".format(data['policy_kwargs'],
kwargs['policy_kwargs'])) kwargs['policy_kwargs']))
model = cls(policy=data["policy_class"], env=data["env"], _init_setup_model=True) if env is None and "env" in data:
env = data["env"]
model = cls(policy=data["policy_class"], env=env, _init_setup_model=True)
model.__dict__.update(data) model.__dict__.update(data)
model.__dict__.update(kwargs) model.__dict__.update(kwargs)
model.set_env(env) model.set_env(env)
@ -298,7 +282,7 @@ class BaseRLModel(object):
def _load_from_file(load_path, load_data=True): def _load_from_file(load_path, load_data=True):
""" Load model data from a .zip archive """ Load model data from a .zip archive
:param load_path: (str or file-like) Where to load the model from :param load_path: (str) Where to load the model from
:param load_data: (bool) Whether we should load and return data :param load_data: (bool) Whether we should load and return data
(class parameters). Mainly used by 'load_parameters' to only load model parameters (weights) (class parameters). Mainly used by 'load_parameters' to only load model parameters (weights)
:return: (dict),(dict),(dict) Class parameters, model parameters (state_dict) and dict of optimizer parameters (dict of state_dict) :return: (dict),(dict),(dict) Class parameters, model parameters (state_dict) and dict of optimizer parameters (dict of state_dict)
@ -477,7 +461,7 @@ class BaseRLModel(object):
def _save_to_file_zip(save_path, data=None, params=None, opt_params=None): def _save_to_file_zip(save_path, data=None, params=None, opt_params=None):
"""Save model to a zip archive """Save model to a zip archive
:param save_path: (str or file-like) Where to store the model :param save_path: (str) Where to store the model
:param data: (dict) Class parameters being stored :param data: (dict) Class parameters being stored
:param params: (dict) Model parameters being stored expected to be state_dict :param params: (dict) Model parameters being stored expected to be state_dict
:param opt_params: (dict) Optimizer parameters being stored expected to contain an entry for every :param opt_params: (dict) Optimizer parameters being stored expected to contain an entry for every
@ -519,7 +503,7 @@ class BaseRLModel(object):
def save(self, path, exclude=None, include=None): def save(self, path, exclude=None, include=None):
""" """
saves all the params from init and pytorch params in a file for continuous learning saves all the params from init and pytorch params in a zip-file for continuous learning
:param path: (str) path to the file where the data should be saved :param path: (str) path to the file where the data should be saved
:param exclude: (list) name of parameters that should be excluded, use standard exclude params if None :param exclude: (list) name of parameters that should be excluded, use standard exclude params if None

View file

@ -315,7 +315,7 @@ class PPO(BaseRLModel):
def get_opt_parameters(self): def get_opt_parameters(self):
""" """
returns a dict of all the optimizers and their parameters Returns a dict of all the optimizers and their parameters
:return: (Dict) of optimizer names and their state_dict :return: (Dict) of optimizer names and their state_dict
""" """
@ -324,12 +324,11 @@ class PPO(BaseRLModel):
def load_parameters(self, load_dict, opt_params): def load_parameters(self, load_dict, opt_params):
""" """
Load model parameters and optimizer parameters from a dictionary Load model parameters and optimizer parameters from a dictionary
Dictionary should be of shape torch model.state_dict() load_dict should contain all keys from torch.model.state_dict()
This does not load agent's hyper-parameters. This does not load agent's hyper-parameters.
:param load_dict: (dict) dict of parameters from model.state_dict() :param load_dict: (dict) dict of parameters from model.state_dict()
:param opt_params: (dict of dicts) dict of optimizer state_dicts should be handled in child_class :param opt_params: (dict of dicts) dict of optimizer state_dicts
""" """
self.policy.optimizer.load_state_dict(opt_params["opt"]) self.policy.optimizer.load_state_dict(opt_params["opt"])
self.policy.load_state_dict(load_dict) self.policy.load_state_dict(load_dict)

View file

@ -277,7 +277,7 @@ class SAC(BaseRLModel):
def get_opt_parameters(self): def get_opt_parameters(self):
""" """
returns a dict of all the optimizers and their parameters Returns a dict of all the optimizers and their parameters
:return: (Dict) of optimizer names and their state_dict :return: (Dict) of optimizer names and their state_dict
""" """
@ -290,7 +290,7 @@ class SAC(BaseRLModel):
def load_parameters(self, load_dict, opt_params): def load_parameters(self, load_dict, opt_params):
""" """
Load model parameters and optimizer parameters from a dictionary Load model parameters and optimizer parameters from a dictionary
Dictionary should be of shape torch model.state_dict() load_dict should contain all keys from torch.model.state_dict()
This does not load agent's hyper-parameters. This does not load agent's hyper-parameters.
:param load_dict: (dict) dict of parameters from model.state_dict() :param load_dict: (dict) dict of parameters from model.state_dict()

View file

@ -239,7 +239,7 @@ class TD3(BaseRLModel):
def get_opt_parameters(self): def get_opt_parameters(self):
""" """
returns a dict of all the optimizers and their parameters Returns a dict of all the optimizers and their parameters
:return: (Dict) of optimizer names and their state_dict :return: (Dict) of optimizer names and their state_dict
""" """
@ -248,7 +248,7 @@ class TD3(BaseRLModel):
def load_parameters(self, load_dict, opt_params): def load_parameters(self, load_dict, opt_params):
""" """
Load model parameters and optimizer parameters from a dictionary Load model parameters and optimizer parameters from a dictionary
Dictionary should be of shape torch model.state_dict() load_dict should contain all keys from torch.model.state_dict()
This does not load agent's hyper-parameters. This does not load agent's hyper-parameters.
:param load_dict: (dict) dict of parameters from model.state_dict() :param load_dict: (dict) dict of parameters from model.state_dict()