mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-04 20:23:54 +00:00
Merge branch 'master' into feat/dropq
This commit is contained in:
commit
8311377634
18 changed files with 354 additions and 64 deletions
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
-------------------------------------
|
-------------------------------------
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
5
setup.py
5
setup.py
|
|
@ -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",
|
||||||
],
|
],
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
"""
|
"""
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -1 +1 @@
|
||||||
1.6.1a0
|
1.6.1a3
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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"))
|
||||||
|
|
|
||||||
|
|
@ -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])
|
||||||
|
|
|
||||||
|
|
@ -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__)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue