mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-24 19:43:50 +00:00
Log more values
This commit is contained in:
parent
5483e02d1a
commit
fe67a98711
3 changed files with 10 additions and 2 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Reference in a new issue