Merge branch 'master' into feat/dropq

This commit is contained in:
Antonin Raffin 2022-08-28 16:51:54 +02:00
commit 8311377634
No known key found for this signature in database
GPG key ID: B8B48F65CAD6232C
18 changed files with 354 additions and 64 deletions

View file

@ -157,6 +157,10 @@ CheckpointCallback
Callback for saving a model every ``save_freq`` calls to ``env.step()``, you must specify a log folder (``save_path``) Callback for saving a model every ``save_freq`` calls to ``env.step()``, you must specify a log folder (``save_path``)
and optionally a prefix for the checkpoints (``rl_model`` by default). and optionally a prefix for the checkpoints (``rl_model`` by default).
If you are using this callback to stop and resume training, you may want to optionally save the replay buffer if the
model has one (``save_replay_buffer``, ``False`` by default).
Additionally, if your environment uses a :ref:`VecNormalize <vec_env>` wrapper, you can save the
corresponding statistics using ``save_vecnormalize`` (``False`` by default).
.. warning:: .. warning::
@ -168,14 +172,20 @@ and optionally a prefix for the checkpoints (``rl_model`` by default).
.. code-block:: python .. code-block:: python
from stable_baselines3 import SAC from stable_baselines3 import SAC
from stable_baselines3.common.callbacks import CheckpointCallback from stable_baselines3.common.callbacks import CheckpointCallback
# Save a checkpoint every 1000 steps
checkpoint_callback = CheckpointCallback(save_freq=1000, save_path='./logs/',
name_prefix='rl_model')
model = SAC('MlpPolicy', 'Pendulum-v1') # Save a checkpoint every 1000 steps
model.learn(2000, callback=checkpoint_callback) checkpoint_callback = CheckpointCallback(
save_freq=1000,
save_path="./logs/",
name_prefix="rl_model",
save_replay_buffer=True,
save_vecnormalize=True,
)
model = SAC("MlpPolicy", "Pendulum-v1")
model.learn(2000, callback=checkpoint_callback)
.. _EvalCallback: .. _EvalCallback:

View file

@ -249,6 +249,55 @@ Here is an example of how to render an episode and log the resulting video to Te
video_recorder = VideoRecorderCallback(gym.make("CartPole-v1"), render_freq=5000) video_recorder = VideoRecorderCallback(gym.make("CartPole-v1"), render_freq=5000)
model.learn(total_timesteps=int(5e4), callback=video_recorder) model.learn(total_timesteps=int(5e4), callback=video_recorder)
Logging Hyperparameters
-----------------------
TensorBoard supports logging of hyperparameters in its HPARAMS tab, which helps comparing agents trainings.
.. warning::
To display hyperparameters in the HPARAMS section, a ``metric_dict`` must be given (as well as a ``hparam_dict``).
Here is an example of how to save hyperparameters in TensorBoard:
.. code-block:: python
from stable_baselines3 import A2C
from stable_baselines3.common.callbacks import BaseCallback
from stable_baselines3.common.logger import HParam
class HParamCallback(BaseCallback):
def __init__(self):
"""
Saves the hyperparameters and metrics at the start of the training, and logs them to TensorBoard.
"""
super().__init__()
def _on_training_start(self) -> None:
hparam_dict = {
"algorithm": self.model.__class__.__name__,
"learning rate": self.model.learning_rate,
"gamma": self.model.gamma,
}
# define the metrics that will appear in the `HPARAMS` Tensorboard tab by referencing their tag
# Tensorbaord will find & display metrics from the `SCALARS` tab
metric_dict = {
"rollout/ep_len_mean": 0,
"train/value_loss": 0,
}
self.logger.record(
"hparams",
HParam(hparam_dict, metric_dict),
exclude=("stdout", "log", "json", "csv"),
)
def _on_step(self) -> bool:
return True
model = A2C("MlpPolicy", "CartPole-v1", tensorboard_log="runs/", verbose=1)
model.learn(total_timesteps=int(5e4), callback=HParamCallback())
Directly Accessing The Summary Writer Directly Accessing The Summary Writer
------------------------------------- -------------------------------------

View file

@ -3,24 +3,29 @@
Changelog Changelog
========== ==========
Release 1.6.1a0 (WIP) Release 1.6.1a3 (WIP)
--------------------------- ---------------------------
Breaking Changes: Breaking Changes:
^^^^^^^^^^^^^^^^^ ^^^^^^^^^^^^^^^^^
- Switched minimum tensorboard version to 2.9.1
New Features: New Features:
^^^^^^^^^^^^^ ^^^^^^^^^^^^^
- Support logging hyperparameters to tensorboard (@timothe-chaumont)
- Added checkpoints for replay buffer and ``VecNormalize`` statistics (@anand-bala)
SB3-Contrib SB3-Contrib
^^^^^^^^^^^ ^^^^^^^^^^^
Bug Fixes: Bug Fixes:
^^^^^^^^^^ ^^^^^^^^^^
- Fixed issue where ``PPO`` gives NaN if rollout buffer provides a batch of size 1 (@hughperkins)
- Fixed the issue that ``predict`` does not always return action as ``np.ndarray`` (@qgallouedec) - Fixed the issue that ``predict`` does not always return action as ``np.ndarray`` (@qgallouedec)
- Fixed division by zero error when computing FPS when a small number of time has elapsed in operating systems with low-precision timers. - Fixed division by zero error when computing FPS when a small number of time has elapsed in operating systems with low-precision timers.
- Added multidimensional action space support (@qgallouedec) - Added multidimensional action space support (@qgallouedec)
- Fixed missing verbose parameter passing in the ``EvalCallback`` constructor (@burakdmb) - Fixed missing verbose parameter passing in the ``EvalCallback`` constructor (@burakdmb)
- Fixed the issue that when updating the target network in DQN, SAC, TD3, the ``running_mean`` and ``running_var`` properties of batch norm layers are not updated (@honglu2875)
Deprecations: Deprecations:
^^^^^^^^^^^^^ ^^^^^^^^^^^^^
@ -33,12 +38,12 @@ Others:
Documentation: Documentation:
^^^^^^^^^^^^^^ ^^^^^^^^^^^^^^
- Added an example of callback that logs hyperparameters to tensorboard. (@timothe-chaumont)
- Fixed typo in docstring "nature" -> "Nature" (@Melanol) - Fixed typo in docstring "nature" -> "Nature" (@Melanol)
- Added info on split tensorboard logs into (@Melanol) - Added info on split tensorboard logs into (@Melanol)
- Fixed typo in ppo doc (@francescoluciano) - Fixed typo in ppo doc (@francescoluciano)
- Fixed typo in install doc(@jlp-ue) - Fixed typo in install doc(@jlp-ue)
Release 1.6.0 (2022-07-11) Release 1.6.0 (2022-07-11)
--------------------------- ---------------------------
@ -1024,4 +1029,5 @@ And all the contributors:
@eleurent @ac-93 @cove9988 @theDebugger811 @hsuehch @Demetrio92 @thomasgubler @IperGiove @ScheiklP @eleurent @ac-93 @cove9988 @theDebugger811 @hsuehch @Demetrio92 @thomasgubler @IperGiove @ScheiklP
@simoninithomas @armandpl @manuel-delverme @Gautam-J @gianlucadecola @buoyancy99 @caburu @xy9485 @simoninithomas @armandpl @manuel-delverme @Gautam-J @gianlucadecola @buoyancy99 @caburu @xy9485
@Gregwar @ycheng517 @quantitative-technologies @bcollazo @git-thor @TibiGG @cool-RR @MWeltevrede @Gregwar @ycheng517 @quantitative-technologies @bcollazo @git-thor @TibiGG @cool-RR @MWeltevrede
@Melanol @qgallouedec @francescoluciano @jlp-ue @burakdmb @Melanol @qgallouedec @francescoluciano @jlp-ue @burakdmb @timothe-chaumont @honglu2875
@anand-bala @hughperkins

View file

@ -122,10 +122,7 @@ setup(
"autorom[accept-rom-license]~=0.4.2", "autorom[accept-rom-license]~=0.4.2",
"pillow", "pillow",
# Tensorboard support # Tensorboard support
"tensorboard>=2.2.0", "tensorboard>=2.9.1",
# Protobuf >= 4 has breaking changes
# which does play well with tensorboard
"protobuf~=3.19.0",
# Checking memory taken by replay buffer # Checking memory taken by replay buffer
"psutil", "psutil",
], ],

View file

@ -15,7 +15,8 @@ class BaseCallback(ABC):
""" """
Base class for callback. Base class for callback.
:param verbose: :param verbose: Verbosity of the output (set to 1 for info messages,
2 for debug)
""" """
def __init__(self, verbose: int = 0): def __init__(self, verbose: int = 0):
@ -214,6 +215,10 @@ class CheckpointCallback(BaseCallback):
""" """
Callback for saving a model every ``save_freq`` calls Callback for saving a model every ``save_freq`` calls
to ``env.step()``. to ``env.step()``.
By default, it only saves model checkpoints,
you need to pass ``save_replay_buffer=True``,
and ``save_vecnormalize=True`` to also save replay buffer checkpoints
and normalization statistics checkpoints.
.. warning:: .. warning::
@ -221,29 +226,67 @@ class CheckpointCallback(BaseCallback):
will effectively correspond to ``n_envs`` steps. will effectively correspond to ``n_envs`` steps.
To account for that, you can use ``save_freq = max(save_freq // n_envs, 1)`` To account for that, you can use ``save_freq = max(save_freq // n_envs, 1)``
:param save_freq: :param save_freq: Save checkpoints every ``save_freq`` call of the callback.
:param save_path: Path to the folder where the model will be saved. :param save_path: Path to the folder where the model will be saved.
:param name_prefix: Common prefix to the saved models :param name_prefix: Common prefix to the saved models
:param verbose: :param save_replay_buffer: Save the model replay buffer
:param save_vecnormalize: Save the ``VecNormalize`` statistics
:param verbose: Verbosity of the output (set to 2 for debug messages)
""" """
def __init__(self, save_freq: int, save_path: str, name_prefix: str = "rl_model", verbose: int = 0): def __init__(
self,
save_freq: int,
save_path: str,
name_prefix: str = "rl_model",
save_replay_buffer: bool = False,
save_vecnormalize: bool = False,
verbose: int = 0,
):
super().__init__(verbose) super().__init__(verbose)
self.save_freq = save_freq self.save_freq = save_freq
self.save_path = save_path self.save_path = save_path
self.name_prefix = name_prefix self.name_prefix = name_prefix
self.save_replay_buffer = save_replay_buffer
self.save_vecnormalize = save_vecnormalize
def _init_callback(self) -> None: def _init_callback(self) -> None:
# Create folder if needed # Create folder if needed
if self.save_path is not None: if self.save_path is not None:
os.makedirs(self.save_path, exist_ok=True) os.makedirs(self.save_path, exist_ok=True)
def _checkpoint_path(self, checkpoint_type: str = "", extension: str = "") -> str:
"""
Helper to get checkpoint path for each type of checkpoint.
:param checkpoint_type: empty for the model, "replay_buffer_"
or "vecnormalize_" for the other checkpoints.
:param extension: Checkpoint file extension (zip for model, pkl for others)
:return: Path to the checkpoint
"""
return os.path.join(self.save_path, f"{self.name_prefix}_{checkpoint_type}{self.num_timesteps}_steps.{extension}")
def _on_step(self) -> bool: def _on_step(self) -> bool:
if self.n_calls % self.save_freq == 0: if self.n_calls % self.save_freq == 0:
path = os.path.join(self.save_path, f"{self.name_prefix}_{self.num_timesteps}_steps") model_path = self._checkpoint_path(extension="zip")
self.model.save(path) self.model.save(model_path)
if self.verbose > 1: if self.verbose > 1:
print(f"Saving model checkpoint to {path}") print(f"Saving model checkpoint to {model_path}")
if self.save_replay_buffer and hasattr(self.model, "replay_buffer") and self.model.replay_buffer is not None:
# If model has a replay buffer, save it too
replay_buffer_path = self._checkpoint_path("replay_buffer_", extension="pkl")
self.model.save_replay_buffer(replay_buffer_path)
if self.verbose > 1:
print(f"Saving model replay buffer checkpoint to {replay_buffer_path}")
if self.save_vecnormalize and self.model.get_vec_normalize_env() is not None:
# Save the VecNormalize statistics
vec_normalize_path = self._checkpoint_path("vecnormalize_", extension="pkl")
self.model.get_vec_normalize_env().save(vec_normalize_path)
if self.verbose > 1:
print(f"Saving model VecNormalize to {vec_normalize_path}")
return True return True

View file

@ -14,6 +14,7 @@ from matplotlib import pyplot as plt
try: try:
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
from torch.utils.tensorboard.summary import hparams
except ImportError: except ImportError:
SummaryWriter = None SummaryWriter = None
@ -66,6 +67,22 @@ class Image:
self.dataformats = dataformats self.dataformats = dataformats
class HParam:
"""
Hyperparameter data class storing hyperparameters and metrics in dictionnaries
:param hparam_dict: key-value pairs of hyperparameters to log
:param metric_dict: key-value pairs of metrics to log
A non-empty metrics dict is required to display hyperparameters in the corresponding Tensorboard section.
"""
def __init__(self, hparam_dict: Dict[str, Union[bool, str, float, int, None]], metric_dict: Dict[str, Union[float, int]]):
self.hparam_dict = hparam_dict
if not metric_dict:
raise Exception("`metric_dict` must not be empty to display hyperparameters to the HPARAMS tensorboard tab.")
self.metric_dict = metric_dict
class FormatUnsupportedError(NotImplementedError): class FormatUnsupportedError(NotImplementedError):
""" """
Custom error to display informative message when Custom error to display informative message when
@ -165,6 +182,9 @@ class HumanOutputFormat(KVWriter, SeqWriter):
elif isinstance(value, Image): elif isinstance(value, Image):
raise FormatUnsupportedError(["stdout", "log"], "image") raise FormatUnsupportedError(["stdout", "log"], "image")
elif isinstance(value, HParam):
raise FormatUnsupportedError(["stdout", "log"], "hparam")
elif isinstance(value, float): elif isinstance(value, float):
# Align left # Align left
value_str = f"{value:<8.3g}" value_str = f"{value:<8.3g}"
@ -264,6 +284,8 @@ class JSONOutputFormat(KVWriter):
raise FormatUnsupportedError(["json"], "figure") raise FormatUnsupportedError(["json"], "figure")
if isinstance(value, Image): if isinstance(value, Image):
raise FormatUnsupportedError(["json"], "image") raise FormatUnsupportedError(["json"], "image")
if isinstance(value, HParam):
raise FormatUnsupportedError(["json"], "hparam")
if hasattr(value, "dtype"): if hasattr(value, "dtype"):
if value.shape == () or len(value) == 1: if value.shape == () or len(value) == 1:
# if value is a dimensionless numpy array or of length 1, serialize as a float # if value is a dimensionless numpy array or of length 1, serialize as a float
@ -333,6 +355,9 @@ class CSVOutputFormat(KVWriter):
elif isinstance(value, Image): elif isinstance(value, Image):
raise FormatUnsupportedError(["csv"], "image") raise FormatUnsupportedError(["csv"], "image")
elif isinstance(value, HParam):
raise FormatUnsupportedError(["csv"], "hparam")
elif isinstance(value, str): elif isinstance(value, str):
# escape quotechars by prepending them with another quotechar # escape quotechars by prepending them with another quotechar
value = value.replace(self.quotechar, self.quotechar + self.quotechar) value = value.replace(self.quotechar, self.quotechar + self.quotechar)
@ -389,6 +414,13 @@ class TensorBoardOutputFormat(KVWriter):
if isinstance(value, Image): if isinstance(value, Image):
self.writer.add_image(key, value.image, step, dataformats=value.dataformats) self.writer.add_image(key, value.image, step, dataformats=value.dataformats)
if isinstance(value, HParam):
# we don't use `self.writer.add_hparams` to have control over the log_dir
experiment, session_start_info, session_end_info = hparams(value.hparam_dict, metric_dict=value.metric_dict)
self.writer.file_writer.add_summary(experiment)
self.writer.file_writer.add_summary(session_start_info)
self.writer.file_writer.add_summary(session_end_info)
# Flush the output to the file # Flush the output to the file
self.writer.flush() self.writer.flush()

View file

@ -4,7 +4,7 @@ import platform
import random import random
from collections import deque from collections import deque
from itertools import zip_longest from itertools import zip_longest
from typing import Dict, Iterable, Optional, Tuple, Union from typing import Dict, Iterable, List, Optional, Tuple, Union
import gym import gym
import numpy as np import numpy as np
@ -67,8 +67,8 @@ def update_learning_rate(optimizer: th.optim.Optimizer, learning_rate: float) ->
Update the learning rate for a given optimizer. Update the learning rate for a given optimizer.
Useful when doing linear schedule. Useful when doing linear schedule.
:param optimizer: :param optimizer: Pytorch optimizer
:param learning_rate: :param learning_rate: New learning rate value
""" """
for param_group in optimizer.param_groups: for param_group in optimizer.param_groups:
param_group["lr"] = learning_rate param_group["lr"] = learning_rate
@ -79,8 +79,8 @@ def get_schedule_fn(value_schedule: Union[Schedule, float, int]) -> Schedule:
Transform (if needed) learning rate and clip range (for PPO) Transform (if needed) learning rate and clip range (for PPO)
to callable. to callable.
:param value_schedule: :param value_schedule: Constant value of schedule function
:return: :return: Schedule function (can return constant value)
""" """
# If the passed schedule is a float # If the passed schedule is a float
# create a constant function # create a constant function
@ -104,7 +104,7 @@ def get_linear_fn(start: float, end: float, end_fraction: float) -> Schedule:
:params end_fraction: fraction of ``progress_remaining`` :params end_fraction: fraction of ``progress_remaining``
where end is reached e.g 0.1 then end is reached after 10% where end is reached e.g 0.1 then end is reached after 10%
of the complete training process. of the complete training process.
:return: :return: Linear schedule function.
""" """
def func(progress_remaining: float) -> float: def func(progress_remaining: float) -> float:
@ -121,8 +121,8 @@ def constant_fn(val: float) -> Schedule:
Create a function that returns a constant Create a function that returns a constant
It is useful for learning rate schedule (to avoid code duplication) It is useful for learning rate schedule (to avoid code duplication)
:param val: :param val: constant value
:return: :return: Constant schedule function.
""" """
def func(_): def func(_):
@ -139,7 +139,7 @@ def get_device(device: Union[th.device, str] = "auto") -> th.device:
By default, it tries to use the gpu. By default, it tries to use the gpu.
:param device: One for 'auto', 'cuda', 'cpu' :param device: One for 'auto', 'cuda', 'cpu'
:return: :return: Supported Pytorch device
""" """
# Cuda by default # Cuda by default
if device == "auto": if device == "auto":
@ -386,12 +386,25 @@ def safe_mean(arr: Union[np.ndarray, list, deque]) -> np.ndarray:
Compute the mean of an array if there is at least one element. Compute the mean of an array if there is at least one element.
For empty array, return NaN. It is used for logging only. For empty array, return NaN. It is used for logging only.
:param arr: :param arr: Numpy array or list of values
:return: :return:
""" """
return np.nan if len(arr) == 0 else np.mean(arr) return np.nan if len(arr) == 0 else np.mean(arr)
def get_parameters_by_name(model: th.nn.Module, included_names: Iterable[str]) -> List[th.Tensor]:
"""
Extract parameters from the state dict of ``model``
if the name contains one of the strings in ``included_names``.
:param model: the model where the parameters come from.
:param included_names: substrings of names to include.
:return: List of parameters values (Pytorch tensors)
that matches the queried names.
"""
return [param for name, param in model.state_dict().items() if any([key in name for key in included_names])]
def zip_strict(*iterables: Iterable) -> Iterable: def zip_strict(*iterables: Iterable) -> Iterable:
r""" r"""
``zip()`` function but enforces that iterables are of equal length. ``zip()`` function but enforces that iterables are of equal length.
@ -411,8 +424,8 @@ def zip_strict(*iterables: Iterable) -> Iterable:
def polyak_update( def polyak_update(
params: Iterable[th.nn.Parameter], params: Iterable[th.Tensor],
target_params: Iterable[th.nn.Parameter], target_params: Iterable[th.Tensor],
tau: float, tau: float,
) -> None: ) -> None:
""" """

View file

@ -11,7 +11,7 @@ from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm
from stable_baselines3.common.policies import BasePolicy from stable_baselines3.common.policies import BasePolicy
from stable_baselines3.common.preprocessing import maybe_transpose from stable_baselines3.common.preprocessing import maybe_transpose
from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, Schedule from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, Schedule
from stable_baselines3.common.utils import get_linear_fn, is_vectorized_observation, polyak_update from stable_baselines3.common.utils import get_linear_fn, get_parameters_by_name, is_vectorized_observation, polyak_update
from stable_baselines3.dqn.policies import CnnPolicy, DQNPolicy, MlpPolicy, MultiInputPolicy from stable_baselines3.dqn.policies import CnnPolicy, DQNPolicy, MlpPolicy, MultiInputPolicy
@ -140,6 +140,9 @@ class DQN(OffPolicyAlgorithm):
def _setup_model(self) -> None: def _setup_model(self) -> None:
super()._setup_model() super()._setup_model()
self._create_aliases() self._create_aliases()
# Copy running stats, see GH issue #996
self.batch_norm_stats = get_parameters_by_name(self.q_net, ["running_"])
self.batch_norm_stats_target = get_parameters_by_name(self.q_net_target, ["running_"])
self.exploration_schedule = get_linear_fn( self.exploration_schedule = get_linear_fn(
self.exploration_initial_eps, self.exploration_initial_eps,
self.exploration_final_eps, self.exploration_final_eps,
@ -170,6 +173,8 @@ class DQN(OffPolicyAlgorithm):
self._n_calls += 1 self._n_calls += 1
if self._n_calls % self.target_update_interval == 0: if self._n_calls % self.target_update_interval == 0:
polyak_update(self.q_net.parameters(), self.q_net_target.parameters(), self.tau) polyak_update(self.q_net.parameters(), self.q_net_target.parameters(), self.tau)
# Copy running stats, see GH issue #996
polyak_update(self.batch_norm_stats, self.batch_norm_stats_target, 1.0)
self.exploration_rate = self.exploration_schedule(self._current_progress_remaining) self.exploration_rate = self.exploration_schedule(self._current_progress_remaining)
self.logger.record("rollout/exploration_rate", self.exploration_rate) self.logger.record("rollout/exploration_rate", self.exploration_rate)

View file

@ -137,8 +137,8 @@ class PPO(OnPolicyAlgorithm):
# Check that `n_steps * n_envs > 1` to avoid NaN # Check that `n_steps * n_envs > 1` to avoid NaN
# when doing advantage normalization # when doing advantage normalization
buffer_size = self.env.num_envs * self.n_steps buffer_size = self.env.num_envs * self.n_steps
assert ( assert buffer_size > 1 or (
buffer_size > 1 not normalize_advantage
), f"`n_steps * n_envs` must be greater than 1. Currently n_steps={self.n_steps} and n_envs={self.env.num_envs}" ), f"`n_steps * n_envs` must be greater than 1. Currently n_steps={self.n_steps} and n_envs={self.env.num_envs}"
# Check that the rollout buffer size is a multiple of the mini-batch size # Check that the rollout buffer size is a multiple of the mini-batch size
untruncated_batches = buffer_size // batch_size untruncated_batches = buffer_size // batch_size
@ -210,7 +210,8 @@ class PPO(OnPolicyAlgorithm):
values = values.flatten() values = values.flatten()
# Normalize advantage # Normalize advantage
advantages = rollout_data.advantages advantages = rollout_data.advantages
if self.normalize_advantage: # Normalization does not make sense if mini batchsize == 1, see GH issue #325
if self.normalize_advantage and len(advantages) > 1:
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8) advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
# ratio between old and new policy, should be one at the first iteration # ratio between old and new policy, should be one at the first iteration

View file

@ -10,7 +10,7 @@ from stable_baselines3.common.noise import ActionNoise
from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm
from stable_baselines3.common.policies import BasePolicy from stable_baselines3.common.policies import BasePolicy
from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, Schedule from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, Schedule
from stable_baselines3.common.utils import polyak_update from stable_baselines3.common.utils import get_parameters_by_name, polyak_update
from stable_baselines3.sac.policies import CnnPolicy, MlpPolicy, MultiInputPolicy, SACPolicy from stable_baselines3.sac.policies import CnnPolicy, MlpPolicy, MultiInputPolicy, SACPolicy
@ -152,6 +152,9 @@ class SAC(OffPolicyAlgorithm):
def _setup_model(self) -> None: def _setup_model(self) -> None:
super()._setup_model() super()._setup_model()
self._create_aliases() self._create_aliases()
# Running mean and running var
self.batch_norm_stats = get_parameters_by_name(self.critic, ["running_"])
self.batch_norm_stats_target = get_parameters_by_name(self.critic_target, ["running_"])
# Target entropy is used when learning the entropy coefficient # Target entropy is used when learning the entropy coefficient
if self.target_entropy == "auto": if self.target_entropy == "auto":
# automatically set target entropy if needed # automatically set target entropy if needed
@ -265,7 +268,6 @@ class SAC(OffPolicyAlgorithm):
# Compute actor loss # Compute actor loss
# Alternative: actor_loss = th.mean(log_prob - qf1_pi) # Alternative: actor_loss = th.mean(log_prob - qf1_pi)
# Min over all critic networks
if update_actor: if update_actor:
q_values_pi = th.cat(self.critic(replay_data.observations, actions_pi), dim=1) q_values_pi = th.cat(self.critic(replay_data.observations, actions_pi), dim=1)
# Note: REDQ and DropQ does a mean here # Note: REDQ and DropQ does a mean here
@ -282,6 +284,8 @@ class SAC(OffPolicyAlgorithm):
# Update target networks # Update target networks
if gradient_step % self.target_update_interval == 0: if gradient_step % self.target_update_interval == 0:
polyak_update(self.critic.parameters(), self.critic_target.parameters(), self.tau) polyak_update(self.critic.parameters(), self.critic_target.parameters(), self.tau)
# Copy running stats, see GH issue #996
polyak_update(self.batch_norm_stats, self.batch_norm_stats_target, 1.0)
self._n_updates += gradient_steps self._n_updates += gradient_steps

View file

@ -10,7 +10,7 @@ from stable_baselines3.common.noise import ActionNoise
from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm
from stable_baselines3.common.policies import BasePolicy from stable_baselines3.common.policies import BasePolicy
from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, Schedule from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, Schedule
from stable_baselines3.common.utils import polyak_update from stable_baselines3.common.utils import get_parameters_by_name, polyak_update
from stable_baselines3.td3.policies import CnnPolicy, MlpPolicy, MultiInputPolicy, TD3Policy from stable_baselines3.td3.policies import CnnPolicy, MlpPolicy, MultiInputPolicy, TD3Policy
@ -131,6 +131,11 @@ class TD3(OffPolicyAlgorithm):
def _setup_model(self) -> None: def _setup_model(self) -> None:
super()._setup_model() super()._setup_model()
self._create_aliases() self._create_aliases()
# Running mean and running var
self.actor_batch_norm_stats = get_parameters_by_name(self.actor, ["running_"])
self.critic_batch_norm_stats = get_parameters_by_name(self.critic, ["running_"])
self.actor_batch_norm_stats_target = get_parameters_by_name(self.actor_target, ["running_"])
self.critic_batch_norm_stats_target = get_parameters_by_name(self.critic_target, ["running_"])
def _create_aliases(self) -> None: def _create_aliases(self) -> None:
self.actor = self.policy.actor self.actor = self.policy.actor
@ -189,6 +194,9 @@ class TD3(OffPolicyAlgorithm):
polyak_update(self.critic.parameters(), self.critic_target.parameters(), self.tau) polyak_update(self.critic.parameters(), self.critic_target.parameters(), self.tau)
polyak_update(self.actor.parameters(), self.actor_target.parameters(), self.tau) polyak_update(self.actor.parameters(), self.actor_target.parameters(), self.tau)
# Copy running stats, see GH issue #996
polyak_update(self.critic_batch_norm_stats, self.critic_batch_norm_stats_target, 1.0)
polyak_update(self.actor_batch_norm_stats, self.actor_batch_norm_stats_target, 1.0)
self.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: if len(actor_losses) > 0:

View file

@ -1 +1 @@
1.6.1a0 1.6.1a3

View file

@ -203,3 +203,29 @@ def test_eval_friendly_error():
with pytest.warns(Warning): with pytest.warns(Warning):
with pytest.raises(AssertionError): with pytest.raises(AssertionError):
model.learn(100, callback=eval_callback) model.learn(100, callback=eval_callback)
def test_checkpoint_additional_info(tmp_path):
# tests if the replay buffer and the VecNormalize stats are saved with every checkpoint
dummy_vec_env = DummyVecEnv([lambda: gym.make("CartPole-v1")])
env = VecNormalize(dummy_vec_env)
checkpoint_dir = tmp_path / "checkpoints"
checkpoint_callback = CheckpointCallback(
save_freq=200,
save_path=checkpoint_dir,
save_replay_buffer=True,
save_vecnormalize=True,
verbose=2,
)
model = DQN("MlpPolicy", env, learning_starts=100, buffer_size=500, seed=0)
model.learn(200, callback=checkpoint_callback)
assert os.path.exists(checkpoint_dir / "rl_model_200_steps.zip")
assert os.path.exists(checkpoint_dir / "rl_model_replay_buffer_200_steps.pkl")
assert os.path.exists(checkpoint_dir / "rl_model_vecnormalize_200_steps.pkl")
# Check that checkpoints can be properly loaded
model = DQN.load(checkpoint_dir / "rl_model_200_steps.zip")
model.load_replay_buffer(checkpoint_dir / "rl_model_replay_buffer_200_steps.pkl")
VecNormalize.load(checkpoint_dir / "rl_model_vecnormalize_200_steps.pkl", dummy_vec_env)

View file

@ -17,6 +17,7 @@ from stable_baselines3.common.logger import (
CSVOutputFormat, CSVOutputFormat,
Figure, Figure,
FormatUnsupportedError, FormatUnsupportedError,
HParam,
HumanOutputFormat, HumanOutputFormat,
Image, Image,
Logger, Logger,
@ -296,6 +297,19 @@ def test_report_figure_to_unsupported_format_raises_error(tmp_path, unsupported_
writer.close() writer.close()
@pytest.mark.parametrize("unsupported_format", ["stdout", "log", "json", "csv"])
def test_report_hparam_to_unsupported_format_raises_error(tmp_path, unsupported_format):
writer = make_output_format(unsupported_format, tmp_path)
with pytest.raises(FormatUnsupportedError) as exec_info:
hparam_dict = {"learning rate": np.random.random()}
metric_dict = {"train/value_loss": 0}
hparam = HParam(hparam_dict=hparam_dict, metric_dict=metric_dict)
writer.write({"hparam": hparam}, key_excluded={"hparam": ()})
assert unsupported_format in str(exec_info.value)
writer.close()
def test_key_length(tmp_path): def test_key_length(tmp_path):
writer = make_output_format("stdout", tmp_path) writer = make_output_format("stdout", tmp_path)
assert writer.max_length == 36 assert writer.max_length == 36

View file

@ -224,3 +224,29 @@ def test_warn_dqn_multi_env():
buffer_size=100, buffer_size=100,
target_update_interval=1, target_update_interval=1,
) )
def test_ppo_warnings():
"""Test that PPO warns and errors correctly on
problematic rollout buffer sizes"""
# Only 1 step: advantage normalization will return NaN
with pytest.raises(AssertionError):
PPO("MlpPolicy", "Pendulum-v1", n_steps=1)
# batch_size of 1 is allowed when normalize_advantage=False
model = PPO("MlpPolicy", "Pendulum-v1", n_steps=1, batch_size=1, normalize_advantage=False)
model.learn(4)
# Truncated mini-batch
# Batch size 1 yields NaN with normalized advantage because
# torch.std(some_length_1_tensor) == NaN
# advantage normalization is automatically deactivated
# in that case
with pytest.warns(UserWarning, match="there will be a truncated mini-batch of size 1"):
model = PPO("MlpPolicy", "Pendulum-v1", n_steps=64, batch_size=63, verbose=1)
model.learn(64)
loss = model.logger.name_to_value["train/loss"]
assert loss > 0
assert not np.isnan(loss) # check not nan (since nan does not equal nan)

View file

@ -3,6 +3,8 @@ import os
import pytest import pytest
from stable_baselines3 import A2C, PPO, SAC, TD3 from stable_baselines3 import A2C, PPO, SAC, TD3
from stable_baselines3.common.callbacks import BaseCallback
from stable_baselines3.common.logger import HParam
from stable_baselines3.common.utils import get_latest_run_id from stable_baselines3.common.utils import get_latest_run_id
MODEL_DICT = { MODEL_DICT = {
@ -15,6 +17,34 @@ MODEL_DICT = {
N_STEPS = 100 N_STEPS = 100
class HParamCallback(BaseCallback):
def __init__(self):
"""
Saves the hyperparameters and metrics at the start of the training, and logs them to TensorBoard.
"""
super().__init__()
def _on_training_start(self) -> None:
hparam_dict = {
"algorithm": self.model.__class__.__name__,
"learning rate": self.model.learning_rate,
"gamma": self.model.gamma,
}
# define the metrics that will appear in the `HPARAMS` Tensorboard tab by referencing their tag
# Tensorbaord will find & display metrics from the `SCALARS` tab
metric_dict = {
"rollout/ep_len_mean": 0,
}
self.logger.record(
"hparams",
HParam(hparam_dict, metric_dict),
exclude=("stdout", "log", "json", "csv"),
)
def _on_step(self) -> bool:
return True
@pytest.mark.parametrize("model_name", MODEL_DICT.keys()) @pytest.mark.parametrize("model_name", MODEL_DICT.keys())
def test_tensorboard(tmp_path, model_name): def test_tensorboard(tmp_path, model_name):
# Skip if no tensorboard installed # Skip if no tensorboard installed
@ -22,8 +52,13 @@ def test_tensorboard(tmp_path, model_name):
logname = model_name.upper() logname = model_name.upper()
algo, env_id = MODEL_DICT[model_name] algo, env_id = MODEL_DICT[model_name]
model = algo("MlpPolicy", env_id, verbose=1, tensorboard_log=tmp_path) kwargs = {}
model.learn(N_STEPS) if model_name == "ppo":
kwargs["n_steps"] = 64
elif model_name in {"sac", "td3"}:
kwargs["train_freq"] = 2
model = algo("MlpPolicy", env_id, verbose=1, tensorboard_log=tmp_path, **kwargs)
model.learn(N_STEPS, callback=HParamCallback())
model.learn(N_STEPS, reset_num_timesteps=False) model.learn(N_STEPS, reset_num_timesteps=False)
assert os.path.isdir(tmp_path / str(logname + "_1")) assert os.path.isdir(tmp_path / str(logname + "_1"))

View file

@ -143,7 +143,8 @@ def test_dqn_train_with_batch_norm():
policy_kwargs=dict(net_arch=[16, 16], features_extractor_class=FlattenBatchNormDropoutExtractor), policy_kwargs=dict(net_arch=[16, 16], features_extractor_class=FlattenBatchNormDropoutExtractor),
learning_starts=0, learning_starts=0,
seed=1, seed=1,
tau=0, # do not clone the target tau=0.0, # do not clone the target
target_update_interval=100, # Copy the stats to the target
) )
( (
@ -154,6 +155,9 @@ def test_dqn_train_with_batch_norm():
) = clone_dqn_batch_norm_stats(model) ) = clone_dqn_batch_norm_stats(model)
model.learn(total_timesteps=200) model.learn(total_timesteps=200)
# Force stats copy
model.target_update_interval = 1
model._on_step()
( (
q_net_bias_after, q_net_bias_after,
@ -165,8 +169,12 @@ def test_dqn_train_with_batch_norm():
assert ~th.isclose(q_net_bias_before, q_net_bias_after).all() assert ~th.isclose(q_net_bias_before, q_net_bias_after).all()
assert ~th.isclose(q_net_running_mean_before, q_net_running_mean_after).all() assert ~th.isclose(q_net_running_mean_before, q_net_running_mean_after).all()
# No weight update
assert th.isclose(q_net_bias_before, q_net_target_bias_after).all()
assert th.isclose(q_net_target_bias_before, q_net_target_bias_after).all() assert th.isclose(q_net_target_bias_before, q_net_target_bias_after).all()
assert th.isclose(q_net_target_running_mean_before, q_net_target_running_mean_after).all() # Running stat should be copied even when tau=0
assert th.isclose(q_net_running_mean_before, q_net_target_running_mean_before).all()
assert th.isclose(q_net_running_mean_after, q_net_target_running_mean_after).all()
def test_td3_train_with_batch_norm(): def test_td3_train_with_batch_norm():
@ -210,10 +218,12 @@ def test_td3_train_with_batch_norm():
assert ~th.isclose(critic_running_mean_before, critic_running_mean_after).all() assert ~th.isclose(critic_running_mean_before, critic_running_mean_after).all()
assert th.isclose(actor_target_bias_before, actor_target_bias_after).all() assert th.isclose(actor_target_bias_before, actor_target_bias_after).all()
assert th.isclose(actor_target_running_mean_before, actor_target_running_mean_after).all() # Running stat should be copied even when tau=0
assert th.isclose(actor_running_mean_after, actor_target_running_mean_after).all()
assert th.isclose(critic_target_bias_before, critic_target_bias_after).all() assert th.isclose(critic_target_bias_before, critic_target_bias_after).all()
assert th.isclose(critic_target_running_mean_before, critic_target_running_mean_after).all() # Running stat should be copied even when tau=0
assert th.isclose(critic_running_mean_after, critic_target_running_mean_after).all()
def test_sac_train_with_batch_norm(): def test_sac_train_with_batch_norm():
@ -250,10 +260,12 @@ def test_sac_train_with_batch_norm():
assert ~th.isclose(actor_running_mean_before, actor_running_mean_after).all() assert ~th.isclose(actor_running_mean_before, actor_running_mean_after).all()
assert ~th.isclose(critic_bias_before, critic_bias_after).all() assert ~th.isclose(critic_bias_before, critic_bias_after).all()
assert ~th.isclose(critic_running_mean_before, critic_running_mean_after).all() # Running stat should be copied even when tau=0
assert th.isclose(critic_running_mean_before, critic_target_running_mean_before).all()
assert th.isclose(critic_target_bias_before, critic_target_bias_after).all() assert th.isclose(critic_target_bias_before, critic_target_bias_after).all()
assert th.isclose(critic_target_running_mean_before, critic_target_running_mean_after).all() # Running stat should be copied even when tau=0
assert th.isclose(critic_running_mean_after, critic_target_running_mean_after).all()
@pytest.mark.parametrize("model_class", [A2C, PPO]) @pytest.mark.parametrize("model_class", [A2C, PPO])

View file

@ -8,13 +8,19 @@ import torch as th
from gym import spaces from gym import spaces
import stable_baselines3 as sb3 import stable_baselines3 as sb3
from stable_baselines3 import A2C, PPO from stable_baselines3 import A2C
from stable_baselines3.common.atari_wrappers import ClipRewardEnv, MaxAndSkipEnv from stable_baselines3.common.atari_wrappers import ClipRewardEnv, MaxAndSkipEnv
from stable_baselines3.common.env_util import is_wrapped, make_atari_env, make_vec_env, unwrap_wrapper from stable_baselines3.common.env_util import is_wrapped, make_atari_env, make_vec_env, unwrap_wrapper
from stable_baselines3.common.evaluation import evaluate_policy from stable_baselines3.common.evaluation import evaluate_policy
from stable_baselines3.common.monitor import Monitor from stable_baselines3.common.monitor import Monitor
from stable_baselines3.common.noise import ActionNoise, OrnsteinUhlenbeckActionNoise, VectorizedActionNoise from stable_baselines3.common.noise import ActionNoise, OrnsteinUhlenbeckActionNoise, VectorizedActionNoise
from stable_baselines3.common.utils import get_system_info, is_vectorized_observation, polyak_update, zip_strict from stable_baselines3.common.utils import (
get_parameters_by_name,
get_system_info,
is_vectorized_observation,
polyak_update,
zip_strict,
)
from stable_baselines3.common.vec_env import DummyVecEnv, SubprocVecEnv from stable_baselines3.common.vec_env import DummyVecEnv, SubprocVecEnv
@ -322,6 +328,22 @@ def test_vec_noise():
assert len(vec.noises) == num_envs assert len(vec.noises) == num_envs
def test_get_parameters_by_name():
model = th.nn.Sequential(th.nn.Linear(5, 5), th.nn.BatchNorm1d(5))
# Initialize stats
model(th.ones(3, 5))
included_names = ["weight", "bias", "running_"]
# 2 x weight, 2 x bias, 1 x running_mean, 1 x running_var; Ignore num_batches_tracked.
parameters = get_parameters_by_name(model, included_names)
assert len(parameters) == 6
assert th.allclose(parameters[4], model[1].running_mean)
assert th.allclose(parameters[5], model[1].running_var)
parameters = get_parameters_by_name(model, ["running_"])
assert len(parameters) == 2
assert th.allclose(parameters[0], model[1].running_mean)
assert th.allclose(parameters[1], model[1].running_var)
def test_polyak(): def test_polyak():
param1, param2 = th.nn.Parameter(th.ones((5, 5))), th.nn.Parameter(th.zeros((5, 5))) param1, param2 = th.nn.Parameter(th.ones((5, 5))), th.nn.Parameter(th.zeros((5, 5)))
target1, target2 = th.nn.Parameter(th.ones((5, 5))), th.nn.Parameter(th.zeros((5, 5))) target1, target2 = th.nn.Parameter(th.ones((5, 5))), th.nn.Parameter(th.zeros((5, 5)))
@ -366,19 +388,6 @@ def test_is_wrapped():
assert unwrap_wrapper(env, Monitor) == monitor_env assert unwrap_wrapper(env, Monitor) == monitor_env
def test_ppo_warnings():
"""Test that PPO warns and errors correctly on
problematic rollour buffer sizes"""
# Only 1 step: advantage normalization will return NaN
with pytest.raises(AssertionError):
PPO("MlpPolicy", "Pendulum-v1", n_steps=1)
# Truncated mini-batch
with pytest.warns(UserWarning):
PPO("MlpPolicy", "Pendulum-v1", n_steps=6, batch_size=8)
def test_get_system_info(): def test_get_system_info():
info, info_str = get_system_info(print_info=True) info, info_str = get_system_info(print_info=True)
assert info["Stable-Baselines3"] == str(sb3.__version__) assert info["Stable-Baselines3"] == str(sb3.__version__)