diff --git a/docs/common/logger.rst b/docs/common/logger.rst index 3ba787c..5a61fe5 100644 --- a/docs/common/logger.rst +++ b/docs/common/logger.rst @@ -3,5 +3,30 @@ Logger ====== +To overwrite the default logger, you can pass one to the algorithm. +Available formats are ``["stdout", "csv", "log", "tensorboard", "json"]``. + + +.. warning:: + + When passing a custom logger object, + this will overwrite ``tensorboard_log`` and ``verbose`` settings + passed to the constructor. + + +.. code-block:: python + + from stable_baselines3 import A2C + from stable_baselines3.common.logger import configure + + tmp_path = "/tmp/sb3_log/" + # set up logger + new_logger = configure(tmp_path, ["stdout", "csv", "tensorboard"]) + + model = A2C("MlpPolicy", "CartPole-v1", verbose=1) + # Set new logger + model.set_logger(new_logger) + model.learn(10000) + .. automodule:: stable_baselines3.common.logger :members: diff --git a/docs/misc/changelog.rst b/docs/misc/changelog.rst index f84c09c..f511159 100644 --- a/docs/misc/changelog.rst +++ b/docs/misc/changelog.rst @@ -4,7 +4,7 @@ Changelog ========== -Release 1.1.0a10 (WIP) +Release 1.1.0a11 (WIP) --------------------------- **Dict observation support, timeout handling and refactored HER** @@ -28,6 +28,11 @@ Breaking Changes: - Updated the KL Divergence estimator in the PPO algorithm to be positive definite and have lower variance (@09tangriro) - Updated the KL Divergence check in the PPO algorithm to be before the gradient update step rather than after end of epoch (@09tangriro) - Removed parameter ``channels_last`` from ``is_image_space`` as it can be inferred. +- The logger object is now an attribute ``model.logger`` that be set by the user using ``model.set_logger()`` +- Changed the signature of ``logger.configure`` and ``utils.configure_logger``, they now return a ``Logger`` object +- Removed ``Logger.CURRENT`` and ``Logger.DEFAULT`` +- Moved ``warn(), debug(), log(), info(), dump()`` methods to the ``Logger`` class +- ``.learn()`` now throws an import error when the user tries to log to tensorboard but the package is not installed New Features: ^^^^^^^^^^^^^ @@ -53,6 +58,7 @@ Bug Fixes: - Fixed potential issue when calling off-policy algorithms with default arguments multiple times (the size of the replay buffer would be the same) - Fixed loading of ``ent_coef`` for ``SAC`` and ``TQC``, it was not optimized anymore (thanks @Atlis) - Fixed saving of ``A2C`` and ``PPO`` policy when using gSDE (thanks @liusida) +- Fixed a bug where no output would be shown even if ``verbose>=1`` after passing ``verbose=0`` once Deprecations: ^^^^^^^^^^^^^ @@ -81,6 +87,7 @@ Documentation: - Updated migration guide (@juancroldan) - Pinned ``docutils==0.16`` to avoid issue with rtd theme - Clarified callback ``save_freq`` definition +- Added doc on how to pass a custom logger - Remove recurrent policies from ``A2C`` docs (@bstee615) diff --git a/stable_baselines3/a2c/a2c.py b/stable_baselines3/a2c/a2c.py index b88b01f..03b1fc8 100644 --- a/stable_baselines3/a2c/a2c.py +++ b/stable_baselines3/a2c/a2c.py @@ -4,7 +4,6 @@ import torch as th from gym import spaces from torch.nn import functional as F -from stable_baselines3.common import logger from stable_baselines3.common.on_policy_algorithm import OnPolicyAlgorithm from stable_baselines3.common.policies import ActorCriticPolicy from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, Schedule @@ -166,13 +165,13 @@ class A2C(OnPolicyAlgorithm): explained_var = explained_variance(self.rollout_buffer.values.flatten(), self.rollout_buffer.returns.flatten()) self._n_updates += 1 - logger.record("train/n_updates", self._n_updates, exclude="tensorboard") - logger.record("train/explained_variance", explained_var) - logger.record("train/entropy_loss", entropy_loss.item()) - logger.record("train/policy_loss", policy_loss.item()) - logger.record("train/value_loss", value_loss.item()) + self.logger.record("train/n_updates", self._n_updates, exclude="tensorboard") + self.logger.record("train/explained_variance", explained_var) + self.logger.record("train/entropy_loss", entropy_loss.item()) + self.logger.record("train/policy_loss", policy_loss.item()) + self.logger.record("train/value_loss", value_loss.item()) if hasattr(self.policy, "log_std"): - logger.record("train/std", th.exp(self.policy.log_std).mean().item()) + self.logger.record("train/std", th.exp(self.policy.log_std).mean().item()) def learn( self, diff --git a/stable_baselines3/common/base_class.py b/stable_baselines3/common/base_class.py index a0af69d..8164504 100644 --- a/stable_baselines3/common/base_class.py +++ b/stable_baselines3/common/base_class.py @@ -11,9 +11,10 @@ import gym import numpy as np import torch as th -from stable_baselines3.common import logger, utils +from stable_baselines3.common import utils from stable_baselines3.common.callbacks import BaseCallback, CallbackList, ConvertCallback, EvalCallback from stable_baselines3.common.env_util import is_wrapped +from stable_baselines3.common.logger import Logger from stable_baselines3.common.monitor import Monitor from stable_baselines3.common.noise import ActionNoise from stable_baselines3.common.policies import BasePolicy, get_policy_from_name @@ -144,6 +145,10 @@ class BaseAlgorithm(ABC): self.ep_success_buffer = None # type: Optional[deque] # For logging (and TD3 delayed updates) self._n_updates = 0 # type: int + # The logger object + self._logger = None # type: Logger + # Whether the user passed a custom logger or not + self._custom_logger = False # Create and wrap the env if needed if env is not None: @@ -228,6 +233,25 @@ class BaseAlgorithm(ABC): def _setup_model(self) -> None: """Create networks, buffer and optimizers.""" + def set_logger(self, logger: Logger) -> None: + """ + Setter for for logger object. + + .. warning:: + + When passing a custom logger object, + this will overwrite ``tensorboard_log`` and ``verbose`` settings + passed to the constructor. + """ + self._logger = logger + # User defined logger + self._custom_logger = True + + @property + def logger(self) -> Logger: + """Getter for the logger object.""" + return self._logger + def _get_eval_env(self, eval_env: Optional[GymEnv]) -> Optional[GymEnv]: """ Return the environment that will be used for evaluation. @@ -265,7 +289,7 @@ class BaseAlgorithm(ABC): An optimizer or a list of optimizers. """ # Log the current learning rate - logger.record("train/learning_rate", self.lr_schedule(self._current_progress_remaining)) + self.logger.record("train/learning_rate", self.lr_schedule(self._current_progress_remaining)) if not isinstance(optimizers, list): optimizers = [optimizers] @@ -290,6 +314,8 @@ class BaseAlgorithm(ABC): "rollout_buffer", "_vec_normalize_env", "_episode_storage", + "_logger", + "_custom_logger", ] def _get_torch_save_params(self) -> Tuple[List[str], List[str]]: @@ -402,8 +428,9 @@ class BaseAlgorithm(ABC): eval_env = self._get_eval_env(eval_env) - # Configure logger's outputs - utils.configure_logger(self.verbose, self.tensorboard_log, tb_log_name, reset_num_timesteps) + # Configure logger's outputs if no logger was passed + if not self._custom_logger: + self._logger = utils.configure_logger(self.verbose, self.tensorboard_log, tb_log_name, reset_num_timesteps) # Create eval callback if needed callback = self._init_callback(callback, eval_env, eval_freq, n_eval_episodes, log_path) diff --git a/stable_baselines3/common/callbacks.py b/stable_baselines3/common/callbacks.py index 0d77feb..4191a16 100644 --- a/stable_baselines3/common/callbacks.py +++ b/stable_baselines3/common/callbacks.py @@ -6,7 +6,7 @@ from typing import Any, Callable, Dict, List, Optional, Union import gym import numpy as np -from stable_baselines3.common import base_class, logger # pytype: disable=pyi-error +from stable_baselines3.common import base_class # pytype: disable=pyi-error from stable_baselines3.common.evaluation import evaluate_policy from stable_baselines3.common.vec_env import DummyVecEnv, VecEnv, sync_envs_normalization @@ -44,7 +44,7 @@ class BaseCallback(ABC): """ self.model = model self.training_env = model.get_env() - self.logger = logger + self.logger = model.logger self._init_callback() def _init_callback(self) -> None: diff --git a/stable_baselines3/common/logger.py b/stable_baselines3/common/logger.py index 3d6b458..d1a32cd 100644 --- a/stable_baselines3/common/logger.py +++ b/stable_baselines3/common/logger.py @@ -398,168 +398,20 @@ def make_output_format(_format: str, log_dir: str, log_suffix: str = "") -> KVWr raise ValueError(f"Unknown format specified: {_format}") -# ================================================================ -# API -# ================================================================ - - -def record(key: str, value: Any, exclude: Optional[Union[str, Tuple[str, ...]]] = None) -> None: - """ - Log a value of some diagnostic - Call this once for each diagnostic quantity, each iteration - If called many times, last value will be used. - - :param key: save to log this key - :param value: save to log this value - :param exclude: outputs to be excluded - """ - Logger.CURRENT.record(key, value, exclude) - - -def record_mean(key: str, value: Union[int, float], exclude: Optional[Union[str, Tuple[str, ...]]] = None) -> None: - """ - The same as record(), but if called many times, values averaged. - - :param key: save to log this key - :param value: save to log this value - :param exclude: outputs to be excluded - """ - Logger.CURRENT.record_mean(key, value, exclude) - - -def record_dict(key_values: Dict[str, Any]) -> None: - """ - Log a dictionary of key-value pairs. - - :param key_values: the list of keys and values to save to log - """ - for key, value in key_values.items(): - record(key, value) - - -def dump(step: int = 0) -> None: - """ - Write all of the diagnostics from the current iteration - """ - Logger.CURRENT.dump(step) - - -def get_log_dict() -> Dict: - """ - get the key values logs - - :return: the logged values - """ - return Logger.CURRENT.name_to_value - - -def log(*args, level: int = INFO) -> None: - """ - Write the sequence of args, with no separators, - to the console and output files (if you've configured an output file). - - level: int. (see logger.py docs) If the global logger level is higher than - the level argument here, don't print to stdout. - - :param args: log the arguments - :param level: the logging level (can be DEBUG=10, INFO=20, WARN=30, ERROR=40, DISABLED=50) - """ - Logger.CURRENT.log(*args, level=level) - - -def debug(*args) -> None: - """ - Write the sequence of args, with no separators, - to the console and output files (if you've configured an output file). - Using the DEBUG level. - - :param args: log the arguments - """ - log(*args, level=DEBUG) - - -def info(*args) -> None: - """ - Write the sequence of args, with no separators, - to the console and output files (if you've configured an output file). - Using the INFO level. - - :param args: log the arguments - """ - log(*args, level=INFO) - - -def warn(*args) -> None: - """ - Write the sequence of args, with no separators, - to the console and output files (if you've configured an output file). - Using the WARN level. - - :param args: log the arguments - """ - log(*args, level=WARN) - - -def error(*args) -> None: - """ - Write the sequence of args, with no separators, - to the console and output files (if you've configured an output file). - Using the ERROR level. - - :param args: log the arguments - """ - log(*args, level=ERROR) - - -def set_level(level: int) -> None: - """ - Set logging threshold on current logger. - - :param level: the logging level (can be DEBUG=10, INFO=20, WARN=30, ERROR=40, DISABLED=50) - """ - Logger.CURRENT.set_level(level) - - -def get_level() -> int: - """ - Get logging threshold on current logger. - :return: the logging level (can be DEBUG=10, INFO=20, WARN=30, ERROR=40, DISABLED=50) - """ - return Logger.CURRENT.level - - -def get_dir() -> str: - """ - Get directory that log files are being written to. - will be None if there is no output directory (i.e., if you didn't call start) - - :return: the logging directory - """ - return Logger.CURRENT.get_dir() - - -record_tabular = record -dump_tabular = dump - - # ================================================================ # Backend # ================================================================ class Logger(object): - # A logger with no output files. (See right below class definition) - # So that you can still log to the terminal without setting up any output files - DEFAULT = None - CURRENT = None # Current logger being used by the free functions above + """ + The logger class. + + :param folder: the logging location + :param output_formats: the list of output formats + """ def __init__(self, folder: Optional[str], output_formats: List[KVWriter]): - """ - the logger class - - :param folder: the logging location - :param output_formats: the list of output format - """ self.name_to_value = defaultdict(float) # values this iteration self.name_to_count = defaultdict(int) self.name_to_excluded = defaultdict(str) @@ -567,8 +419,6 @@ class Logger(object): self.dir = folder self.output_formats = output_formats - # Logging API, forwarded - # ---------------------------------------- def record(self, key: str, value: Any, exclude: Optional[Union[str, Tuple[str, ...]]] = None) -> None: """ Log a value of some diagnostic @@ -626,6 +476,46 @@ class Logger(object): if self.level <= level: self._do_log(args) + def debug(self, *args) -> None: + """ + Write the sequence of args, with no separators, + to the console and output files (if you've configured an output file). + Using the DEBUG level. + + :param args: log the arguments + """ + self.log(*args, level=DEBUG) + + def info(self, *args) -> None: + """ + Write the sequence of args, with no separators, + to the console and output files (if you've configured an output file). + Using the INFO level. + + :param args: log the arguments + """ + self.log(*args, level=INFO) + + def warn(self, *args) -> None: + """ + Write the sequence of args, with no separators, + to the console and output files (if you've configured an output file). + Using the WARN level. + + :param args: log the arguments + """ + self.log(*args, level=WARN) + + def error(self, *args) -> None: + """ + Write the sequence of args, with no separators, + to the console and output files (if you've configured an output file). + Using the ERROR level. + + :param args: log the arguments + """ + self.log(*args, level=ERROR) + # Configuration # ---------------------------------------- def set_level(self, level: int) -> None: @@ -665,18 +555,15 @@ class Logger(object): _format.write_sequence(map(str, args)) -# Initialize logger -Logger.DEFAULT = Logger.CURRENT = Logger(folder=None, output_formats=[HumanOutputFormat(sys.stdout)]) - - -def configure(folder: Optional[str] = None, format_strings: Optional[List[str]] = None) -> None: +def configure(folder: Optional[str] = None, format_strings: Optional[List[str]] = None) -> Logger: """ - configure the current logger + Configure the current logger. :param folder: the save location - (if None, $SB3_LOGDIR, if still None, tempdir/baselines-[date & time]) + (if None, $SB3_LOGDIR, if still None, tempdir/SB3-[date & time]) :param format_strings: the output logging format (if None, $SB3_LOG_FORMAT, if still None, ['stdout', 'log', 'csv']) + :return: The logger object. """ if folder is None: folder = os.getenv("SB3_LOGDIR") @@ -689,46 +576,14 @@ def configure(folder: Optional[str] = None, format_strings: Optional[List[str]] if format_strings is None: format_strings = os.getenv("SB3_LOG_FORMAT", "stdout,log,csv").split(",") - format_strings = filter(None, format_strings) + format_strings = list(filter(None, format_strings)) output_formats = [make_output_format(f, folder, log_suffix) for f in format_strings] - Logger.CURRENT = Logger(folder=folder, output_formats=output_formats) - log(f"Logging to {folder}") - - -def reset() -> None: - """ - reset the current logger - """ - if Logger.CURRENT is not Logger.DEFAULT: - Logger.CURRENT.close() - Logger.CURRENT = Logger.DEFAULT - log("Reset logger") - - -class ScopedConfigure(object): - def __init__(self, folder: Optional[str] = None, format_strings: Optional[List[str]] = None): - """ - Class for using context manager while logging - - usage: - with ScopedConfigure(folder=None, format_strings=None): - {code} - - :param folder: the logging folder - :param format_strings: the list of output logging format - """ - self.dir = folder - self.format_strings = format_strings - self.prev_logger = None - - def __enter__(self) -> None: - self.prev_logger = Logger.CURRENT - configure(folder=self.dir, format_strings=self.format_strings) - - def __exit__(self, *args) -> None: - Logger.CURRENT.close() - Logger.CURRENT = self.prev_logger + logger = Logger(folder=folder, output_formats=output_formats) + # Only print when some files will be saved + if len(format_strings) > 0 and format_strings != ["stdout"]: + logger.log(f"Logging to {folder}") + return logger # ================================================================ diff --git a/stable_baselines3/common/off_policy_algorithm.py b/stable_baselines3/common/off_policy_algorithm.py index 46e6d56..2f99f0b 100644 --- a/stable_baselines3/common/off_policy_algorithm.py +++ b/stable_baselines3/common/off_policy_algorithm.py @@ -8,7 +8,6 @@ import gym import numpy as np import torch as th -from stable_baselines3.common import logger from stable_baselines3.common.base_class import BaseAlgorithm from stable_baselines3.common.buffers import DictReplayBuffer, ReplayBuffer from stable_baselines3.common.callbacks import BaseCallback @@ -430,20 +429,20 @@ class OffPolicyAlgorithm(BaseAlgorithm): """ time_elapsed = time.time() - self.start_time fps = int(self.num_timesteps / (time_elapsed + 1e-8)) - logger.record("time/episodes", self._episode_num, exclude="tensorboard") + self.logger.record("time/episodes", self._episode_num, exclude="tensorboard") if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0: - logger.record("rollout/ep_rew_mean", safe_mean([ep_info["r"] for ep_info in self.ep_info_buffer])) - logger.record("rollout/ep_len_mean", safe_mean([ep_info["l"] for ep_info in self.ep_info_buffer])) - logger.record("time/fps", fps) - logger.record("time/time_elapsed", int(time_elapsed), exclude="tensorboard") - logger.record("time/total timesteps", self.num_timesteps, exclude="tensorboard") + self.logger.record("rollout/ep_rew_mean", safe_mean([ep_info["r"] for ep_info in self.ep_info_buffer])) + self.logger.record("rollout/ep_len_mean", safe_mean([ep_info["l"] for ep_info in self.ep_info_buffer])) + self.logger.record("time/fps", fps) + self.logger.record("time/time_elapsed", int(time_elapsed), exclude="tensorboard") + self.logger.record("time/total timesteps", self.num_timesteps, exclude="tensorboard") if self.use_sde: - logger.record("train/std", (self.actor.get_std()).mean().item()) + self.logger.record("train/std", (self.actor.get_std()).mean().item()) if len(self.ep_success_buffer) > 0: - logger.record("rollout/success rate", safe_mean(self.ep_success_buffer)) + self.logger.record("rollout/success rate", safe_mean(self.ep_success_buffer)) # Pass the number of timesteps for tensorboard - logger.dump(step=self.num_timesteps) + self.logger.dump(step=self.num_timesteps) def _on_step(self) -> None: """ diff --git a/stable_baselines3/common/on_policy_algorithm.py b/stable_baselines3/common/on_policy_algorithm.py index 924788d..5d872a9 100644 --- a/stable_baselines3/common/on_policy_algorithm.py +++ b/stable_baselines3/common/on_policy_algorithm.py @@ -5,7 +5,6 @@ import gym import numpy as np import torch as th -from stable_baselines3.common import logger from stable_baselines3.common.base_class import BaseAlgorithm from stable_baselines3.common.buffers import DictRolloutBuffer, RolloutBuffer from stable_baselines3.common.callbacks import BaseCallback @@ -243,14 +242,14 @@ class OnPolicyAlgorithm(BaseAlgorithm): # Display training infos if log_interval is not None and iteration % log_interval == 0: fps = int(self.num_timesteps / (time.time() - self.start_time)) - logger.record("time/iterations", iteration, exclude="tensorboard") + self.logger.record("time/iterations", iteration, exclude="tensorboard") if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0: - logger.record("rollout/ep_rew_mean", safe_mean([ep_info["r"] for ep_info in self.ep_info_buffer])) - logger.record("rollout/ep_len_mean", safe_mean([ep_info["l"] for ep_info in self.ep_info_buffer])) - logger.record("time/fps", fps) - logger.record("time/time_elapsed", int(time.time() - self.start_time), exclude="tensorboard") - logger.record("time/total_timesteps", self.num_timesteps, exclude="tensorboard") - logger.dump(step=self.num_timesteps) + self.logger.record("rollout/ep_rew_mean", safe_mean([ep_info["r"] for ep_info in self.ep_info_buffer])) + self.logger.record("rollout/ep_len_mean", safe_mean([ep_info["l"] for ep_info in self.ep_info_buffer])) + self.logger.record("time/fps", fps) + self.logger.record("time/time_elapsed", int(time.time() - self.start_time), exclude="tensorboard") + self.logger.record("time/total_timesteps", self.num_timesteps, exclude="tensorboard") + self.logger.dump(step=self.num_timesteps) self.train() diff --git a/stable_baselines3/common/utils.py b/stable_baselines3/common/utils.py index 9b1a6f3..8711591 100644 --- a/stable_baselines3/common/utils.py +++ b/stable_baselines3/common/utils.py @@ -15,7 +15,7 @@ try: except ImportError: SummaryWriter = None -from stable_baselines3.common import logger +from stable_baselines3.common.logger import Logger, configure from stable_baselines3.common.type_aliases import GymEnv, Schedule, TensorDict, TrainFreq, TrainFrequencyUnit @@ -172,14 +172,23 @@ def configure_logger( tensorboard_log: Optional[str] = None, tb_log_name: str = "", reset_num_timesteps: bool = True, -) -> None: +) -> Logger: """ Configure the logger's outputs. :param verbose: the verbosity level: 0 no output, 1 info, 2 debug :param tensorboard_log: the log location for tensorboard (if None, no logging) :param tb_log_name: tensorboard log + :param reset_num_timesteps: Whether the ``num_timesteps`` attribute is reset or not. + It allows to continue a previous learning curve (``reset_num_timesteps=False``) + or start from t=0 (``reset_num_timesteps=True``, the default). + :return: The logger object """ + save_path, format_strings = None, ["stdout"] + + if tensorboard_log is not None and SummaryWriter is None: + raise ImportError("Trying to log data to tensorboard but tensorboard is not installed.") + if tensorboard_log is not None and SummaryWriter is not None: latest_run_id = get_latest_run_id(tensorboard_log, tb_log_name) if not reset_num_timesteps: @@ -187,11 +196,12 @@ def configure_logger( latest_run_id -= 1 save_path = os.path.join(tensorboard_log, f"{tb_log_name}_{latest_run_id + 1}") if verbose >= 1: - logger.configure(save_path, ["stdout", "tensorboard"]) + format_strings = ["stdout", "tensorboard"] else: - logger.configure(save_path, ["tensorboard"]) + format_strings = ["tensorboard"] elif verbose == 0: - logger.configure(format_strings=[""]) + format_strings = [""] + return configure(save_path, format_strings=format_strings) def check_for_correct_spaces(env: GymEnv, observation_space: gym.spaces.Space, action_space: gym.spaces.Space) -> None: diff --git a/stable_baselines3/common/vec_env/base_vec_env.py b/stable_baselines3/common/vec_env/base_vec_env.py index c7bd7ac..d3e624a 100644 --- a/stable_baselines3/common/vec_env/base_vec_env.py +++ b/stable_baselines3/common/vec_env/base_vec_env.py @@ -1,4 +1,5 @@ import inspect +import warnings from abc import ABC, abstractmethod from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple, Type, Union @@ -6,8 +7,6 @@ import cloudpickle import gym import numpy as np -from stable_baselines3.common import logger - # Define type aliases here to avoid circular import # Used when we want to access one or more VecEnv VecEnvIndices = Union[None, int, Iterable[int]] @@ -177,7 +176,7 @@ class VecEnv(ABC): try: imgs = self.get_images() except NotImplementedError: - logger.warn(f"Render not defined for {self}") + warnings.warn(f"Render not defined for {self}") return # Create a big image by tiling images from subprocesses diff --git a/stable_baselines3/common/vec_env/vec_video_recorder.py b/stable_baselines3/common/vec_env/vec_video_recorder.py index 08a2950..70d74eb 100644 --- a/stable_baselines3/common/vec_env/vec_video_recorder.py +++ b/stable_baselines3/common/vec_env/vec_video_recorder.py @@ -3,7 +3,6 @@ from typing import Callable from gym.wrappers.monitoring import video_recorder -from stable_baselines3.common import logger from stable_baselines3.common.vec_env.base_vec_env import VecEnv, VecEnvObs, VecEnvStepReturn, VecEnvWrapper from stable_baselines3.common.vec_env.dummy_vec_env import DummyVecEnv from stable_baselines3.common.vec_env.subproc_vec_env import SubprocVecEnv @@ -93,7 +92,7 @@ class VecVideoRecorder(VecEnvWrapper): self.video_recorder.capture_frame() self.recorded_frames += 1 if self.recorded_frames > self.video_length: - logger.info("Saving video to ", self.video_recorder.path) + print(f"Saving video to {self.video_recorder.path}") self.close_video_recorder() elif self._video_enabled(): self.start_video_recorder() diff --git a/stable_baselines3/dqn/dqn.py b/stable_baselines3/dqn/dqn.py index 6dffaa7..d68a643 100644 --- a/stable_baselines3/dqn/dqn.py +++ b/stable_baselines3/dqn/dqn.py @@ -5,7 +5,6 @@ import numpy as np import torch as th from torch.nn import functional as F -from stable_baselines3.common import logger from stable_baselines3.common.buffers import ReplayBuffer from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm from stable_baselines3.common.preprocessing import maybe_transpose @@ -150,7 +149,7 @@ class DQN(OffPolicyAlgorithm): polyak_update(self.q_net.parameters(), self.q_net_target.parameters(), self.tau) self.exploration_rate = self.exploration_schedule(self._current_progress_remaining) - logger.record("rollout/exploration rate", self.exploration_rate) + self.logger.record("rollout/exploration rate", self.exploration_rate) def train(self, gradient_steps: int, batch_size: int = 100) -> None: # Update learning rate according to schedule @@ -191,8 +190,8 @@ class DQN(OffPolicyAlgorithm): # Increase update counter self._n_updates += gradient_steps - logger.record("train/n_updates", self._n_updates, exclude="tensorboard") - logger.record("train/loss", np.mean(losses)) + self.logger.record("train/n_updates", self._n_updates, exclude="tensorboard") + self.logger.record("train/loss", np.mean(losses)) def predict( self, diff --git a/stable_baselines3/ppo/ppo.py b/stable_baselines3/ppo/ppo.py index 660eb32..28f8777 100644 --- a/stable_baselines3/ppo/ppo.py +++ b/stable_baselines3/ppo/ppo.py @@ -6,7 +6,6 @@ import torch as th from gym import spaces from torch.nn import functional as F -from stable_baselines3.common import logger from stable_baselines3.common.on_policy_algorithm import OnPolicyAlgorithm from stable_baselines3.common.policies import ActorCriticPolicy from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, Schedule @@ -269,20 +268,20 @@ class PPO(OnPolicyAlgorithm): explained_var = explained_variance(self.rollout_buffer.values.flatten(), self.rollout_buffer.returns.flatten()) # Logs - logger.record("train/entropy_loss", np.mean(entropy_losses)) - logger.record("train/policy_gradient_loss", np.mean(pg_losses)) - logger.record("train/value_loss", np.mean(value_losses)) - logger.record("train/approx_kl", np.mean(approx_kl_divs)) - logger.record("train/clip_fraction", np.mean(clip_fractions)) - logger.record("train/loss", loss.item()) - logger.record("train/explained_variance", explained_var) + self.logger.record("train/entropy_loss", np.mean(entropy_losses)) + self.logger.record("train/policy_gradient_loss", np.mean(pg_losses)) + self.logger.record("train/value_loss", np.mean(value_losses)) + self.logger.record("train/approx_kl", np.mean(approx_kl_divs)) + self.logger.record("train/clip_fraction", np.mean(clip_fractions)) + self.logger.record("train/loss", loss.item()) + self.logger.record("train/explained_variance", explained_var) if hasattr(self.policy, "log_std"): - logger.record("train/std", th.exp(self.policy.log_std).mean().item()) + self.logger.record("train/std", th.exp(self.policy.log_std).mean().item()) - logger.record("train/n_updates", self._n_updates, exclude="tensorboard") - logger.record("train/clip_range", clip_range) + self.logger.record("train/n_updates", self._n_updates, exclude="tensorboard") + self.logger.record("train/clip_range", clip_range) if self.clip_range_vf is not None: - logger.record("train/clip_range_vf", clip_range_vf) + self.logger.record("train/clip_range_vf", clip_range_vf) def learn( self, diff --git a/stable_baselines3/sac/sac.py b/stable_baselines3/sac/sac.py index bcbd165..dd3c501 100644 --- a/stable_baselines3/sac/sac.py +++ b/stable_baselines3/sac/sac.py @@ -5,7 +5,6 @@ import numpy as np import torch as th from torch.nn import functional as F -from stable_baselines3.common import logger from stable_baselines3.common.buffers import ReplayBuffer from stable_baselines3.common.noise import ActionNoise from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm @@ -267,12 +266,12 @@ class SAC(OffPolicyAlgorithm): self._n_updates += gradient_steps - logger.record("train/n_updates", self._n_updates, exclude="tensorboard") - logger.record("train/ent_coef", np.mean(ent_coefs)) - logger.record("train/actor_loss", np.mean(actor_losses)) - logger.record("train/critic_loss", np.mean(critic_losses)) + self.logger.record("train/n_updates", self._n_updates, exclude="tensorboard") + self.logger.record("train/ent_coef", np.mean(ent_coefs)) + self.logger.record("train/actor_loss", np.mean(actor_losses)) + self.logger.record("train/critic_loss", np.mean(critic_losses)) if len(ent_coef_losses) > 0: - logger.record("train/ent_coef_loss", np.mean(ent_coef_losses)) + self.logger.record("train/ent_coef_loss", np.mean(ent_coef_losses)) def learn( self, diff --git a/stable_baselines3/td3/td3.py b/stable_baselines3/td3/td3.py index 2b165c0..9227910 100644 --- a/stable_baselines3/td3/td3.py +++ b/stable_baselines3/td3/td3.py @@ -5,7 +5,6 @@ import numpy as np import torch as th from torch.nn import functional as F -from stable_baselines3.common import logger from stable_baselines3.common.buffers import ReplayBuffer from stable_baselines3.common.noise import ActionNoise from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm @@ -182,10 +181,10 @@ class TD3(OffPolicyAlgorithm): polyak_update(self.critic.parameters(), self.critic_target.parameters(), self.tau) polyak_update(self.actor.parameters(), self.actor_target.parameters(), self.tau) - logger.record("train/n_updates", self._n_updates, exclude="tensorboard") + self.logger.record("train/n_updates", self._n_updates, exclude="tensorboard") if len(actor_losses) > 0: - logger.record("train/actor_loss", np.mean(actor_losses)) - logger.record("train/critic_loss", np.mean(critic_losses)) + self.logger.record("train/actor_loss", np.mean(actor_losses)) + self.logger.record("train/critic_loss", np.mean(critic_losses)) def learn( self, diff --git a/stable_baselines3/version.txt b/stable_baselines3/version.txt index 481e045..a149840 100644 --- a/stable_baselines3/version.txt +++ b/stable_baselines3/version.txt @@ -1 +1 @@ -1.1.0a10 +1.1.0a11 diff --git a/tests/test_logger.py b/tests/test_logger.py index 74a52e5..e516171 100644 --- a/tests/test_logger.py +++ b/tests/test_logger.py @@ -1,3 +1,4 @@ +import os from typing import Sequence import numpy as np @@ -6,31 +7,21 @@ import torch as th from matplotlib import pyplot as plt from pandas.errors import EmptyDataError +from stable_baselines3 import A2C from stable_baselines3.common.logger import ( DEBUG, INFO, + CSVOutputFormat, Figure, FormatUnsupportedError, + HumanOutputFormat, Image, - ScopedConfigure, + TensorBoardOutputFormat, Video, configure, - debug, - dump, - error, - get_dir, - get_level, - get_log_dict, - info, make_output_format, read_csv, read_json, - record, - record_dict, - record_mean, - reset, - set_level, - warn, ) KEY_VALUES = { @@ -103,46 +94,83 @@ def read_log(tmp_path, capsys): return read_fn +def test_set_logger(tmp_path): + # set up logger + new_logger = configure(str(tmp_path), ["stdout", "csv", "tensorboard"]) + # Default outputs with verbose=0 + model = A2C("MlpPolicy", "CartPole-v1", verbose=0).learn(4) + assert model.logger.output_formats == [] + + model = A2C("MlpPolicy", "CartPole-v1", verbose=0, tensorboard_log=str(tmp_path)).learn(4) + assert str(tmp_path) in model.logger.dir + assert isinstance(model.logger.output_formats[0], TensorBoardOutputFormat) + + # Check that env variable work + new_tmp_path = str(tmp_path / "new_tmp") + os.environ["SB3_LOGDIR"] = new_tmp_path + model = A2C("MlpPolicy", "CartPole-v1", verbose=0).learn(4) + assert model.logger.dir == new_tmp_path + + # Default outputs with verbose=1 + model = A2C("MlpPolicy", "CartPole-v1", verbose=1).learn(4) + assert isinstance(model.logger.output_formats[0], HumanOutputFormat) + # with tensorboard + model = A2C("MlpPolicy", "CartPole-v1", verbose=1, tensorboard_log=str(tmp_path)).learn(4) + assert isinstance(model.logger.output_formats[0], HumanOutputFormat) + assert isinstance(model.logger.output_formats[1], TensorBoardOutputFormat) + assert len(model.logger.output_formats) == 2 + model.learn(32) + # set new logger + model.set_logger(new_logger) + # Check that the new logger is correctly setup + assert isinstance(model.logger.output_formats[0], HumanOutputFormat) + assert isinstance(model.logger.output_formats[1], CSVOutputFormat) + assert isinstance(model.logger.output_formats[2], TensorBoardOutputFormat) + assert len(model.logger.output_formats) == 3 + model.learn(32) + + model = A2C("MlpPolicy", "CartPole-v1", verbose=1) + model.set_logger(new_logger) + model.learn(32) + # Check that the new logger is not overwritten + assert isinstance(model.logger.output_formats[0], HumanOutputFormat) + assert isinstance(model.logger.output_formats[1], CSVOutputFormat) + assert isinstance(model.logger.output_formats[2], TensorBoardOutputFormat) + assert len(model.logger.output_formats) == 3 + + def test_main(tmp_path): """ tests for the logger module """ - info("hi") - debug("shouldn't appear") - assert get_level() == INFO - set_level(DEBUG) - assert get_level() == DEBUG - debug("should appear") - configure(folder=str(tmp_path)) - assert get_dir() == str(tmp_path) - record("a", 3) - record("b", 2.5) - dump() - record("b", -2.5) - record("a", 5.5) - dump() - info("^^^ should see a = 5.5") - record("f", "this text \n \r should appear in one line") - dump() - info('^^^ should see f = "this text \n \r should appear in one line"') - record_mean("b", -22.5) - record_mean("b", -44.4) - record("a", 5.5) - dump() - with ScopedConfigure(None, None): - info("^^^ should see b = 33.3") + logger = configure(None, ["stdout"]) + logger.info("hi") + logger.debug("shouldn't appear") + assert logger.level == INFO + logger.set_level(DEBUG) + assert logger.level == DEBUG + logger.debug("should appear") + logger = configure(folder=str(tmp_path)) + assert logger.dir == str(tmp_path) + logger.record("a", 3) + logger.record("b", 2.5) + logger.dump() + logger.record("b", -2.5) + logger.record("a", 5.5) + logger.dump() + logger.info("^^^ should see a = 5.5") + logger.record("f", "this text \n \r should appear in one line") + logger.dump() + logger.info('^^^ should see f = "this text \n \r should appear in one line"') + logger.record_mean("b", -22.5) + logger.record_mean("b", -44.4) + logger.record("a", 5.5) + logger.dump() - with ScopedConfigure(str(tmp_path / "test-logger"), ["json"]): - record("b", -2.5) - dump() - - reset() - record("a", "longasslongasslongasslongasslongasslongassvalue") - dump() - warn("hey") - error("oh") - record_dict({"test": 1}) - assert isinstance(get_log_dict(), dict) and set(get_log_dict().keys()) == {"test"} + logger.record("a", "longasslongasslongasslongasslongasslongassvalue") + logger.dump() + logger.warn("hey") + logger.error("oh") @pytest.mark.parametrize("_format", ["stdout", "log", "json", "csv", "tensorboard"])