From f4546837c37ad90a95b6f5ba0ef1a363ce533e48 Mon Sep 17 00:00:00 2001 From: Antonin Raffin Date: Thu, 7 Nov 2019 17:41:28 +0100 Subject: [PATCH] Add std to logger --- torchy_baselines/common/base_class.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/torchy_baselines/common/base_class.py b/torchy_baselines/common/base_class.py index d96950d..2f8ef21 100644 --- a/torchy_baselines/common/base_class.py +++ b/torchy_baselines/common/base_class.py @@ -328,7 +328,7 @@ class BaseRLModel(object): assert env.num_envs == 1 if hasattr(self, 'use_sde') and self.use_sde: - self.policy.reset_noise() + self.actor.reset_noise() while total_steps < n_steps or total_episodes < n_episodes: done = False @@ -394,6 +394,8 @@ class BaseRLModel(object): logger.logkv("fps", fps) logger.logkv('time_elapsed', int(time.time() - self.start_time)) logger.logkv("total timesteps", num_timesteps) + if hasattr(self, 'use_sde') and self.use_sde: + logger.logkv("std", th.exp(self.actor.log_std).mean().item()) logger.dumpkvs() mean_reward = np.mean(episode_rewards) if total_episodes > 0 else 0.0