From fe67a98711d00d0001d763d6546ef7257de74aff Mon Sep 17 00:00:00 2001 From: Antonin Raffin Date: Tue, 26 Nov 2019 17:44:06 +0100 Subject: [PATCH] Log more values --- README.md | 4 +++- torchy_baselines/sac/policies.py | 2 +- torchy_baselines/sac/sac.py | 6 ++++++ 3 files changed, 10 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index 0b19832..bf14287 100644 --- a/README.md +++ b/README.md @@ -21,13 +21,15 @@ TODO: - save/load - better predict - complete logger -- SDE: learn the feature extractor? - Refactor: buffer with numpy array instead of pytorch - Refactor: remove duplicated code for evaluation + - plotting? -> zoo Later: - get_parameters / set_parameters +- SDE: use [affine transform](https://www.tensorflow.org/probability/api_docs/python/tfp/bijectors/Affine) + to scale the noise after a tanh transform? - CNN policies + normalization - tensorboard support - DQN diff --git a/torchy_baselines/sac/policies.py b/torchy_baselines/sac/policies.py index dd644bd..34bb3e4 100644 --- a/torchy_baselines/sac/policies.py +++ b/torchy_baselines/sac/policies.py @@ -23,7 +23,7 @@ class Actor(BaseNetwork): for the std instead of only (n_features,) when using SDE. """ def __init__(self, obs_dim, action_dim, net_arch, activation_fn=nn.ReLU, - use_sde=False, log_std_init=-3, full_std=False): + use_sde=False, log_std_init=-3, full_std=True): super(Actor, self).__init__() actor_net = create_mlp(obs_dim, -1, net_arch, activation_fn) diff --git a/torchy_baselines/sac/sac.py b/torchy_baselines/sac/sac.py index aae0443..49f562f 100644 --- a/torchy_baselines/sac/sac.py +++ b/torchy_baselines/sac/sac.py @@ -9,6 +9,7 @@ from torchy_baselines.common.buffers import ReplayBuffer from torchy_baselines.common.evaluation import evaluate_policy from torchy_baselines.sac.policies import SACPolicy from torchy_baselines.common.vec_env import sync_envs_normalization +from torchy_baselines.common import logger class SAC(BaseRLModel): @@ -234,6 +235,11 @@ class SAC(BaseRLModel): for param, target_param in zip(self.critic.parameters(), self.critic_target.parameters()): target_param.data.copy_(self.tau * param.data + (1 - self.tau) * target_param.data) + # TODO: average + logger.logkv("ent_coef", ent_coef.item()) + logger.logkv("actor_loss", actor_loss.item()) + logger.logkv("critic_loss", critic_loss.item()) + def learn(self, total_timesteps, callback=None, log_interval=4, eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="SAC", reset_num_timesteps=True):