From 4c0f6cbe53c31070d316f5c45d7c568a2e51d9d6 Mon Sep 17 00:00:00 2001 From: Noah Dormann Date: Thu, 5 Dec 2019 08:43:12 +0100 Subject: [PATCH] update get_opt_parameters to remove duplicate code --- torchy_baselines/sac/sac.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/torchy_baselines/sac/sac.py b/torchy_baselines/sac/sac.py index bfd873b..5222754 100644 --- a/torchy_baselines/sac/sac.py +++ b/torchy_baselines/sac/sac.py @@ -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): """