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
"""
opt_dict = {"actor": self.actor.optimizer.state_dict(), "critic": self.critic.optimizer.state_dict()}
if self.ent_coef_optimizer is not None:
return {"actor": self.actor.optimizer.state_dict(), "critic": self.critic.optimizer.state_dict(),
"ent_coef_optimizer": self.ent_coef_optimizer.state_dict()}
else:
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()})
return opt_dict
def load_parameters(self, load_dict, opt_params):
"""