From c3b0398d56f57bd2a300ad096ca7c120c79b489a Mon Sep 17 00:00:00 2001 From: Noah Dormann Date: Thu, 5 Dec 2019 08:40:28 +0100 Subject: [PATCH] Changed load so it still works when env not saved improved save function --- tests/test_save_load.py | 8 +++----- torchy_baselines/common/base_class.py | 15 +++++++++++---- 2 files changed, 14 insertions(+), 9 deletions(-) diff --git a/tests/test_save_load.py b/tests/test_save_load.py index 1ade709..bae9e0e 100644 --- a/tests/test_save_load.py +++ b/tests/test_save_load.py @@ -61,7 +61,7 @@ def test_save_load(model_class): # Check model.save("test_save.zip") del model - model = model_class.load("test_save") + model = model_class.load("test_save", env=env) # check if params are still the same after load new_params = model.get_policy_parameters() @@ -108,10 +108,8 @@ def test_exclude_include_saved_params(model_class): """ env = DummyVecEnv([lambda: IdentityEnvBox(10)]) - # create model - model = model_class('MlpPolicy', env, policy_kwargs=dict(net_arch=[16]), verbose=1, create_eval_env=True) - # set verbose as something different then standard settings - model.verbose = 2 + # create model, set verbose as 2, which is not standard + model = model_class('MlpPolicy', env, policy_kwargs=dict(net_arch=[16]), verbose=2, create_eval_env=True) # Check if exclude works model.save("test_save.zip", exclude=["verbose"]) diff --git a/torchy_baselines/common/base_class.py b/torchy_baselines/common/base_class.py index 0719771..a195580 100644 --- a/torchy_baselines/common/base_class.py +++ b/torchy_baselines/common/base_class.py @@ -271,7 +271,10 @@ class BaseRLModel(object): if env is None and "env" in data: env = data["env"] - model = cls(policy=data["policy_class"], env=env, _init_setup_model=True) + if env is not None: + model = cls(policy=data["policy_class"], env=env, _init_setup_model=True) + else: + model = cls(policy=data["policy_class"], env=env, _init_setup_model=False) model.__dict__.update(data) model.__dict__.update(kwargs) model.set_env(env) @@ -508,19 +511,23 @@ class BaseRLModel(object): :param path: (str) path to the file where the data should be saved :param exclude: ([str]) name of parameters that should be excluded in addition to the default one :param include: ([str]) name of parameters that might be excluded but should be included anyway - :return: """ - data = self.__dict__ + # 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 if include is not None: exclude = [param_name for param_name in exclude if param_name not in include] # remove parameter entries of parameters which are to be excluded for param_name in exclude: - data.pop(param_name, None) + if param_name in data: + data.pop(param_name, None) params_to_save = self.get_policy_parameters() opt_params_to_save = self.get_opt_parameters()