update get_opt_parameters to remove duplicate code

This commit is contained in:
Noah Dormann 2019-12-05 08:43:12 +01:00
parent c3b0398d56
commit 4c0f6cbe53

View file

@ -281,11 +281,10 @@ class SAC(BaseRLModel):
:return: (Dict) of optimizer names and their state_dict :return: (Dict) of optimizer names and their state_dict
""" """
opt_dict = {"actor": self.actor.optimizer.state_dict(), "critic": self.critic.optimizer.state_dict()}
if self.ent_coef_optimizer is not None: if self.ent_coef_optimizer is not None:
return {"actor": self.actor.optimizer.state_dict(), "critic": self.critic.optimizer.state_dict(), opt_dict.update({"ent_coef_optimizer": self.ent_coef_optimizer.state_dict()})
"ent_coef_optimizer": self.ent_coef_optimizer.state_dict()} return opt_dict
else:
return {"actor": self.actor.optimizer.state_dict(), "critic": self.critic.optimizer.state_dict()}
def load_parameters(self, load_dict, opt_params): def load_parameters(self, load_dict, opt_params):
""" """