refactored the assets in test_save_load

fixed base_class 'params.pth'
This commit is contained in:
Noah Dormann 2019-11-21 15:44:57 +01:00
parent 26f31fd25b
commit 526c37bf1f
2 changed files with 16 additions and 13 deletions

View file

@ -38,15 +38,11 @@ def test_save_load(model_class):
# Update model parameters with the new random values
model.load_parameters(random_params, opt_params)
# Get items that are the same in params and new_params
new_params = model.get_policy_parameters()
shared_items = {k: params[k] for k in params if k in new_params and th.all(th.eq(params[k], new_params[k]))}
# Check that the there are at least some parameters new random parameters
#for k in params.key():
# assert not th.allclose(params[k], new_params[k])
assert not len(shared_items) == len(new_params), "Selected actions did not change " \
"after changing model parameters."
# Check that all params are different now
for k in params:
assert not th.allclose(params[k], new_params[k]), "Selected actions did not change " \
"after changing model parameters."
params = new_params
@ -56,9 +52,16 @@ def test_save_load(model_class):
del model
model = model_class.load("test_save")
#check if params are still the same after load
# check if params are still the same after load
new_params = model.get_policy_parameters()
shared_items = {k: params[k] for k in params if k in new_params and th.all(th.eq(params[k], new_params[k]))}
# Check that at least some actions are chosen different now
assert len(shared_items) == len(new_params), "Parameters not the same after save and load."
# Check that all params are the same as before save load procedure now
for k in params:
assert th.allclose(params[k], new_params[k]), "Model parameters not the same after save and load."
# check if optimizer params are still the same after load
new_opt_params = model.get_opt_parameters()
# check if keys are the same
assert opt_params.keys() == new_opt_params.keys()
# check if values are the same: don't know how to to that
os.remove("test_save.zip")

View file

@ -505,7 +505,7 @@ class BaseRLModel(object):
if data is not None:
file_.writestr("data", serialized_data)
if params is not None:
with file_.open('param.pth', mode="w") as param_file:
with file_.open('params.pth', mode="w") as param_file:
th.save(params, param_file)
if opt_params is not None:
for file_name, dict in opt_params.items():