mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-29 20:14:17 +00:00
Implemented Changes suggested from Antonin-Raffin
Added Optimizer saving
This commit is contained in:
parent
a7655ca6e1
commit
fb5f192fc4
4 changed files with 101 additions and 65 deletions
|
|
@ -18,14 +18,6 @@ def test_td3():
|
|||
os.remove("test_save.pth")
|
||||
|
||||
|
||||
def test_a2c():
|
||||
model = A2C('MlpPolicy', 'CartPole-v1', policy_kwargs=dict(net_arch=[16]), verbose=1, create_eval_env=True)
|
||||
model.learn(total_timesteps=1000, eval_freq=500)
|
||||
model.save("test_save")
|
||||
model.load("test_save")
|
||||
os.remove("test_save.pth")
|
||||
|
||||
|
||||
def test_cemrl():
|
||||
model = CEMRL('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[16]), pop_size=2, n_grad=1,
|
||||
learning_starts=100, verbose=1, create_eval_env=True, action_noise=action_noise)
|
||||
|
|
@ -40,9 +32,9 @@ def test_cemrl():
|
|||
def test_onpolicy(model_class, env_id):
|
||||
model = model_class('MlpPolicy', env_id, policy_kwargs=dict(net_arch=[16]), verbose=1, create_eval_env=True)
|
||||
model.learn(total_timesteps=1000, eval_freq=500)
|
||||
# model.save("test_save")
|
||||
# model.load("test_save")
|
||||
# os.remove("test_save.pth")
|
||||
model.save("test_save")
|
||||
model.load("test_save")
|
||||
#os.remove("test_save.pth")
|
||||
|
||||
|
||||
def test_sac():
|
||||
|
|
|
|||
|
|
@ -18,6 +18,8 @@ def test_save_load(model_class):
|
|||
"""
|
||||
Test if 'save' and 'load' saves and loads model correctly
|
||||
and if 'load_parameters' and 'get_policy_parameters' work correctly
|
||||
|
||||
''warning does not test function of optimizer parameter load
|
||||
|
||||
:param model_class: (BaseRLModel) A RL model
|
||||
"""
|
||||
|
|
@ -26,20 +28,23 @@ def test_save_load(model_class):
|
|||
# create model
|
||||
model = model_class('MlpPolicy', env, policy_kwargs=dict(net_arch=[16]), verbose=1, create_eval_env=True)
|
||||
|
||||
# test action probability for given (obs, action) pair
|
||||
|
||||
# Get dictionary of current parameters
|
||||
params = deepcopy(model.get_policy_parameters())
|
||||
opt_params = deepcopy(model.get_opt_parameters())
|
||||
|
||||
# Modify all parameters to be random values
|
||||
random_params = dict((param_name, th.rand_like(param)) for param_name, param in params.items())
|
||||
# Update model parameters with the new zeroed values
|
||||
model.load_parameters(random_params)
|
||||
|
||||
# shared items
|
||||
# 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 at least some actions are chosen different now
|
||||
|
||||
# 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."
|
||||
|
||||
|
|
@ -48,9 +53,10 @@ def test_save_load(model_class):
|
|||
# Check
|
||||
model.learn(total_timesteps=1000, eval_freq=500)
|
||||
model.save("test_save.zip")
|
||||
model = model.load("test_save")
|
||||
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
|
||||
|
|
|
|||
|
|
@ -194,7 +194,8 @@ class BaseRLModel(object):
|
|||
"""
|
||||
return self.policy.state_dict()
|
||||
|
||||
def get_optim_parameters(self):
|
||||
@abstractmethod
|
||||
def get_opt_parameters(self):
|
||||
"""
|
||||
Get current model optimizer parameters as dictionary of variable names -> tensors
|
||||
:return: (OrderedDict) Dictionary of variable name -> tensor of model's optimizer parameters
|
||||
|
|
@ -249,7 +250,7 @@ class BaseRLModel(object):
|
|||
"""
|
||||
pass
|
||||
|
||||
def load_parameters(self, load_dict, exact_match=True):
|
||||
def load_parameters(self, load_dict, opt_params=None, exact_match=True):
|
||||
"""
|
||||
Load model parameters from a dictionary
|
||||
|
||||
|
|
@ -262,8 +263,11 @@ class BaseRLModel(object):
|
|||
As such training after using this function may lead to less-than-optimal results.
|
||||
|
||||
|
||||
:param load_path_or_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
|
||||
"""
|
||||
if opt_params is not None:
|
||||
raise ValueError("Optimizer Parameters where given but no overloaded load function exists for this class")
|
||||
self.policy.load_state_dict(load_dict)
|
||||
|
||||
@abstractmethod
|
||||
|
|
@ -276,7 +280,6 @@ class BaseRLModel(object):
|
|||
raise NotImplementedError()
|
||||
|
||||
@classmethod
|
||||
@abstractmethod
|
||||
def load(cls, load_path, env=None, **kwargs):
|
||||
"""
|
||||
Load the model from file
|
||||
|
|
@ -286,7 +289,7 @@ class BaseRLModel(object):
|
|||
(can be None if you only need prediction from a trained model)
|
||||
:param kwargs: extra arguments to change the model when loading
|
||||
"""
|
||||
data, params = cls._load_from_file(load_path)
|
||||
data, params, opt_params = cls._load_from_file(load_path)
|
||||
|
||||
if 'policy_kwargs' in kwargs and kwargs['policy_kwargs'] != data['policy_kwargs']:
|
||||
raise ValueError("The specified policy kwargs do not equal the stored policy kwargs."
|
||||
|
|
@ -297,7 +300,7 @@ class BaseRLModel(object):
|
|||
model.__dict__.update(data)
|
||||
model.__dict__.update(kwargs)
|
||||
model.set_env(env)
|
||||
model.load_parameters(params)
|
||||
model.load_parameters(params, opt_params)
|
||||
|
||||
return model
|
||||
|
||||
|
|
@ -308,7 +311,7 @@ class BaseRLModel(object):
|
|||
:param load_path: (str or file-like) Where to load the model from
|
||||
:param load_data: (bool) Whether we should load and return data
|
||||
(class parameters). Mainly used by 'load_parameters' to only load model parameters (weights)
|
||||
:return: (dict. OrderedDict),(dict. OrderedDict) Class parameters and model parameters (state_dict)
|
||||
:return: (dict. OrderedDict),(dict. OrderedDict),(dict. OrderedDict) Class parameters, model parameters (state_dict) and dict of optimizer parameters (dict of state_dict)
|
||||
"""
|
||||
# Check if file exists if load_path is a string
|
||||
if isinstance(load_path, str):
|
||||
|
|
@ -327,6 +330,7 @@ class BaseRLModel(object):
|
|||
# as None (_save_to_file_zip allows this).
|
||||
data = None
|
||||
params = None
|
||||
opt_params = None
|
||||
if "data" in namelist and load_data:
|
||||
# Load class parameters and convert to string
|
||||
json_data = file_.read("data").decode()
|
||||
|
|
@ -341,11 +345,24 @@ class BaseRLModel(object):
|
|||
# go to start of file
|
||||
file_content.seek(0)
|
||||
params = th.load(file_content)
|
||||
# check for all other .pth files
|
||||
other_files = [file_name for file_name in namelist if
|
||||
os.path.splitext(file_name)[1] == ".pth" and file_name != "param.pth"]
|
||||
if len(other_files) > 0:
|
||||
opt_params = dict()
|
||||
for file in other_files:
|
||||
with file_.open(file, mode="r") as opt_param_file:
|
||||
# File has to be seekable so load in BytesIO first
|
||||
file_content = io.BytesIO()
|
||||
file_content.write(opt_param_file.read())
|
||||
# go to start of file
|
||||
file_content.seek(0)
|
||||
opt_params[os.path.splitext(file)[0]] = th.load(file_content)
|
||||
except zipfile.BadZipFile:
|
||||
# load_path wasn't a zip file
|
||||
raise ValueError("Error: the file {} wasn't a zip-file".format(load_path))
|
||||
|
||||
return data, params
|
||||
return data, params, opt_params
|
||||
|
||||
def set_random_seed(self, seed=0):
|
||||
set_random_seed(seed, using_cuda=self.device == th.device('cuda'))
|
||||
|
|
@ -456,36 +473,41 @@ class BaseRLModel(object):
|
|||
logger.dumpkvs()
|
||||
|
||||
mean_reward = np.mean(episode_rewards) if total_episodes > 0 else 0.0
|
||||
|
||||
return mean_reward, total_steps, total_episodes, obs
|
||||
|
||||
@staticmethod
|
||||
def _save_to_file_zip(save_path, data=None, params=None, opt_params=None):
|
||||
"""Save model to a zip archive
|
||||
|
||||
:param save_path: (str or file-like) Where to store the model
|
||||
:param data: (OrderedDict) Class parameters being stored
|
||||
:param params: (OrderedDict) Model parameters being stored expected to be state_dict
|
||||
:param opt_params: (OrderedDict) Optimizer parameters being stored expected to contain an entry for every
|
||||
optimizer with its name and the state_dict
|
||||
"""
|
||||
|
||||
def _save_to_file_zip(save_path, data=None, params=None):
|
||||
"""Save model to a zip archive
|
||||
|
||||
:param save_path: (str or file-like) Where to store the model
|
||||
:param data: (OrderedDict) Class parameters being stored
|
||||
:param params: (OrderedDict) Model parameters being stored expexted to be state_dict
|
||||
"""
|
||||
|
||||
# data/params can be None, so do not
|
||||
# try to serialize them blindly
|
||||
if data is not None:
|
||||
serialized_data = data_to_json(data)
|
||||
|
||||
# Check postfix if save_path is a string
|
||||
if isinstance(save_path, str):
|
||||
_, ext = os.path.splitext(save_path)
|
||||
if ext == "":
|
||||
save_path += ".zip"
|
||||
|
||||
# Create a zip-archive and write our objects
|
||||
# there. This works when save_path is either
|
||||
# str or a file-like
|
||||
with zipfile.ZipFile(save_path, "w") as file_:
|
||||
# Do not try to save "None" elements
|
||||
# data/params can be None, so do not
|
||||
# try to serialize them blindly
|
||||
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:
|
||||
th.save(params, param_file)
|
||||
serialized_data = data_to_json(data)
|
||||
|
||||
# Check postfix if save_path is a string
|
||||
if isinstance(save_path, str):
|
||||
_, ext = os.path.splitext(save_path)
|
||||
if ext == "":
|
||||
save_path += ".zip"
|
||||
|
||||
# Create a zip-archive and write our objects
|
||||
# there. This works when save_path is either
|
||||
# str or a file-like
|
||||
with zipfile.ZipFile(save_path, "w") as file_:
|
||||
# Do not try to save "None" elements
|
||||
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:
|
||||
th.save(params, param_file)
|
||||
if opt_params is not None:
|
||||
for file_name, dict in opt_params.items():
|
||||
with file_.open(file_name + '.pth', mode="w") as opt_param_file:
|
||||
th.save(dict, opt_param_file)
|
||||
|
|
|
|||
|
|
@ -292,6 +292,29 @@ class PPO(BaseRLModel):
|
|||
|
||||
return self
|
||||
|
||||
def get_opt_parameters(self):
|
||||
"""
|
||||
returns a dict of all the optimizers and their parameters
|
||||
|
||||
:return: (Dict) of optimizer names and their state_dict
|
||||
"""
|
||||
return {"opt": self.policy.optimizer.state_dict()}
|
||||
|
||||
def load_parameters(self, load_dict, opt_params):
|
||||
"""
|
||||
Load model parameters and optimizer parameters from a dictionary
|
||||
|
||||
Dictionary should be of shape torch model.state_dict()
|
||||
|
||||
This does not load agent's hyper-parameters.
|
||||
|
||||
|
||||
: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
|
||||
"""
|
||||
self.policy.optimizer.load_state_dict(opt_params["opt"])
|
||||
self.policy.load_state_dict(load_dict)
|
||||
|
||||
def save(self, path):
|
||||
"""
|
||||
saves all the params from init and pytorch params in a file for continous learning
|
||||
|
|
@ -320,12 +343,5 @@ class PPO(BaseRLModel):
|
|||
}
|
||||
|
||||
params_to_save = self.get_policy_parameters()
|
||||
|
||||
_save_to_file_zip(path, data=data, params=params_to_save)
|
||||
|
||||
"""def load(self, path, env=None, **_kwargs):
|
||||
if not path.endswith('.pth'):
|
||||
path += '.pth'
|
||||
if env is not None:
|
||||
pass
|
||||
self.policy.load_state_dict(th.load(path))"""
|
||||
opt_params_to_save = self.get_opt_parameters()
|
||||
self._save_to_file_zip(path, data=data, params=params_to_save,opt_params=opt_params_to_save)
|
||||
Loading…
Reference in a new issue