mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-05 20:30:42 +00:00
Merge branch 'master' into feat/mps-support
This commit is contained in:
commit
d47c5867dc
26 changed files with 438 additions and 100 deletions
8
Makefile
8
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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -31,53 +31,52 @@ to do inference in another framework.
|
|||
Export to ONNX
|
||||
-----------------
|
||||
|
||||
As of June 2021, ONNX format `doesn't support <https://github.com/onnx/onnx/issues/3033>`_ 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 <https://github.com/DLR-RM/stable-baselines3/blob/f3a35aa786ee41ffff599b99fa1607c067e89074/stable_baselines3/common/policies.py#L621-L637>`_
|
||||
to obtain the action (e.g., convert action logits to action).
|
||||
The following returns normalized actions and doesn't include the `post-processing <https://github.com/DLR-RM/stable-baselines3/blob/a9273f968eaf8c6e04302a07d803eebfca6e7e86/stable_baselines3/common/policies.py#L370-L377>`_ 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 <https://github.com/DLR-RM/stable-baselines3/issues/383>`_ and `GH#1349 <https://github.com/DLR-RM/stable-baselines3/issues/1349>`_.
|
||||
|
||||
|
||||
For more discussion around the topic refer to this `issue. <https://github.com/DLR-RM/stable-baselines3/issues/383>`_
|
||||
|
||||
Trace/Export to C++
|
||||
-------------------
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 <https://www.youtube.com/watch?v=7-PUg9EAa3Y>`_ that covers
|
||||
this section in more details. You can also find the `slides online <https://araffin.github.io/slides/tips-reliable-rl/>`_.
|
||||
|
||||
|
||||
When you try to reproduce a RL paper by implementing the algorithm, the `nuts and bolts of RL research <http://joschu.net/docs/nuts-and-bolts.pdf>`_
|
||||
by John Schulman are quite useful (`video <https://www.youtube.com/watch?v=8EcdaCk9KaQ>`_).
|
||||
|
||||
|
|
@ -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
|
||||
.. _SBX: https://github.com/araffin/sbx
|
||||
|
|
|
|||
|
|
@ -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 <https://github.com/DLR-RM/stable-baselines3/issues/1573>`_ 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 <TimeLimit<OrderEnforcing<PassiveEnvChecker<CartPoleEnv<CartPole-v1>>>>>
|
||||
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 <https://github.com/DLR-RM/stable-baselines3/issues/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
|
||||
--------------------------------
|
||||
|
||||
|
|
|
|||
|
|
@ -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 <https://wandb.ai/openrlbenchmark/sbx/reports/SBX-TD3-Influence-of-policy-net--Vmlldzo2NDg1Mzk3>`_) 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 <https://wandb.ai/openrlbenchmark/sbx/reports/SBX-TD3-RL-Zoo-v2-3-0a0-vs-SB3-TD3-RL-Zoo-2-2-1---Vmlldzo2MjUyNTQx>`_)
|
||||
|
||||
|
||||
|
||||
- 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
|
||||
|
|
|
|||
|
|
@ -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/
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
6
setup.py
6
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",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
"""
|
||||
Helpers for dealing with vectorized environments.
|
||||
"""
|
||||
|
||||
from collections import OrderedDict
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -1 +1 @@
|
|||
2.2.1
|
||||
2.3.0
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in a new issue