Log more values

This commit is contained in:
Antonin Raffin 2019-11-26 17:44:06 +01:00
parent 5483e02d1a
commit fe67a98711
3 changed files with 10 additions and 2 deletions

View file

@ -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

View file

@ -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)

View file

@ -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):