mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-29 20:14:17 +00:00
Changed load so it still works when env not saved
improved save function
This commit is contained in:
parent
bea2799691
commit
c3b0398d56
2 changed files with 14 additions and 9 deletions
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Reference in a new issue