diff --git a/Makefile b/Makefile index fe9f6ae..e0f6b2b 100644 --- a/Makefile +++ b/Makefile @@ -18,19 +18,19 @@ type: mypy lint: # stop the build if there are Python syntax errors or undefined names # see https://www.flake8rules.com/ - ruff ${LINT_PATHS} --select=E9,F63,F7,F82 --show-source + ruff check ${LINT_PATHS} --select=E9,F63,F7,F82 --output-format=full # exit-zero treats all errors as warnings. - ruff ${LINT_PATHS} --exit-zero + ruff check ${LINT_PATHS} --exit-zero format: # Sort imports - ruff --select I ${LINT_PATHS} --fix + ruff check --select I ${LINT_PATHS} --fix # Reformat using black black ${LINT_PATHS} check-codestyle: # Sort imports - ruff --select I ${LINT_PATHS} + ruff check --select I ${LINT_PATHS} # Reformat using black black --check ${LINT_PATHS} diff --git a/README.md b/README.md index 4f42708..6e55f10 100644 --- a/README.md +++ b/README.md @@ -127,7 +127,7 @@ import gymnasium as gym from stable_baselines3 import PPO -env = gym.make("CartPole-v1") +env = gym.make("CartPole-v1", render_mode="human") model = PPO("MlpPolicy", env, verbose=1) model.learn(total_timesteps=10_000) diff --git a/docs/guide/callbacks.rst b/docs/guide/callbacks.rst index 239966a..472f421 100644 --- a/docs/guide/callbacks.rst +++ b/docs/guide/callbacks.rst @@ -29,24 +29,25 @@ You can find two examples of custom callbacks in the documentation: one for savi :param verbose: Verbosity level: 0 for no output, 1 for info messages, 2 for debug messages """ - def __init__(self, verbose=0): + def __init__(self, verbose: int = 0): super().__init__(verbose) # Those variables will be accessible in the callback # (they are defined in the base class) # The RL model # self.model = None # type: BaseAlgorithm # An alias for self.model.get_env(), the environment used for training - # self.training_env = None # type: Union[gym.Env, VecEnv, None] + # self.training_env # type: VecEnv # Number of time the callback was called # self.n_calls = 0 # type: int + # num_timesteps = n_envs * n times env.step() was called # self.num_timesteps = 0 # type: int # local and global variables - # self.locals = None # type: Dict[str, Any] - # self.globals = None # type: Dict[str, Any] + # self.locals = {} # type: Dict[str, Any] + # self.globals = {} # type: Dict[str, Any] # The logger object, used to report things in the terminal - # self.logger = None # stable_baselines3.common.logger - # # Sometimes, for event callback, it is useful - # # to have access to the parent object + # self.logger # type: stable_baselines3.common.logger.Logger + # Sometimes, for event callback, it is useful + # to have access to the parent object # self.parent = None # type: Optional[BaseCallback] def _on_training_start(self) -> None: diff --git a/docs/guide/export.rst b/docs/guide/export.rst index cccf300..88a02fe 100644 --- a/docs/guide/export.rst +++ b/docs/guide/export.rst @@ -31,53 +31,52 @@ to do inference in another framework. Export to ONNX ----------------- -As of June 2021, ONNX format `doesn't support `_ exporting models that use the ``broadcast_tensors`` functionality of pytorch. So in order to export the trained stable-baseline3 models in the ONNX format, we need to first remove the layers that use broadcasting. This can be done by creating a class that removes the unsupported layers. -The following examples are for ``MlpPolicy`` only, and are general examples. Note that you have to preprocess the observation the same way stable-baselines3 agent does (see ``common.preprocessing.preprocess_obs``). +If you are using PyTorch 2.0+ and ONNX Opset 14+, you can easily export SB3 policies using the following code: -For PPO, assuming a shared feature extractor. .. warning:: - The following example is for continuous actions only. - When using discrete or binary actions, you must do some `post-processing `_ - to obtain the action (e.g., convert action logits to action). + The following returns normalized actions and doesn't include the `post-processing `_ step that is done with continuous actions + (clip or unscale the action to the correct space). .. code-block:: python import torch as th + from typing import Tuple from stable_baselines3 import PPO + from stable_baselines3.common.policies import BasePolicy - class OnnxablePolicy(th.nn.Module): - def __init__(self, extractor, action_net, value_net): + class OnnxableSB3Policy(th.nn.Module): + def __init__(self, policy: BasePolicy): super().__init__() - self.extractor = extractor - self.action_net = action_net - self.value_net = value_net + self.policy = policy - def forward(self, observation): - # NOTE: You may have to process (normalize) observation in the correct - # way before using this. See `common.preprocessing.preprocess_obs` - action_hidden, value_hidden = self.extractor(observation) - return self.action_net(action_hidden), self.value_net(value_hidden) + def forward(self, observation: th.Tensor) -> Tuple[th.Tensor, th.Tensor, th.Tensor]: + # NOTE: Preprocessing is included, but postprocessing + # (clipping/inscaling actions) is not, + # If needed, you also need to transpose the images so that they are channel first + # use deterministic=False if you want to export the stochastic policy + # policy() returns `actions, values, log_prob` for PPO + return self.policy(observation, deterministic=True) # Example: model = PPO("MlpPolicy", "Pendulum-v1") + PPO("MlpPolicy", "Pendulum-v1").save("PathToTrainedModel") model = PPO.load("PathToTrainedModel.zip", device="cpu") - onnxable_model = OnnxablePolicy( - model.policy.mlp_extractor, model.policy.action_net, model.policy.value_net - ) + + onnx_policy = OnnxableSB3Policy(model.policy) observation_size = model.observation_space.shape dummy_input = th.randn(1, *observation_size) th.onnx.export( - onnxable_model, + onnx_policy, dummy_input, "my_ppo_model.onnx", - opset_version=9, + opset_version=17, input_names=["input"], ) @@ -93,7 +92,13 @@ For PPO, assuming a shared feature extractor. observation = np.zeros((1, *observation_size)).astype(np.float32) ort_sess = ort.InferenceSession(onnx_path) - action, value = ort_sess.run(None, {"input": observation}) + actions, values, log_prob = ort_sess.run(None, {"input": observation}) + + print(actions, values, log_prob) + + # Check that the predictions are the same + with th.no_grad(): + print(model.policy(th.as_tensor(observation), deterministic=True)) For SAC the procedure is similar. The example shown only exports the actor network as the actor is sufficient to roll out the trained policies. @@ -108,23 +113,16 @@ For SAC the procedure is similar. The example shown only exports the actor netwo class OnnxablePolicy(th.nn.Module): def __init__(self, actor: th.nn.Module): super().__init__() - # Removing the flatten layer because it can't be onnxed - self.actor = th.nn.Sequential( - actor.latent_pi, - actor.mu, - # For gSDE - # th.nn.Hardtanh(min_val=-actor.clip_mean, max_val=actor.clip_mean), - # Squash the output - th.nn.Tanh(), - ) + self.actor = actor def forward(self, observation: th.Tensor) -> th.Tensor: - # NOTE: You may have to process (normalize) observation in the correct - # way before using this. See `common.preprocessing.preprocess_obs` - return self.actor(observation) + # NOTE: You may have to postprocess (unnormalize) actions + # to the correct bounds (see commented code below) + return self.actor(observation, deterministic=True) # Example: model = SAC("MlpPolicy", "Pendulum-v1") + SAC("MlpPolicy", "Pendulum-v1").save("PathToTrainedModel.zip") model = SAC.load("PathToTrainedModel.zip", device="cpu") onnxable_model = OnnxablePolicy(model.policy.actor) @@ -134,7 +132,7 @@ For SAC the procedure is similar. The example shown only exports the actor netwo onnxable_model, dummy_input, "my_sac_actor.onnx", - opset_version=9, + opset_version=17, input_names=["input"], ) @@ -147,10 +145,23 @@ For SAC the procedure is similar. The example shown only exports the actor netwo observation = np.zeros((1, *observation_size)).astype(np.float32) ort_sess = ort.InferenceSession(onnx_path) - action = ort_sess.run(None, {"input": observation}) + scaled_action = ort_sess.run(None, {"input": observation})[0] + + print(scaled_action) + + # Post-process: rescale to correct space + # Rescale the action from [-1, 1] to [low, high] + # low, high = model.action_space.low, model.action_space.high + # post_processed_action = low + (0.5 * (scaled_action + 1.0) * (high - low)) + + # Check that the predictions are the same + with th.no_grad(): + print(model.actor(th.as_tensor(observation), deterministic=True)) + + +For more discussion around the topic, please refer to `GH#383 `_ and `GH#1349 `_. -For more discussion around the topic refer to this `issue. `_ Trace/Export to C++ ------------------- diff --git a/docs/guide/integrations.rst b/docs/guide/integrations.rst index 14573cd..9f864a2 100644 --- a/docs/guide/integrations.rst +++ b/docs/guide/integrations.rst @@ -70,8 +70,10 @@ Installation .. code-block:: bash + # Download model and save it into the logs/ folder - python -m rl_zoo3.load_from_hub --algo a2c --env LunarLander-v2 -orga sb3 -f logs/ + # Only use TRUST_REMOTE_CODE=True with HF models that can be trusted (here the SB3 organization) + TRUST_REMOTE_CODE=True python -m rl_zoo3.load_from_hub --algo a2c --env LunarLander-v2 -orga sb3 -f logs/ # Test the agent python -m rl_zoo3.enjoy --algo a2c --env LunarLander-v2 -f logs/ # Push model, config and hyperparameters to the hub @@ -86,12 +88,19 @@ For instance ``sb3/demo-hf-CartPole-v1``: .. code-block:: python + import os + import gymnasium as gym from huggingface_sb3 import load_from_hub from stable_baselines3 import PPO from stable_baselines3.common.evaluation import evaluate_policy + + # Allow the use of `pickle.load()` when downloading model from the hub + # Please make sure that the organization from which you download can be trusted + os.environ["TRUST_REMOTE_CODE"] = "True" + # Retrieve the model from the hub ## repo_id = id of the model repository from the Hugging Face Hub (repo_id = {organization}/{repo_name}) ## filename = name of the model zip file from the repository diff --git a/docs/guide/rl_tips.rst b/docs/guide/rl_tips.rst index ce6f43e..ae37640 100644 --- a/docs/guide/rl_tips.rst +++ b/docs/guide/rl_tips.rst @@ -252,6 +252,12 @@ A better solution would be to use a squashing function (cf ``SAC``) or a Beta di Tips and Tricks when implementing an RL algorithm ================================================= +.. note:: + + We have a `video on YouTube about reliable RL `_ that covers + this section in more details. You can also find the `slides online `_. + + When you try to reproduce a RL paper by implementing the algorithm, the `nuts and bolts of RL research `_ by John Schulman are quite useful (`video `_). @@ -282,4 +288,4 @@ in RL with discrete actions: 3. Pong (one of the easiest Atari game) 4. other Atari games (e.g. Breakout) -.. _SBX: https://github.com/araffin/sbx \ No newline at end of file +.. _SBX: https://github.com/araffin/sbx diff --git a/docs/guide/vec_envs.rst b/docs/guide/vec_envs.rst index 792fede..c04001c 100644 --- a/docs/guide/vec_envs.rst +++ b/docs/guide/vec_envs.rst @@ -96,6 +96,90 @@ SB3 VecEnv API is actually close to Gym 0.21 API but differs to Gym 0.26+ API: ``vec_env.env_method("method_name", args1, args2, kwargs1=kwargs1)`` and ``vec_env.set_attr("attribute_name", new_value)``. +Modifying Vectorized Environments Attributes +-------------------------------------------- + +If you plan to `modify the attributes of an environment `_ while it is used (e.g., modifying an attribute specifying the task carried out for a portion of training when doing multi-task learning, or +a parameter of the environment dynamics), you must expose a setter method. +In fact, directly accessing the environment attribute in the callback can lead to unexpected behavior because environments can be wrapped (using gym or VecEnv wrappers, the ``Monitor`` wrapper being one example). + +Consider the following example for a custom env: + +.. code-block:: python + + import gymnasium as gym + from gymnasium import spaces + + from stable_baselines3.common.env_util import make_vec_env + + + class MyMultiTaskEnv(gym.Env): + + def __init__(self): + super().__init__() + """ + A state and action space for robotic locomotion. + The multi-task twist is that the policy would need to adapt to different terrains, each with its own + friction coefficient, mu. + The friction coefficient is the only parameter that changes between tasks. + mu is a scalar between 0 and 1, and during training a callback is used to update mu. + """ + ... + + def step(self, action): + # Do something, depending on the action and current value of mu the next state is computed + return self._get_obs(), reward, done, truncated, info + + def set_mu(self, new_mu: float) -> None: + # Note: this value should be used only at the next reset + self.mu = new_mu + + # Example of wrapped env + # env is of type >>>> + env = gym.make("CartPole-v1") + # To access the base env, without wrapper, you should use `.unwrapped` + # or env.get_wrapper_attr("gravity") to include wrappers + env.unwrapped.gravity + # SB3 uses VecEnv for training, where `env.unwrapped.x = new_value` cannot be used to set an attribute + # therefore, you should expose a setter like `set_mu` to properly set an attribute + vec_env = make_vec_env(MyMultiTaskEnv) + # Print current mu value + # Note: you should use vec_env.env_method("get_wrapper_attr", "mu") in Gymnasium v1.0 + print(vec_env.env_method("get_wrapper_attr", "mu")) + # Change `mu` attribute via the setter + vec_env.env_method("set_mu", "mu", 0.1) + + +In this example ``env.mu`` cannot be accessed/changed directly because it is wrapped in a ``VecEnv`` and because it could be wrapped with other wrappers (see `GH#1573 `_ for a longer explanation). +Instead, the callback should use the ``set_mu`` method via the ``env_method`` method for Vectorized Environments. + +.. code-block:: python + + from itertools import cycle + + class ChangeMuCallback(BaseCallback): + """ + This callback changes the value of mu during training looping + through a list of values until training is aborted. + The environment is implemented so that the impact of changing + the value of mu mid-episode is visible only after the episode is over + and the reset method has been called. + """" + def __init__(self): + super().__init__() + # An iterator that contains the different of the friction coefficient + self.mus = cycle([0.1, 0.2, 0.5, 0.13, 0.9]) + + def _on_step(self): + # Note: in practice, you should not change this value at every step + # but rather depending on some events/metrics like agent performance/episode termination + # both accessible via the `self.logger` or `self.locals` variables + self.training_env.env_method("set_mu", next(self.mus)) + +This callback can then be used to safely modify environment attributes during training since +it calls the environment setter method. + + Vectorized Environments Wrappers -------------------------------- diff --git a/docs/misc/changelog.rst b/docs/misc/changelog.rst index 410a5df..90a1953 100644 --- a/docs/misc/changelog.rst +++ b/docs/misc/changelog.rst @@ -3,6 +3,97 @@ Changelog ========== +Release 2.3.0 (2024-03-31) +-------------------------- + +**New defaults hyperparameters for DDPG, TD3 and DQN** + + +Breaking Changes: +^^^^^^^^^^^^^^^^^ +- The defaults hyperparameters of ``TD3`` and ``DDPG`` have been changed to be more consistent with ``SAC`` + +.. code-block:: python + + # SB3 < 2.3.0 default hyperparameters + # model = TD3("MlpPolicy", env, train_freq=(1, "episode"), gradient_steps=-1, batch_size=100) + # SB3 >= 2.3.0: + model = TD3("MlpPolicy", env, train_freq=1, gradient_steps=1, batch_size=256) + +.. note:: + + Two inconsistencies remain: the default network architecture for ``TD3/DDPG`` is ``[400, 300]`` instead of ``[256, 256]`` for SAC (for backward compatibility reasons, see `report on the influence of the network size `_) and the default learning rate is 1e-3 instead of 3e-4 for SAC (for performance reasons, see `W&B report on the influence of the lr `_) + + + +- The default ``learning_starts`` parameter of ``DQN`` have been changed to be consistent with the other offpolicy algorithms + + +.. code-block:: python + + # SB3 < 2.3.0 default hyperparameters, 50_000 corresponded to Atari defaults hyperparameters + # model = DQN("MlpPolicy", env, learning_starts=50_000) + # SB3 >= 2.3.0: + model = DQN("MlpPolicy", env, learning_starts=100) + +- For safety, ``torch.load()`` is now called with ``weights_only=True`` when loading torch tensors, + policy ``load()`` still uses ``weights_only=False`` as gymnasium imports are required for it to work +- When using ``huggingface_sb3``, you will now need to set ``TRUST_REMOTE_CODE=True`` when downloading models from the hub, as ``pickle.load`` is not safe. + + +New Features: +^^^^^^^^^^^^^ +- Log success rate ``rollout/success_rate`` when available for on policy algorithms (@corentinlger) + +Bug Fixes: +^^^^^^^^^^ +- Fixed ``monitor_wrapper`` argument that was not passed to the parent class, and dones argument that wasn't passed to ``_update_into_buffer`` (@corentinlger) + +`SB3-Contrib`_ +^^^^^^^^^^^^^^ +- Added ``rollout_buffer_class`` and ``rollout_buffer_kwargs`` arguments to MaskablePPO +- Fixed ``train_freq`` type annotation for tqc and qrdqn (@Armandpl) +- Fixed ``sb3_contrib/common/maskable/*.py`` type annotations +- Fixed ``sb3_contrib/ppo_mask/ppo_mask.py`` type annotations +- Fixed ``sb3_contrib/common/vec_env/async_eval.py`` type annotations +- Add some additional notes about ``MaskablePPO`` (evaluation and multi-process) (@icheered) + + +`RL Zoo`_ +^^^^^^^^^ +- Updated defaults hyperparameters for TD3/DDPG to be more consistent with SAC +- Upgraded MuJoCo envs hyperparameters to v4 (pre-trained agents need to be updated) +- Added test dependencies to `setup.py` (@power-edge) +- Simplify dependencies of `requirements.txt` (remove duplicates from `setup.py`) + +`SBX`_ (SB3 + Jax) +^^^^^^^^^^^^^^^^^^ +- Added support for ``MultiDiscrete`` and ``MultiBinary`` action spaces to PPO +- Added support for large values for gradient_steps to SAC, TD3, and TQC +- Fix ``train()`` signature and update type hints +- Fix replay buffer device at load time +- Added flatten layer +- Added ``CrossQ`` + +Deprecations: +^^^^^^^^^^^^^ + +Others: +^^^^^^^ +- Updated black from v23 to v24 +- Updated ruff to >= v0.3.1 +- Updated env checker for (multi)discrete spaces with non-zero start. + +Documentation: +^^^^^^^^^^^^^^ +- Added a paragraph on modifying vectorized environment parameters via setters (@fracapuano) +- Updated callback code example +- Updated export to ONNX documentation, it is now much simpler to export SB3 models with newer ONNX Opset! +- Added video link to "Practical Tips for Reliable Reinforcement Learning" video +- Added ``render_mode="human"`` in the README example (@marekm4) +- Fixed docstring signature for sum_independent_dims (@stagoverflow) +- Updated docstring description for ``log_interval`` in the base class (@rushitnshah). + Release 2.2.1 (2023-11-17) -------------------------- **Support for options at reset, bug fixes and better error messages** @@ -1492,7 +1583,7 @@ And all the contributors: @flodorner @KuKuXia @NeoExtended @PartiallyTyped @mmcenta @richardwu @kinalmehta @rolandgvc @tkelestemur @mloo3 @tirafesi @blurLake @koulakis @joeljosephjin @shwang @rk37 @andyshih12 @RaphaelWag @xicocaio @diditforlulz273 @liorcohen5 @ManifoldFR @mloo3 @SwamyDev @wmmc88 @megan-klaiber @thisray -@tfederico @hn2 @LucasAlegre @AptX395 @zampanteymedio @JadenTravnik @decodyng @ardabbour @lorenz-h @mschweizer @lorepieri8 @vwxyzjn +@tfederico @hn2 @LucasAlegre @AptX395 @zampanteymedio @fracapuano @JadenTravnik @decodyng @ardabbour @lorenz-h @mschweizer @lorepieri8 @vwxyzjn @ShangqunYu @PierreExeter @JacopoPan @ltbd78 @tom-doerr @Atlis @liusida @09tangriro @amy12xx @juancroldan @benblack769 @bstee615 @c-rizz @skandermoalla @MihaiAnca13 @davidblom603 @ayeright @cyprienc @wkirgsn @AechPro @CUN-bjy @batu @IljaAvadiev @timokau @kachayev @cleversonahum @@ -1504,3 +1595,4 @@ And all the contributors: @anand-bala @hughperkins @sidney-tio @AlexPasqua @dominicgkerr @Akhilez @Rocamonde @tobirohrer @ZikangXiong @ReHoss @DavyMorgan @luizapozzobon @Bonifatius94 @theSquaredError @harveybellini @DavyMorgan @FieteO @jonasreiher @npit @WeberSamuel @troiganto @lutogniew @lbergmann1 @lukashass @BertrandDecoster @pseudo-rnd-thoughts @stefanbschneider @kyle-he @PatrickHelm @corentinlger +@marekm4 @stagoverflow @rushitnshah diff --git a/docs/modules/ppo.rst b/docs/modules/ppo.rst index ace2fcc..b5e6672 100644 --- a/docs/modules/ppo.rst +++ b/docs/modules/ppo.rst @@ -23,7 +23,7 @@ Notes - Original paper: https://arxiv.org/abs/1707.06347 - Clear explanation of PPO on Arxiv Insights channel: https://www.youtube.com/watch?v=5P7I-xPq8u8 -- OpenAI blog post: https://blog.openai.com/openai-baselines-ppo/ +- OpenAI blog post: https://openai.com/research/openai-baselines-ppo - Spinning Up guide: https://spinningup.openai.com/en/latest/algorithms/ppo.html - 37 implementation details blog: https://iclr-blog-track.github.io/2022/03/25/ppo-implementation-details/ diff --git a/pyproject.toml b/pyproject.toml index 1195687..ce0a14e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -3,13 +3,15 @@ line-length = 127 # Assume Python 3.8 target-version = "py38" + +[tool.ruff.lint] # See https://beta.ruff.rs/docs/rules/ select = ["E", "F", "B", "UP", "C90", "RUF"] # B028: Ignore explicit stacklevel` # RUF013: Too many false positives (implicit optional) ignore = ["B028", "RUF013"] -[tool.ruff.per-file-ignores] +[tool.ruff.lint.per-file-ignores] # Default implementation in abstract methods "./stable_baselines3/common/callbacks.py"= ["B027"] "./stable_baselines3/common/noise.py"= ["B027"] @@ -17,7 +19,7 @@ ignore = ["B028", "RUF013"] "./tests/*.py"= ["RUF012", "RUF013"] -[tool.ruff.mccabe] +[tool.ruff.lint.mccabe] # Unlike Flake8, default to a complexity level of 10. max-complexity = 15 diff --git a/setup.py b/setup.py index 5e10ed6..161539a 100644 --- a/setup.py +++ b/setup.py @@ -43,7 +43,7 @@ import gymnasium from stable_baselines3 import PPO -env = gymnasium.make("CartPole-v1") +env = gymnasium.make("CartPole-v1", render_mode="human") model = PPO("MlpPolicy", env, verbose=1) model.learn(total_timesteps=10_000) @@ -120,9 +120,9 @@ setup( # Type check "mypy", # Lint code and sort imports (flake8 and isort replacement) - "ruff>=0.0.288", + "ruff>=0.3.1", # Reformat - "black>=23.9.1,<24", + "black>=24.2.0,<25", ], "docs": [ "sphinx>=5,<8", diff --git a/stable_baselines3/common/base_class.py b/stable_baselines3/common/base_class.py index 5e87599..e6c7d3c 100644 --- a/stable_baselines3/common/base_class.py +++ b/stable_baselines3/common/base_class.py @@ -523,7 +523,10 @@ class BaseAlgorithm(ABC): :param total_timesteps: The total number of samples (env steps) to train on :param callback: callback(s) called at every step with state of the algorithm. - :param log_interval: The number of episodes before logging. + :param log_interval: for on-policy algos (e.g., PPO, A2C, ...) this is the number of + training iterations (i.e., log_interval * n_steps * n_envs timesteps) before logging; + for off-policy algos (e.g., TD3, SAC, ...) this is the number of episodes before + logging. :param tb_log_name: the name of the run for TensorBoard logging :param reset_num_timesteps: whether or not to reset the current timestep number (used in logging) :param progress_bar: Display a progress bar using tqdm and rich. diff --git a/stable_baselines3/common/distributions.py b/stable_baselines3/common/distributions.py index 149345d..132a353 100644 --- a/stable_baselines3/common/distributions.py +++ b/stable_baselines3/common/distributions.py @@ -113,7 +113,7 @@ def sum_independent_dims(tensor: th.Tensor) -> th.Tensor: so we can sum components of the ``log_prob`` or the entropy. :param tensor: shape: (n_batch, n_actions) or (n_batch,) - :return: shape: (n_batch,) + :return: shape: (n_batch,) for (n_batch, n_actions) input, scalar for (n_batch,) input """ if len(tensor.shape) > 1: tensor = tensor.sum(dim=1) diff --git a/stable_baselines3/common/env_checker.py b/stable_baselines3/common/env_checker.py index dc465a1..f24c86e 100644 --- a/stable_baselines3/common/env_checker.py +++ b/stable_baselines3/common/env_checker.py @@ -17,13 +17,37 @@ def _is_numpy_array_space(space: spaces.Space) -> bool: return not isinstance(space, (spaces.Dict, spaces.Tuple)) +def _starts_at_zero(space: Union[spaces.Discrete, spaces.MultiDiscrete]) -> bool: + """ + Return False if a (Multi)Discrete space has a non-zero start. + """ + return np.allclose(space.start, np.zeros_like(space.start)) + + +def _check_non_zero_start(space: spaces.Space, space_type: str = "observation", key: str = "") -> None: + """ + :param space: Observation or action space + :param space_type: information about whether it is an observation or action space + (for the warning message) + :param key: When the observation space comes from a Dict space, we pass the + corresponding key to have more precise warning messages. Defaults to "". + """ + if isinstance(space, (spaces.Discrete, spaces.MultiDiscrete)) and not _starts_at_zero(space): + maybe_key = f"(key='{key}')" if key else "" + warnings.warn( + f"{type(space).__name__} {space_type} space {maybe_key} with a non-zero start (start={space.start}) " + "is not supported by Stable-Baselines3. " + f"You can use a wrapper or update your {space_type} space." + ) + + def _check_image_input(observation_space: spaces.Box, key: str = "") -> None: """ Check that the input will be compatible with Stable-Baselines when the observation is apparently an image. :param observation_space: Observation space - :key: When the observation space comes from a Dict space, we pass the + :param key: When the observation space comes from a Dict space, we pass the corresponding key to have more precise warning messages. Defaults to "". """ if observation_space.dtype != np.uint8: @@ -63,11 +87,7 @@ def _check_unsupported_spaces(env: gym.Env, observation_space: spaces.Space, act for key, space in observation_space.spaces.items(): if isinstance(space, spaces.Dict): nested_dict = True - if isinstance(space, spaces.Discrete) and space.start != 0: - warnings.warn( - f"Discrete observation space (key '{key}') with a non-zero start is not supported by Stable-Baselines3. " - "You can use a wrapper or update your observation space." - ) + _check_non_zero_start(space, "observation", key) if nested_dict: warnings.warn( @@ -87,11 +107,7 @@ def _check_unsupported_spaces(env: gym.Env, observation_space: spaces.Space, act "which is supported by SB3." ) - if isinstance(observation_space, spaces.Discrete) and observation_space.start != 0: - warnings.warn( - "Discrete observation space with a non-zero start is not supported by Stable-Baselines3. " - "You can use a wrapper or update your observation space." - ) + _check_non_zero_start(observation_space, "observation") if isinstance(observation_space, spaces.Sequence): warnings.warn( @@ -100,11 +116,7 @@ def _check_unsupported_spaces(env: gym.Env, observation_space: spaces.Space, act "Note: The checks for returned values are skipped." ) - if isinstance(action_space, spaces.Discrete) and action_space.start != 0: - warnings.warn( - "Discrete action space with a non-zero start is not supported by Stable-Baselines3. " - "You can use a wrapper or update your action space." - ) + _check_non_zero_start(action_space, "action") if not _is_numpy_array_space(action_space): warnings.warn( diff --git a/stable_baselines3/common/on_policy_algorithm.py b/stable_baselines3/common/on_policy_algorithm.py index ddd0f8d..1ba36d5 100644 --- a/stable_baselines3/common/on_policy_algorithm.py +++ b/stable_baselines3/common/on_policy_algorithm.py @@ -92,6 +92,7 @@ class OnPolicyAlgorithm(BaseAlgorithm): use_sde=use_sde, sde_sample_freq=sde_sample_freq, support_multi_env=True, + monitor_wrapper=monitor_wrapper, seed=seed, stats_window_size=stats_window_size, tensorboard_log=tensorboard_log, @@ -200,7 +201,7 @@ class OnPolicyAlgorithm(BaseAlgorithm): if not callback.on_step(): return False - self._update_info_buffer(infos) + self._update_info_buffer(infos, dones) n_steps += 1 if isinstance(self.action_space, spaces.Discrete): @@ -250,6 +251,28 @@ class OnPolicyAlgorithm(BaseAlgorithm): """ raise NotImplementedError + def _dump_logs(self, iteration: int) -> None: + """ + Write log. + + :param iteration: Current logging iteration + """ + assert self.ep_info_buffer is not None + assert self.ep_success_buffer is not None + + time_elapsed = max((time.time_ns() - self.start_time) / 1e9, sys.float_info.epsilon) + fps = int((self.num_timesteps - self._num_timesteps_at_start) / time_elapsed) + self.logger.record("time/iterations", iteration, exclude="tensorboard") + if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0: + 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 len(self.ep_success_buffer) > 0: + self.logger.record("rollout/success_rate", safe_mean(self.ep_success_buffer)) + self.logger.dump(step=self.num_timesteps) + def learn( self: SelfOnPolicyAlgorithm, total_timesteps: int, @@ -285,16 +308,7 @@ class OnPolicyAlgorithm(BaseAlgorithm): # Display training infos if log_interval is not None and iteration % log_interval == 0: assert self.ep_info_buffer is not None - time_elapsed = max((time.time_ns() - self.start_time) / 1e9, sys.float_info.epsilon) - fps = int((self.num_timesteps - self._num_timesteps_at_start) / time_elapsed) - self.logger.record("time/iterations", iteration, exclude="tensorboard") - if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0: - 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") - self.logger.dump(step=self.num_timesteps) + self._dump_logs(iteration) self.train() diff --git a/stable_baselines3/common/policies.py b/stable_baselines3/common/policies.py index 50be01c..e4d62ef 100644 --- a/stable_baselines3/common/policies.py +++ b/stable_baselines3/common/policies.py @@ -173,7 +173,9 @@ class BaseModel(nn.Module): :return: """ device = get_device(device) - saved_variables = th.load(path, map_location=device) + # Note(antonin): we cannot use `weights_only=True` here because we need to allow + # gymnasium imports for the policy to be loaded successfully + saved_variables = th.load(path, map_location=device, weights_only=False) # Create policy object model = cls(**saved_variables["data"]) diff --git a/stable_baselines3/common/save_util.py b/stable_baselines3/common/save_util.py index 0cbf6d4..2d86520 100644 --- a/stable_baselines3/common/save_util.py +++ b/stable_baselines3/common/save_util.py @@ -2,6 +2,7 @@ Save util taken from stable_baselines used to serialize data (class parameters) of model classes """ + import base64 import functools import io @@ -446,7 +447,7 @@ def load_from_zip_file( file_content.seek(0) # Load the parameters with the right ``map_location``. # Remove ".pth" ending with splitext - th_object = th.load(file_content, map_location=device) + th_object = th.load(file_content, map_location=device, weights_only=True) # "tensors.pth" was renamed "pytorch_variables.pth" in v0.9.0, see PR #138 if file_path == "pytorch_variables.pth" or file_path == "tensors.pth": # PyTorch variables (not state_dicts) diff --git a/stable_baselines3/common/type_aliases.py b/stable_baselines3/common/type_aliases.py index d75e115..85d0906 100644 --- a/stable_baselines3/common/type_aliases.py +++ b/stable_baselines3/common/type_aliases.py @@ -1,4 +1,5 @@ """Common aliases for type hints""" + from enum import Enum from typing import TYPE_CHECKING, Any, Callable, Dict, List, NamedTuple, Optional, Protocol, SupportsFloat, Tuple, Union diff --git a/stable_baselines3/common/vec_env/util.py b/stable_baselines3/common/vec_env/util.py index 2a03d8e..855f50e 100644 --- a/stable_baselines3/common/vec_env/util.py +++ b/stable_baselines3/common/vec_env/util.py @@ -1,6 +1,7 @@ """ Helpers for dealing with vectorized environments. """ + from collections import OrderedDict from typing import Any, Dict, List, Tuple diff --git a/stable_baselines3/common/vec_env/vec_frame_stack.py b/stable_baselines3/common/vec_env/vec_frame_stack.py index d412a96..daa2b36 100644 --- a/stable_baselines3/common/vec_env/vec_frame_stack.py +++ b/stable_baselines3/common/vec_env/vec_frame_stack.py @@ -29,7 +29,12 @@ class VecFrameStack(VecEnvWrapper): def step_wait( self, - ) -> Tuple[Union[np.ndarray, Dict[str, np.ndarray]], np.ndarray, np.ndarray, List[Dict[str, Any]],]: + ) -> Tuple[ + Union[np.ndarray, Dict[str, np.ndarray]], + np.ndarray, + np.ndarray, + List[Dict[str, Any]], + ]: observations, rewards, dones, infos = self.venv.step_wait() observations, infos = self.stacked_obs.update(observations, dones, infos) # type: ignore[arg-type] return observations, rewards, dones, infos diff --git a/stable_baselines3/ddpg/ddpg.py b/stable_baselines3/ddpg/ddpg.py index c311b23..2fe2fdf 100644 --- a/stable_baselines3/ddpg/ddpg.py +++ b/stable_baselines3/ddpg/ddpg.py @@ -60,11 +60,11 @@ class DDPG(TD3): learning_rate: Union[float, Schedule] = 1e-3, buffer_size: int = 1_000_000, # 1e6 learning_starts: int = 100, - batch_size: int = 100, + batch_size: int = 256, tau: float = 0.005, gamma: float = 0.99, - train_freq: Union[int, Tuple[int, str]] = (1, "episode"), - gradient_steps: int = -1, + train_freq: Union[int, Tuple[int, str]] = 1, + gradient_steps: int = 1, action_noise: Optional[ActionNoise] = None, replay_buffer_class: Optional[Type[ReplayBuffer]] = None, replay_buffer_kwargs: Optional[Dict[str, Any]] = None, diff --git a/stable_baselines3/dqn/dqn.py b/stable_baselines3/dqn/dqn.py index 42e3d0d..894ed9f 100644 --- a/stable_baselines3/dqn/dqn.py +++ b/stable_baselines3/dqn/dqn.py @@ -79,7 +79,7 @@ class DQN(OffPolicyAlgorithm): env: Union[GymEnv, str], learning_rate: Union[float, Schedule] = 1e-4, buffer_size: int = 1_000_000, # 1e6 - learning_starts: int = 50000, + learning_starts: int = 100, batch_size: int = 32, tau: float = 1.0, gamma: float = 0.99, diff --git a/stable_baselines3/td3/td3.py b/stable_baselines3/td3/td3.py index a06ce67..a61d954 100644 --- a/stable_baselines3/td3/td3.py +++ b/stable_baselines3/td3/td3.py @@ -83,11 +83,11 @@ class TD3(OffPolicyAlgorithm): learning_rate: Union[float, Schedule] = 1e-3, buffer_size: int = 1_000_000, # 1e6 learning_starts: int = 100, - batch_size: int = 100, + batch_size: int = 256, tau: float = 0.005, gamma: float = 0.99, - train_freq: Union[int, Tuple[int, str]] = (1, "episode"), - gradient_steps: int = -1, + train_freq: Union[int, Tuple[int, str]] = 1, + gradient_steps: int = 1, action_noise: Optional[ActionNoise] = None, replay_buffer_class: Optional[Type[ReplayBuffer]] = None, replay_buffer_kwargs: Optional[Dict[str, Any]] = None, diff --git a/stable_baselines3/version.txt b/stable_baselines3/version.txt index c043eea..276cbf9 100644 --- a/stable_baselines3/version.txt +++ b/stable_baselines3/version.txt @@ -1 +1 @@ -2.2.1 +2.3.0 diff --git a/tests/test_envs.py b/tests/test_envs.py index e82ef57..9a61eee 100644 --- a/tests/test_envs.py +++ b/tests/test_envs.py @@ -123,6 +123,8 @@ def test_high_dimension_action_space(): spaces.Dict({"img": spaces.Box(low=0, high=255, shape=(32, 32, 3), dtype=np.uint8)}), # Non zero start index spaces.Discrete(3, start=-1), + # Non zero start index (MultiDiscrete) + spaces.MultiDiscrete([4, 4], start=[1, 0]), # Non zero start index inside a Dict spaces.Dict({"obs": spaces.Discrete(3, start=1)}), ], @@ -164,6 +166,8 @@ def test_non_default_spaces(new_obs_space): spaces.Box(low=np.array([-1, -1, -1]), high=np.array([1, 1, 0.99]), dtype=np.float32), # Non zero start index spaces.Discrete(3, start=-1), + # Non zero start index (MultiDiscrete) + spaces.MultiDiscrete([4, 4], start=[1, 0]), ], ) def test_non_default_action_spaces(new_action_space): @@ -179,7 +183,7 @@ def test_non_default_action_spaces(new_action_space): env.action_space = new_action_space # Discrete action space - if isinstance(new_action_space, spaces.Discrete): + if isinstance(new_action_space, (spaces.Discrete, spaces.MultiDiscrete)): with pytest.warns(UserWarning): check_env(env) return diff --git a/tests/test_logger.py b/tests/test_logger.py index 05bf196..dfd9e55 100644 --- a/tests/test_logger.py +++ b/tests/test_logger.py @@ -14,7 +14,7 @@ from gymnasium import spaces from matplotlib import pyplot as plt from pandas.errors import EmptyDataError -from stable_baselines3 import A2C, DQN +from stable_baselines3 import A2C, DQN, PPO from stable_baselines3.common.env_checker import check_env from stable_baselines3.common.logger import ( DEBUG, @@ -33,6 +33,7 @@ from stable_baselines3.common.logger import ( read_csv, read_json, ) +from stable_baselines3.common.monitor import Monitor KEY_VALUES = { "test": 1, @@ -474,3 +475,92 @@ def test_human_output_format_custom_test_io(base_class): """ assert printed == desired_printed + + +class DummySuccessEnv(gym.Env): + """ + Create a dummy success environment that returns wether True or False for info['is_success'] + at the end of an episode according to its dummy successes list + """ + + def __init__(self, dummy_successes, ep_steps): + """Init the dummy success env + + :param dummy_successes: list of size (n_logs_iterations, n_episodes_per_log) that specifies + the success value of log iteration i at episode j + :param ep_steps: number of steps per episode (to activate truncated) + """ + self.n_steps = 0 + self.log_id = 0 + self.ep_id = 0 + + self.ep_steps = ep_steps + + self.dummy_success = dummy_successes + self.num_logs = len(dummy_successes) + self.ep_per_log = len(dummy_successes[0]) + self.steps_per_log = self.ep_per_log * self.ep_steps + + self.action_space = spaces.Discrete(2) + self.observation_space = spaces.Discrete(2) + + def reset(self, seed=None, options=None): + """ + Reset the env and advance to the next episode_id to get the next dummy success + """ + self.n_steps = 0 + + if self.ep_id == self.ep_per_log: + self.ep_id = 0 + self.log_id = (self.log_id + 1) % self.num_logs + + return self.observation_space.sample(), {} + + def step(self, action): + """ + Step and return a dummy success when an episode is truncated + """ + self.n_steps += 1 + truncated = self.n_steps >= self.ep_steps + + info = {} + if truncated: + maybe_success = self.dummy_success[self.log_id][self.ep_id] + info["is_success"] = maybe_success + self.ep_id += 1 + return self.observation_space.sample(), 0.0, False, truncated, info + + +def test_rollout_success_rate_on_policy_algorithm(tmp_path): + """ + Test if the rollout/success_rate information is correctly logged with on policy algorithms + + To do so, create a dummy environment that takes as argument dummy successes (i.e when an episode) + is going to be successfull or not. + """ + + STATS_WINDOW_SIZE = 10 + # Add dummy successes with 0.3, 0.5 and 0.8 success_rate of length STATS_WINDOW_SIZE + dummy_successes = [ + [True] * 3 + [False] * 7, + [True] * 5 + [False] * 5, + [True] * 8 + [False] * 2, + ] + ep_steps = 64 + + # Monitor the env to track the success info + monitor_file = str(tmp_path / "monitor.csv") + env = Monitor(DummySuccessEnv(dummy_successes, ep_steps), filename=monitor_file, info_keywords=("is_success",)) + + # Equip the model of a custom logger to check the success_rate info + model = PPO("MlpPolicy", env=env, stats_window_size=STATS_WINDOW_SIZE, n_steps=env.steps_per_log, verbose=1) + logger = InMemoryLogger() + model.set_logger(logger) + + # Make the model learn and check that the success rate corresponds to the ratio of dummy successes + model.learn(total_timesteps=env.ep_per_log * ep_steps, log_interval=1) + assert logger.name_to_value["rollout/success_rate"] == 0.3 + model.learn(total_timesteps=env.ep_per_log * ep_steps, log_interval=1) + assert logger.name_to_value["rollout/success_rate"] == 0.5 + model.learn(total_timesteps=env.ep_per_log * ep_steps, log_interval=1) + assert logger.name_to_value["rollout/success_rate"] == 0.8