Implement DQN (#28)

* Created DQN template according to the paper.
Next steps:
- Create Policy
- Complete Training
- Debug

* Changed Base Class

* refactor save, to be consistence with overriding the excluded_save_params function. Do not try to exclude the parameters twice.

* Added simple DQN policy

* Finished learn and train function
- missing correct loss computation

* changed collect_rollouts to work with discrete space

* moved discrete space collect_rollouts to dqn

* basic dqn working

* deleted SDE related code

* added gradient clipping and moved greedy policy to policy

* changed policy to implement target network
and added soft update(in fact standart tau is 1 so hard update)

* fixed policy setup

* rebase target_update_intervall on _n_updates

* adapted all tests
all tests passing

* Move to stable-baseline3

* Fixes for DQN

* Fix tests + add CNNPolicy

* Allow any optimizer for DQN

* added some util functions to create a arbitrary linear schedule, fixed pickle problem with old exploration schedule

* more documentation

* changed buffer dtype

* refactor and document

* Added Sphinx Documentation
Updated changelog.rst

* removed custom collect_rollouts as it is no longer necessary

* Implemented suggestions to clean code and documentation.

* extracted some functions on tests to reduce duplicated code

* added support for exploration_fraction

* Fixed exploration_fraction

* Added documentation

* Fixed get_linear_fn -> proper progress scaling

* Merged master

* Added nature reference

* Changed default parameters to https://www.nature.com/articles/nature14236/tables/1

* Fixed n_updates to be incremented correctly

* Correct train_freq

* Doc update

* added special parameter for DQN in tests

* different fix for test_discrete

* Update docs/modules/dqn.rst

Co-authored-by: Antonin RAFFIN <antonin.raffin@ensta.org>

* Update docs/modules/dqn.rst

Co-authored-by: Antonin RAFFIN <antonin.raffin@ensta.org>

* Update docs/modules/dqn.rst

Co-authored-by: Antonin RAFFIN <antonin.raffin@ensta.org>

* Added RMSProp in optimizer_kwargs, as described in nature paper

* Exploration fraction is inverse of 50.000.000 (total frames) / 1.000.000 (frames with linear schedule) according to nature paper

* Changelog update for buffer dtype

* standard exlude parameters should be always excluded to assure proper saving only if intentionally included by ``include`` parameter

* slightly more iterations on test_discrete to pass the test

* added param use_rms_prop instead of mutable default argument

* forgot alpha

* using huber loss, adam and learning rate 1e-4

* account for train_freq in update_target_network

* Added memory check for both buffers

* Doc updated for buffer allocation

* Added psutil Requirement

* Adapted test_identity.py

* Fixes with new SB3 version

* Fix for tensorboard name

* Convert assert to warning and fix tests

* Refactor off-policy algorithms

* Fixes

* test: remove next_obs in replay buffer

* Update changelog

* Fix tests and use tmp_path where possible

* Fix sampling bug in buffer

* Do not store next obs on episode termination

* Fix replay buffer sampling

* Update comment

* moved epsilon from policy to model

* Update predict method

* Update atari wrappers to match SB2

* Minor edit in the buffers

* Update changelog

* Merge branch 'master' into dqn

* Update DQN to new structure

* Fix tests and remove hardcoded path

* Fix for DQN

* Disable memory efficient replay buffer by default

* Fix docstring

* Add tests for memory efficient buffer

* Update changelog

* Split collect rollout

* Move target update outside `train()` for DQN

* Update changelog

* Update linear schedule doc

* Cleanup DQN code

* Minor edit

* Update version and docker images

Co-authored-by: Antonin RAFFIN <antonin.raffin@ensta.org>
This commit is contained in:
Noah 2020-06-29 11:16:54 +02:00 committed by GitHub
parent e47da426c1
commit 96b771f24e
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
32 changed files with 1278 additions and 273 deletions

View file

@ -1,4 +1,4 @@
image: stablebaselines/stable-baselines3-cpu:0.6.0 image: stablebaselines/stable-baselines3-cpu:0.8.0a1
type-check: type-check:
script: script:

View file

@ -40,7 +40,6 @@ These algorithms will make it easier for the research community and industry to
Please look at the issue for more details. Please look at the issue for more details.
Planned features: Planned features:
- [ ] DQN (almost ready, currently in testing phase)
- [ ] DDPG (you can use its successor TD3 for now) - [ ] DDPG (you can use its successor TD3 for now)
- [ ] HER - [ ] HER

View file

@ -12,6 +12,7 @@ A2C ✔️ ✔️ ✔️ ✔️
PPO ✔️ ✔️ ✔️ ✔️ ✔️ PPO ✔️ ✔️ ✔️ ✔️ ✔️
SAC ✔️ ❌ ❌ ❌ ❌ SAC ✔️ ❌ ❌ ❌ ❌
TD3 ✔️ ❌ ❌ ❌ ❌ TD3 ✔️ ❌ ❌ ❌ ❌
DQN ❌ ✔️ ❌ ❌ ❌
============ =========== ============ ================= =============== ================ ============ =========== ============ ================= =============== ================

View file

@ -33,7 +33,7 @@ notebooks:
Basic Usage: Training, Saving, Loading Basic Usage: Training, Saving, Loading
-------------------------------------- --------------------------------------
In the following example, we will train, save and load a A2C model on the Lunar Lander environment. In the following example, we will train, save and load a DQN model on the Lunar Lander environment.
.. image:: ../_static/img/colab-badge.svg .. image:: ../_static/img/colab-badge.svg
:target: https://colab.research.google.com/github/Stable-Baselines-Team/rl-colab-notebooks/blob/sb3/saving_loading_dqn.ipynb :target: https://colab.research.google.com/github/Stable-Baselines-Team/rl-colab-notebooks/blob/sb3/saving_loading_dqn.ipynb
@ -57,7 +57,7 @@ In the following example, we will train, save and load a A2C model on the Lunar
import gym import gym
from stable_baselines3 import A2C from stable_baselines3 import DQN
from stable_baselines3.common.evaluation import evaluate_policy from stable_baselines3.common.evaluation import evaluate_policy
@ -65,15 +65,15 @@ In the following example, we will train, save and load a A2C model on the Lunar
env = gym.make('LunarLander-v2') env = gym.make('LunarLander-v2')
# Instantiate the agent # Instantiate the agent
model = A2C('MlpPolicy', env, verbose=1) model = DQN('MlpPolicy', env, verbose=1)
# Train the agent # Train the agent
model.learn(total_timesteps=int(2e5)) model.learn(total_timesteps=int(2e5))
# Save the agent # Save the agent
model.save("a2c_lunar") model.save("dqn_lunar")
del model # delete trained model to demonstrate loading del model # delete trained model to demonstrate loading
# Load the trained agent # Load the trained agent
model = A2C.load("a2c_lunar") model = DQN.load("dqn_lunar")
# Evaluate the agent # Evaluate the agent
mean_reward, std_reward = evaluate_policy(model, model.get_env(), n_eval_episodes=10) mean_reward, std_reward = evaluate_policy(model, model.get_env(), n_eval_episodes=10)
@ -81,7 +81,7 @@ In the following example, we will train, save and load a A2C model on the Lunar
# Enjoy trained agent # Enjoy trained agent
obs = env.reset() obs = env.reset()
for i in range(1000): for i in range(1000):
action, _states = model.predict(obs) action, _states = model.predict(obs, deterministic=True)
obs, rewards, dones, info = env.step(action) obs, rewards, dones, info = env.step(action)
env.render() env.render()

View file

@ -58,6 +58,7 @@ Main Features
modules/ppo modules/ppo
modules/sac modules/sac
modules/td3 modules/td3
modules/dqn
.. toctree:: .. toctree::
:maxdepth: 1 :maxdepth: 1

View file

@ -3,15 +3,20 @@
Changelog Changelog
========== ==========
Pre-Release 0.8.0a0 (WIP) Pre-Release 0.8.0a1 (WIP)
------------------------------ ------------------------------
Breaking Changes: Breaking Changes:
^^^^^^^^^^^^^^^^^ ^^^^^^^^^^^^^^^^^
- ``AtariWrapper`` and other Atari wrappers were updated to match SB2 ones
- ``save_replay_buffer`` now receives as argument the file path instead of the folder path (@tirafesi) - ``save_replay_buffer`` now receives as argument the file path instead of the folder path (@tirafesi)
New Features: New Features:
^^^^^^^^^^^^^ ^^^^^^^^^^^^^
- Added ``DQN`` Algorithm (@Artemis-Skade)
- Buffer dtype is now set according to action and observation spaces for ``ReplayBuffer``
- Added warning when allocation of a buffer may exceed the available memory of the system
when ``psutil`` is available
Bug Fixes: Bug Fixes:
^^^^^^^^^^ ^^^^^^^^^^
@ -22,6 +27,10 @@ Deprecations:
Others: Others:
^^^^^^^ ^^^^^^^
- Refactored off-policy algorithm to share the same ``.learn()`` method
- Split the ``collect_rollout()`` method for off-policy algorithms
- Added ``_on_step()`` for off-policy base class
- Optimized replay buffer size by removing the need of ``next_observations`` numpy array
Documentation: Documentation:
^^^^^^^^^^^^^^ ^^^^^^^^^^^^^^
@ -29,6 +38,7 @@ Documentation:
- Fixed a typo in the section of Enjoy a Trained Agent, in RL Baselines3 Zoo README. (@blurLake) - Fixed a typo in the section of Enjoy a Trained Agent, in RL Baselines3 Zoo README. (@blurLake)
Pre-Release 0.7.0 (2020-06-10) Pre-Release 0.7.0 (2020-06-10)
------------------------------ ------------------------------

94
docs/modules/dqn.rst Normal file
View file

@ -0,0 +1,94 @@
.. _dqn:
.. automodule:: stable_baselines3.dqn
DQN
===
`Deep Q Network (DQN) <https://arxiv.org/abs/1312.5602>`_
.. rubric:: Available Policies
.. autosummary::
:nosignatures:
MlpPolicy
CnnPolicy
Notes
-----
- Original paper: https://arxiv.org/abs/1312.5602
- Further reference: https://www.nature.com/articles/nature14236
.. note::
This implementation provides only vanilla Deep Q-Learning and has no extensions such as Double-DQN, Dueling-DQN and Prioritized Experience Replay.
Can I use?
----------
- Recurrent policies: ❌
- Multi processing: ❌
- Gym spaces:
============= ====== ===========
Space Action Observation
============= ====== ===========
Discrete ✔ ✔
Box ❌ ✔
MultiDiscrete ❌ ✔
MultiBinary ❌ ✔
============= ====== ===========
Example
-------
.. code-block:: python
import gym
import numpy as np
from stable_baselines3 import DQN
from stable_baselines3.dqn import MlpPolicy
env = gym.make('Pendulum-v0')
model = DQN(MlpPolicy, env, verbose=1)
model.learn(total_timesteps=10000, log_interval=4)
model.save("dqn_pendulum")
del model # remove to demonstrate saving and loading
model = DQN.load("dqn_pendulum")
obs = env.reset()
while True:
action, _states = model.predict(obs, deterministic=True)
obs, reward, done, info = env.step(action)
env.render()
if done:
obs = env.reset()
Parameters
----------
.. autoclass:: DQN
:members:
:inherited-members:
.. _dqn_policies:
DQN Policies
-------------
.. autoclass:: MlpPolicy
:members:
:inherited-members:
.. autoclass:: CnnPolicy
:members:

View file

@ -27,6 +27,7 @@ per-file-ignores =
./stable_baselines3/__init__.py:F401 ./stable_baselines3/__init__.py:F401
./stable_baselines3/common/__init__.py:F401 ./stable_baselines3/common/__init__.py:F401
./stable_baselines3/a2c/__init__.py:F401 ./stable_baselines3/a2c/__init__.py:F401
./stable_baselines3/dqn/__init__.py:F401
./stable_baselines3/ppo/__init__.py:F401 ./stable_baselines3/ppo/__init__.py:F401
./stable_baselines3/sac/__init__.py:F401 ./stable_baselines3/sac/__init__.py:F401
./stable_baselines3/td3/__init__.py:F401 ./stable_baselines3/td3/__init__.py:F401

View file

@ -108,7 +108,9 @@ setup(name='stable_baselines3',
# For atari games, # For atari games,
'atari_py~=0.2.0', 'pillow', 'atari_py~=0.2.0', 'pillow',
# Tensorboard support # Tensorboard support
'tensorboard' 'tensorboard',
# Checking memory taken by replay buffer
'psutil'
] ]
}, },
description='Pytorch version of Stable Baselines, implementations of reinforcement learning algorithms.', description='Pytorch version of Stable Baselines, implementations of reinforcement learning algorithms.',

View file

@ -4,6 +4,7 @@ from stable_baselines3.a2c import A2C
from stable_baselines3.ppo import PPO from stable_baselines3.ppo import PPO
from stable_baselines3.sac import SAC from stable_baselines3.sac import SAC
from stable_baselines3.td3 import TD3 from stable_baselines3.td3 import TD3
from stable_baselines3.dqn import DQN
# Read version from file # Read version from file
version_file = os.path.join(os.path.dirname(__file__), 'version.txt') version_file = os.path.join(os.path.dirname(__file__), 'version.txt')

View file

@ -1,31 +1,213 @@
import gym import gym
from gym.wrappers import AtariPreprocessing from gym import spaces
import numpy as np import numpy as np
try:
import cv2 # pytype:disable=import-error
cv2.ocl.setUseOpenCL(False)
except ImportError:
cv2 = None
from stable_baselines3.common.type_aliases import GymStepReturn from stable_baselines3.common.type_aliases import GymStepReturn
class NoopResetEnv(gym.Wrapper):
def __init__(self, env: gym.Env, noop_max: int = 30):
"""
Sample initial states by taking random number of no-ops on reset.
No-op is assumed to be action 0.
:param env: (gym.Env) the environment to wrap
:param noop_max: (int) the maximum value of no-ops to run
"""
gym.Wrapper.__init__(self, env)
self.noop_max = noop_max
self.override_num_noops = None
self.noop_action = 0
assert env.unwrapped.get_action_meanings()[0] == 'NOOP'
def reset(self, **kwargs) -> np.ndarray:
self.env.reset(**kwargs)
if self.override_num_noops is not None:
noops = self.override_num_noops
else:
noops = self.unwrapped.np_random.randint(1, self.noop_max + 1)
assert noops > 0
obs = np.empty(0)
for _ in range(noops):
obs, _, done, _ = self.env.step(self.noop_action)
if done:
obs = self.env.reset(**kwargs)
return obs
class FireResetEnv(gym.Wrapper):
def __init__(self, env: gym.Env):
"""
Take action on reset for environments that are fixed until firing.
:param env: (gym.Env) the environment to wrap
"""
gym.Wrapper.__init__(self, env)
assert env.unwrapped.get_action_meanings()[1] == 'FIRE'
assert len(env.unwrapped.get_action_meanings()) >= 3
def reset(self, **kwargs) -> np.ndarray:
self.env.reset(**kwargs)
obs, _, done, _ = self.env.step(1)
if done:
self.env.reset(**kwargs)
obs, _, done, _ = self.env.step(2)
if done:
self.env.reset(**kwargs)
return obs
class EpisodicLifeEnv(gym.Wrapper):
def __init__(self, env: gym.Env):
"""
Make end-of-life == end-of-episode, but only reset on true game over.
Done by DeepMind for the DQN and co. since it helps value estimation.
:param env: (gym.Env) the environment to wrap
"""
gym.Wrapper.__init__(self, env)
self.lives = 0
self.was_real_done = True
def step(self, action: int) -> GymStepReturn:
obs, reward, done, info = self.env.step(action)
self.was_real_done = done
# check current lives, make loss of life terminal,
# then update lives to handle bonus lives
lives = self.env.unwrapped.ale.lives()
if 0 < lives < self.lives:
# for Qbert sometimes we stay in lives == 0 condtion for a few frames
# so its important to keep lives > 0, so that we only reset once
# the environment advertises done.
done = True
self.lives = lives
return obs, reward, done, info
def reset(self, **kwargs) -> np.ndarray:
"""
Calls the Gym environment reset, only when lives are exhausted.
This way all states are still reachable even though lives are episodic,
and the learner need not know about any of this behind-the-scenes.
:param kwargs: Extra keywords passed to env.reset() call
:return: (np.ndarray) the first observation of the environment
"""
if self.was_real_done:
obs = self.env.reset(**kwargs)
else:
# no-op step to advance from terminal/lost life state
obs, _, _, _ = self.env.step(0)
self.lives = self.env.unwrapped.ale.lives()
return obs
class MaxAndSkipEnv(gym.Wrapper):
def __init__(self, env: gym.Env, skip: int = 4):
"""
Return only every ``skip``-th frame (frameskipping)
:param env: (gym.Env) the environment
:param skip: (int) number of ``skip``-th frame
"""
gym.Wrapper.__init__(self, env)
# most recent raw observations (for max pooling across time steps)
self._obs_buffer = np.zeros((2,) + env.observation_space.shape, dtype=env.observation_space.dtype)
self._skip = skip
def step(self, action: int) -> GymStepReturn:
"""
Step the environment with the given action
Repeat action, sum reward, and max over last observations.
:param action: ([int] or [float]) the action
:return: ([int] or [float], [float], [bool], dict) observation, reward, done, information
"""
total_reward = 0.0
done = None
for i in range(self._skip):
obs, reward, done, info = self.env.step(action)
if i == self._skip - 2:
self._obs_buffer[0] = obs
if i == self._skip - 1:
self._obs_buffer[1] = obs
total_reward += reward
if done:
break
# Note that the observation on the done=True frame
# doesn't matter
max_frame = self._obs_buffer.max(axis=0)
return max_frame, total_reward, done, info
def reset(self, **kwargs):
return self.env.reset(**kwargs)
class ClipRewardEnv(gym.RewardWrapper):
def __init__(self, env: gym.Env):
"""
Clips the reward to {+1, 0, -1} by its sign.
:param env: (gym.Env) the environment
"""
gym.RewardWrapper.__init__(self, env)
def reward(self, reward: float) -> float:
"""
Bin reward to {+1, 0, -1} by its sign.
:param reward: (float)
:return: (float)
"""
return np.sign(reward)
class WarpFrame(gym.ObservationWrapper):
def __init__(self, env: gym.Env, width: int = 84, height: int = 84):
"""
Convert to grayscale and warp frames to 84x84 (default)
as done in the Nature paper and later work.
:param env: (gym.Env) the environment
:param width: (int)
:param height: (int)
"""
gym.ObservationWrapper.__init__(self, env)
self.width = width
self.height = height
self.observation_space = spaces.Box(low=0, high=255, shape=(self.height, self.width, 1),
dtype=env.observation_space.dtype)
def observation(self, frame: np.ndarray) -> np.ndarray:
"""
returns the current observation from a frame
:param frame: (np.ndarray) environment frame
:return: (np.ndarray) the observation
"""
frame = cv2.cvtColor(frame, cv2.COLOR_RGB2GRAY)
frame = cv2.resize(frame, (self.width, self.height), interpolation=cv2.INTER_AREA)
return frame[:, :, None]
class AtariWrapper(gym.Wrapper): class AtariWrapper(gym.Wrapper):
""" """
Atari 2600 preprocessings Atari 2600 preprocessings
It is a wrapper around the one found in gym.
It reshapes the observation to have an additional dimension and clip the reward.
See https://github.com/openai/gym/blob/master/gym/wrappers/atari_preprocessing.py
.
This class follows the guidelines in
Machado et al. (2018), "Revisiting the Arcade Learning Environment:
Evaluation Protocols and Open Problems for General Agents".
Specifically: Specifically:
* NoopReset: obtain initial state by taking random number of no-ops on reset. * NoopReset: obtain initial state by taking random number of no-ops on reset.
* Frame skipping: 4 by default * Frame skipping: 4 by default
* Max-pooling: most recent two observations * Max-pooling: most recent two observations
* Termination signal when a life is lost: turned off by default. Not recommended by Machado et al. (2018). * Termination signal when a life is lost.
* Resize to a square image: 84x84 by default * Resize to a square image: 84x84 by default
* Grayscale observation: by default * Grayscale observation
* Scale observation: optional * Clip reward to {-1, 0, 1}
:param env: (gym.Env) gym environment :param env: (gym.Env) gym environment
:param noop_max: (int): max number of no-ops :param noop_max: (int): max number of no-ops
@ -33,44 +215,22 @@ class AtariWrapper(gym.Wrapper):
:param screen_size: (int): resize Atari frame :param screen_size: (int): resize Atari frame
:param terminal_on_life_loss: (bool): if True, then step() returns done=True whenever a :param terminal_on_life_loss: (bool): if True, then step() returns done=True whenever a
life is lost. life is lost.
:param grayscale_obs: (bool): if True (default), then gray scale observation is returned, otherwise, RGB observation :param clip_reward: (bool) If True (default), the reward is clip to {-1, 0, 1} depending on its sign.
is returned.
:param scale_obs: (bool): if True, then observation normalized in range [0,1] is returned. It also limits memory
optimization benefits of FrameStack Wrapper.
:param scale_obs: (bool) If True (default), the reward is clip to {-1, 0, 1} depending on its sign.
""" """
def __init__(self, env: gym.Env, def __init__(self, env: gym.Env,
noop_max: int = 30, noop_max: int = 30,
frame_skip: int = 4, frame_skip: int = 4,
screen_size: int = 84, screen_size: int = 84,
terminal_on_life_loss: bool = False, terminal_on_life_loss: bool = True,
grayscale_obs: bool = True,
scale_obs: bool = False,
clip_reward: bool = True): clip_reward: bool = True):
env = AtariPreprocessing(env, noop_max=noop_max, frame_skip=frame_skip, screen_size=screen_size, env = NoopResetEnv(env, noop_max=noop_max)
terminal_on_life_loss=terminal_on_life_loss, grayscale_obs=grayscale_obs, env = MaxAndSkipEnv(env, skip=frame_skip)
scale_obs=scale_obs) if terminal_on_life_loss:
# Add channel dimension env = EpisodicLifeEnv(env)
if grayscale_obs: if 'FIRE' in env.unwrapped.get_action_meanings():
obs_space = env.observation_space env = FireResetEnv(env)
_low, _high, _obs_dtype = (0, 255, np.uint8) if not scale_obs else (0, 1, np.float32) env = WarpFrame(env, width=screen_size, height=screen_size)
env.observation_space = gym.spaces.Box(low=_low, high=_high, shape=obs_space.shape + (1,), if clip_reward:
dtype=_obs_dtype) env = ClipRewardEnv(env)
super(AtariWrapper, self).__init__(env) super(AtariWrapper, self).__init__(env)
self.clip_reward = clip_reward
def _add_axis(self, obs: np.ndarray) -> np.ndarray:
if self.env.grayscale_obs:
return obs[..., np.newaxis]
return obs
def reset(self) -> np.ndarray:
return self._add_axis(self.env.reset())
def step(self, action: int) -> GymStepReturn:
obs, reward, done, info = self.env.step(action)
# Bin reward to {+1, 0, -1} by its sign.
if self.clip_reward:
reward = np.sign(reward)
return self._add_axis(obs), reward, done, info

View file

@ -84,6 +84,8 @@ class BaseAlgorithm(ABC):
self.action_space = None # type: Optional[gym.spaces.Space] self.action_space = None # type: Optional[gym.spaces.Space]
self.n_envs = None self.n_envs = None
self.num_timesteps = 0 self.num_timesteps = 0
# Used for updating schedules
self._total_timesteps = 0
self.eval_env = None self.eval_env = None
self.seed = seed self.seed = seed
self.action_noise = None # type: Optional[ActionNoise] self.action_noise = None # type: Optional[ActionNoise]
@ -398,7 +400,7 @@ class BaseAlgorithm(ABC):
log_path: Optional[str] = None, log_path: Optional[str] = None,
reset_num_timesteps: bool = True, reset_num_timesteps: bool = True,
tb_log_name: str = 'run', tb_log_name: str = 'run',
) -> Tuple[int, 'BaseCallback']: ) -> Tuple[int, BaseCallback]:
""" """
Initialize different variables needed for training. Initialize different variables needed for training.
@ -410,7 +412,7 @@ class BaseAlgorithm(ABC):
:param log_path (Optional[str]): Path to a log folder :param log_path (Optional[str]): Path to a log folder
:param reset_num_timesteps: (bool) Whether to reset or not the ``num_timesteps`` attribute :param reset_num_timesteps: (bool) Whether to reset or not the ``num_timesteps`` attribute
:param tb_log_name: (str) the name of the run for tensorboard log :param tb_log_name: (str) the name of the run for tensorboard log
:return: (int, Tuple[BaseCallback]) :return: (Tuple[int, BaseCallback])
""" """
self.start_time = time.time() self.start_time = time.time()
self.ep_info_buffer = deque(maxlen=100) self.ep_info_buffer = deque(maxlen=100)
@ -425,6 +427,7 @@ class BaseAlgorithm(ABC):
else: else:
# Make sure training timesteps are ahead of the internal counter # Make sure training timesteps are ahead of the internal counter
total_timesteps += self.num_timesteps total_timesteps += self.num_timesteps
self._total_timesteps = total_timesteps
# Avoid resetting the environment when calling ``.learn()`` consecutive times # Avoid resetting the environment when calling ``.learn()`` consecutive times
if reset_num_timesteps or self._last_obs is None: if reset_num_timesteps or self._last_obs is None:

View file

@ -1,9 +1,16 @@
from typing import Union, Optional, Generator from typing import Union, Optional, Generator
import warnings
import numpy as np import numpy as np
import torch as th import torch as th
from gym import spaces from gym import spaces
try:
# Check memory used by replay buffer when possible
import psutil
except ImportError:
psutil = None
from stable_baselines3.common.vec_env import VecNormalize from stable_baselines3.common.vec_env import VecNormalize
from stable_baselines3.common.type_aliases import RolloutBufferSamples, ReplayBufferSamples from stable_baselines3.common.type_aliases import RolloutBufferSamples, ReplayBufferSamples
from stable_baselines3.common.preprocessing import get_action_dim, get_obs_shape from stable_baselines3.common.preprocessing import get_action_dim, get_obs_shape
@ -145,25 +152,52 @@ class ReplayBuffer(BaseBuffer):
:param action_space: (spaces.Space) Action space :param action_space: (spaces.Space) Action space
:param device: (th.device) :param device: (th.device)
:param n_envs: (int) Number of parallel environments :param n_envs: (int) Number of parallel environments
:param optimize_memory_usage: (bool) Enable a memory efficient variant
of the replay buffer which reduces by almost a factor two the memory used,
at a cost of more complexity.
See https://github.com/DLR-RM/stable-baselines3/issues/37#issuecomment-637501195
and https://github.com/DLR-RM/stable-baselines3/pull/28#issuecomment-637559274
""" """
def __init__(self, def __init__(self,
buffer_size: int, buffer_size: int,
observation_space: spaces.Space, observation_space: spaces.Space,
action_space: spaces.Space, action_space: spaces.Space,
device: Union[th.device, str] = 'cpu', device: Union[th.device, str] = 'cpu',
n_envs: int = 1): n_envs: int = 1,
optimize_memory_usage: bool = False):
super(ReplayBuffer, self).__init__(buffer_size, observation_space, super(ReplayBuffer, self).__init__(buffer_size, observation_space,
action_space, device, n_envs=n_envs) action_space, device, n_envs=n_envs)
assert n_envs == 1, "Replay buffer only support single environment for now" assert n_envs == 1, "Replay buffer only support single environment for now"
self.observations = np.zeros((self.buffer_size, self.n_envs,) + self.obs_shape, dtype=np.float32) # Check that the replay buffer can fit into the memory
self.actions = np.zeros((self.buffer_size, self.n_envs, self.action_dim), dtype=np.float32) if psutil is not None:
self.next_observations = np.zeros((self.buffer_size, self.n_envs,) + self.obs_shape, dtype=np.float32) mem_available = psutil.virtual_memory().available
self.optimize_memory_usage = optimize_memory_usage
self.observations = np.zeros((self.buffer_size, self.n_envs,) + self.obs_shape, dtype=observation_space.dtype)
if optimize_memory_usage:
# `observations` contains also the next observation
self.next_observations = None
else:
self.next_observations = np.zeros((self.buffer_size, self.n_envs,) + self.obs_shape, dtype=observation_space.dtype)
self.actions = np.zeros((self.buffer_size, self.n_envs, self.action_dim), dtype=action_space.dtype)
self.rewards = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32) self.rewards = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32)
self.dones = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32) self.dones = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32)
if psutil is not None:
total_memory_usage = (self.observations.nbytes + self.actions.nbytes
+ self.rewards.nbytes + self.dones.nbytes)
if self.next_observations is not None:
total_memory_usage += self.next_observations.nbytes
if total_memory_usage > mem_available:
# Convert to GB
total_memory_usage /= 1e9
mem_available /= 1e9
warnings.warn("This system does not have apparently enough memory to store the complete "
f"replay buffer {total_memory_usage:.2f}GB > {mem_available:.2f}GB")
def add(self, def add(self,
obs: np.ndarray, obs: np.ndarray,
next_obs: np.ndarray, next_obs: np.ndarray,
@ -172,7 +206,11 @@ class ReplayBuffer(BaseBuffer):
done: np.ndarray) -> None: done: np.ndarray) -> None:
# Copy to avoid modification by reference # Copy to avoid modification by reference
self.observations[self.pos] = np.array(obs).copy() self.observations[self.pos] = np.array(obs).copy()
self.next_observations[self.pos] = np.array(next_obs).copy() if self.optimize_memory_usage:
self.observations[(self.pos + 1) % self.buffer_size] = np.array(next_obs).copy()
else:
self.next_observations[self.pos] = np.array(next_obs).copy()
self.actions[self.pos] = np.array(action).copy() self.actions[self.pos] = np.array(action).copy()
self.rewards[self.pos] = np.array(reward).copy() self.rewards[self.pos] = np.array(reward).copy()
self.dones[self.pos] = np.array(done).copy() self.dones[self.pos] = np.array(done).copy()
@ -182,13 +220,43 @@ class ReplayBuffer(BaseBuffer):
self.full = True self.full = True
self.pos = 0 self.pos = 0
def sample(self,
batch_size: int,
env: Optional[VecNormalize] = None
) -> ReplayBufferSamples:
"""
Sample elements from the replay buffer.
Custom sampling when using memory efficient variant,
as we should not sample the element with index `self.pos`
See https://github.com/DLR-RM/stable-baselines3/pull/28#issuecomment-637559274
:param batch_size: (int) Number of element to sample
:param env: (Optional[VecNormalize]) associated gym VecEnv
to normalize the observations/rewards when sampling
:return: (Union[RolloutBufferSamples, ReplayBufferSamples])
"""
if not self.optimize_memory_usage:
return super().sample(batch_size=batch_size, env=env)
# Do not sample the element with index `self.pos` as the transitions is invalid
# (we use only one array to store `obs` and `next_obs`)
if self.full:
batch_inds = (np.random.randint(1, self.buffer_size, size=batch_size) + self.pos) % self.buffer_size
else:
batch_inds = np.random.randint(0, self.pos, size=batch_size)
return self._get_samples(batch_inds, env=env)
def _get_samples(self, def _get_samples(self,
batch_inds: np.ndarray, batch_inds: np.ndarray,
env: Optional[VecNormalize] = None env: Optional[VecNormalize] = None
) -> ReplayBufferSamples: ) -> ReplayBufferSamples:
if self.optimize_memory_usage:
next_obs = self._normalize_obs(self.observations[(batch_inds + 1) % self.buffer_size, 0, :], env)
else:
next_obs = self._normalize_obs(self.next_observations[batch_inds, 0, :], env)
data = (self._normalize_obs(self.observations[batch_inds, 0, :], env), data = (self._normalize_obs(self.observations[batch_inds, 0, :], env),
self.actions[batch_inds, 0, :], self.actions[batch_inds, 0, :],
self._normalize_obs(self.next_observations[batch_inds, 0, :], env), next_obs,
self.dones[batch_inds], self.dones[batch_inds],
self._normalize_reward(self.rewards[batch_inds], env)) self._normalize_reward(self.rewards[batch_inds], env))
return ReplayBufferSamples(*tuple(map(self.to_torch, data))) return ReplayBufferSamples(*tuple(map(self.to_torch, data)))
@ -227,7 +295,8 @@ class RolloutBuffer(BaseBuffer):
self.reset() self.reset()
def reset(self) -> None: def reset(self) -> None:
self.observations = np.zeros((self.buffer_size, self.n_envs,) + self.obs_shape, dtype=np.float32) self.observations = np.zeros((self.buffer_size, self.n_envs,) + self.obs_shape,
dtype=np.float32)
self.actions = np.zeros((self.buffer_size, self.n_envs, self.action_dim), dtype=np.float32) self.actions = np.zeros((self.buffer_size, self.n_envs, self.action_dim), dtype=np.float32)
self.rewards = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32) self.rewards = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32)
self.returns = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32) self.returns = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32)

View file

@ -94,8 +94,8 @@ def make_atari_env(env_id: Union[str, Type[gym.Env]],
in a Monitor wrapper to provide additional information about training. in a Monitor wrapper to provide additional information about training.
:param wrapper_kwargs: (Dict[str, Any]) Optional keyword argument to pass to the ``AtariWrapper`` :param wrapper_kwargs: (Dict[str, Any]) Optional keyword argument to pass to the ``AtariWrapper``
:param env_kwargs: (Dict[str, Any]) Optional keyword argument to pass to the env constructor :param env_kwargs: (Dict[str, Any]) Optional keyword argument to pass to the env constructor
:param vec_env_cls: (Type[VecEnv]) A custom `VecEnv` class constructor. Default: None. :param vec_env_cls: (Type[VecEnv]) A custom ``VecEnv`` class constructor. Default: None.
:param vec_env_kwargs: (Dict[str, Any]) Keyword arguments to pass to the `VecEnv` class constructor. :param vec_env_kwargs: (Dict[str, Any]) Keyword arguments to pass to the ``VecEnv`` class constructor.
:return: (VecEnv) The wrapped environment :return: (VecEnv) The wrapped environment
""" """
if wrapper_kwargs is None: if wrapper_kwargs is None:

View file

@ -1,8 +1,7 @@
import time import time
import os
import pickle import pickle
import warnings import warnings
from typing import Union, Type, Optional, Dict, Any, Callable from typing import Union, Type, Optional, Dict, Any, Callable, List, Tuple
import gym import gym
import torch as th import torch as th
@ -13,7 +12,7 @@ from stable_baselines3.common.base_class import BaseAlgorithm
from stable_baselines3.common.policies import BasePolicy from stable_baselines3.common.policies import BasePolicy
from stable_baselines3.common.utils import safe_mean from stable_baselines3.common.utils import safe_mean
from stable_baselines3.common.vec_env import VecEnv from stable_baselines3.common.vec_env import VecEnv
from stable_baselines3.common.type_aliases import GymEnv, RolloutReturn from stable_baselines3.common.type_aliases import GymEnv, RolloutReturn, MaybeCallback
from stable_baselines3.common.callbacks import BaseCallback from stable_baselines3.common.callbacks import BaseCallback
from stable_baselines3.common.noise import ActionNoise from stable_baselines3.common.noise import ActionNoise
from stable_baselines3.common.buffers import ReplayBuffer from stable_baselines3.common.buffers import ReplayBuffer
@ -32,6 +31,17 @@ class OffPolicyAlgorithm(BaseAlgorithm):
:param buffer_size: (int) size of the replay buffer :param buffer_size: (int) size of the replay buffer
:param learning_starts: (int) how many steps of the model to collect transitions for before learning starts :param learning_starts: (int) how many steps of the model to collect transitions for before learning starts
:param batch_size: (int) Minibatch size for each gradient update :param batch_size: (int) Minibatch size for each gradient update
:param tau: (float) the soft update coefficient ("Polyak update", between 0 and 1)
:param gamma: (float) the discount factor
:param train_freq: (int) Update the model every ``train_freq`` steps.
:param gradient_steps: (int) How many gradient update after each step
:param n_episodes_rollout: (int) Update the model every ``n_episodes_rollout`` episodes.
Note that this cannot be used at the same time as ``train_freq``
:param action_noise: (ActionNoise) the action noise type (None by default), this can help
for hard exploration problem. Cf common.noise for the different action noise type.
:param optimize_memory_usage: (bool) Enable a memory efficient variant of the replay buffer
at a cost of more complexity.
See https://github.com/DLR-RM/stable-baselines3/issues/37#issuecomment-637501195
:param policy_kwargs: Additional arguments to be passed to the policy on creation :param policy_kwargs: Additional arguments to be passed to the policy on creation
:param tensorboard_log: (str) the log location for tensorboard (if None, no logging) :param tensorboard_log: (str) the log location for tensorboard (if None, no logging)
:param verbose: The verbosity level: 0 none, 1 training information, 2 debug :param verbose: The verbosity level: 0 none, 1 training information, 2 debug
@ -62,6 +72,13 @@ class OffPolicyAlgorithm(BaseAlgorithm):
buffer_size: int = int(1e6), buffer_size: int = int(1e6),
learning_starts: int = 100, learning_starts: int = 100,
batch_size: int = 256, batch_size: int = 256,
tau: float = 0.005,
gamma: float = 0.99,
train_freq: int = 1,
gradient_steps: int = 1,
n_episodes_rollout: int = -1,
action_noise: Optional[ActionNoise] = None,
optimize_memory_usage: bool = False,
policy_kwargs: Dict[str, Any] = None, policy_kwargs: Dict[str, Any] = None,
tensorboard_log: Optional[str] = None, tensorboard_log: Optional[str] = None,
verbose: int = 0, verbose: int = 0,
@ -84,6 +101,22 @@ class OffPolicyAlgorithm(BaseAlgorithm):
self.buffer_size = buffer_size self.buffer_size = buffer_size
self.batch_size = batch_size self.batch_size = batch_size
self.learning_starts = learning_starts self.learning_starts = learning_starts
self.tau = tau
self.gamma = gamma
self.train_freq = train_freq
self.gradient_steps = gradient_steps
self.n_episodes_rollout = n_episodes_rollout
self.action_noise = action_noise
self.optimize_memory_usage = optimize_memory_usage
if train_freq > 0 and n_episodes_rollout > 0:
warnings.warn("You passed a positive value for `train_freq` and `n_episodes_rollout`."
"Please make sure this is intended. "
"The agent will collect data by stepping in the environment "
"until both conditions are true: "
"`number of steps in the env` >= `train_freq` and "
"`number of episodes` > `n_episodes_rollout`")
self.actor = None # type: Optional[th.nn.Module] self.actor = None # type: Optional[th.nn.Module]
self.replay_buffer = None # type: Optional[ReplayBuffer] self.replay_buffer = None # type: Optional[ReplayBuffer]
# Update policy keyword arguments # Update policy keyword arguments
@ -97,7 +130,8 @@ class OffPolicyAlgorithm(BaseAlgorithm):
self._setup_lr_schedule() self._setup_lr_schedule()
self.set_random_seed(self.seed) self.set_random_seed(self.seed)
self.replay_buffer = ReplayBuffer(self.buffer_size, self.observation_space, self.replay_buffer = ReplayBuffer(self.buffer_size, self.observation_space,
self.action_space, self.device) self.action_space, self.device,
optimize_memory_usage=self.optimize_memory_usage)
self.policy = self.policy_class(self.observation_space, self.action_space, self.policy = self.policy_class(self.observation_space, self.action_space,
self.lr_schedule, **self.policy_kwargs) self.lr_schedule, **self.policy_kwargs)
self.policy = self.policy.to(self.device) self.policy = self.policy.to(self.device)
@ -122,10 +156,158 @@ class OffPolicyAlgorithm(BaseAlgorithm):
self.replay_buffer = pickle.load(file_handler) self.replay_buffer = pickle.load(file_handler)
assert isinstance(self.replay_buffer, ReplayBuffer), 'The replay buffer must inherit from ReplayBuffer class' assert isinstance(self.replay_buffer, ReplayBuffer), 'The replay buffer must inherit from ReplayBuffer class'
def collect_rollouts(self, # noqa: C901 def _setup_learn(self,
total_timesteps: int,
eval_env: Optional[GymEnv],
callback: Union[None, Callable, List[BaseCallback], BaseCallback] = None,
eval_freq: int = 10000,
n_eval_episodes: int = 5,
log_path: Optional[str] = None,
reset_num_timesteps: bool = True,
tb_log_name: str = 'run',
) -> Tuple[int, BaseCallback]:
"""
cf `BaseAlgorithm`.
"""
# Prevent continuity issue by truncating trajectory
# when using memory efficient replay buffer
# see https://github.com/DLR-RM/stable-baselines3/issues/46
truncate_last_traj = (self.optimize_memory_usage and reset_num_timesteps
and self.replay_buffer is not None
and (self.replay_buffer.full or self.replay_buffer.pos > 0))
if truncate_last_traj:
warnings.warn("The last trajectory in the replay buffer will be truncated, "
"see https://github.com/DLR-RM/stable-baselines3/issues/46."
"You should use `reset_num_timesteps=False` or `optimize_memory_usage=False`"
"to avoid that issue.")
# Go to the previous index
pos = (self.replay_buffer.pos - 1) % self.replay_buffer.buffer_size
self.replay_buffer.dones[pos] = True
return super()._setup_learn(total_timesteps, eval_env, callback, eval_freq,
n_eval_episodes, log_path, reset_num_timesteps, tb_log_name)
def learn(self,
total_timesteps: int,
callback: MaybeCallback = None,
log_interval: int = 4,
eval_env: Optional[GymEnv] = None,
eval_freq: int = -1,
n_eval_episodes: int = 5,
tb_log_name: str = "run",
eval_log_path: Optional[str] = None,
reset_num_timesteps: bool = True) -> 'OffPolicyAlgorithm':
total_timesteps, callback = self._setup_learn(total_timesteps, eval_env, callback, eval_freq,
n_eval_episodes, eval_log_path, reset_num_timesteps,
tb_log_name)
callback.on_training_start(locals(), globals())
while self.num_timesteps < total_timesteps:
rollout = self.collect_rollouts(self.env, n_episodes=self.n_episodes_rollout,
n_steps=self.train_freq, action_noise=self.action_noise,
callback=callback,
learning_starts=self.learning_starts,
replay_buffer=self.replay_buffer,
log_interval=log_interval)
if rollout.continue_training is False:
break
if self.num_timesteps > 0 and self.num_timesteps > self.learning_starts:
# If no `gradient_steps` is specified,
# do as many gradients steps as steps performed during the rollout
gradient_steps = self.gradient_steps if self.gradient_steps > 0 else rollout.episode_timesteps
self.train(batch_size=self.batch_size, gradient_steps=gradient_steps)
callback.on_training_end()
return self
def train(self, gradient_steps: int, batch_size: int) -> None:
"""
Sample the replay buffer and do the updates
(gradient descent and update target networks)
"""
raise NotImplementedError()
def _sample_action(self, learning_starts: int,
action_noise: Optional[ActionNoise] = None) -> Tuple[np.ndarray, np.ndarray]:
"""
Sample an action according to the exploration policy.
This is either done by sampling the probability distribution of the policy,
or sampling a random action (from a uniform distribution over the action space)
or by adding noise to the deterministic output.
:param action_noise: (Optional[ActionNoise]) Action noise that will be used for exploration
Required for deterministic policy (e.g. TD3). This can also be used
in addition to the stochastic policy for SAC.
:param learning_starts: (int) Number of steps before learning for the warm-up phase.
:return: (Tuple[np.ndarray, np.ndarray]) action to take in the environment
and scaled action that will be stored in the replay buffer.
The two differs when the action space is not normalized (bounds are not [-1, 1]).
"""
# Select action randomly or according to policy
if self.num_timesteps < learning_starts and not (self.use_sde and self.use_sde_at_warmup):
# Warmup phase
unscaled_action = np.array([self.action_space.sample()])
else:
# Note: when using continuous actions,
# we assume that the policy uses tanh to scale the action
# We use non-deterministic action in the case of SAC, for TD3, it does not matter
unscaled_action, _ = self.predict(self._last_obs, deterministic=False)
# Rescale the action from [low, high] to [-1, 1]
if isinstance(self.action_space, gym.spaces.Box):
scaled_action = self.policy.scale_action(unscaled_action)
# Add noise to the action (improve exploration)
if action_noise is not None:
scaled_action = np.clip(scaled_action + action_noise(), -1, 1)
# We store the scaled action in the buffer
buffer_action = scaled_action
action = self.policy.unscale_action(scaled_action)
else:
# Discrete case, no need to normalize or clip
buffer_action = unscaled_action
action = buffer_action
return action, buffer_action
def _dump_logs(self) -> None:
"""
Write log.
"""
fps = int(self.num_timesteps / (time.time() - self.start_time))
logger.record("time/episodes", self._episode_num, exclude="tensorboard")
if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0:
logger.record('rollout/ep_rew_mean', safe_mean([ep_info['r'] for ep_info in self.ep_info_buffer]))
logger.record('rollout/ep_len_mean', safe_mean([ep_info['l'] for ep_info in self.ep_info_buffer]))
logger.record("time/fps", fps)
logger.record('time/time_elapsed', int(time.time() - self.start_time), exclude="tensorboard")
logger.record("time/total timesteps", self.num_timesteps, exclude="tensorboard")
if self.use_sde:
logger.record("train/std", (self.actor.get_std()).mean().item())
if len(self.ep_success_buffer) > 0:
logger.record('rollout/success rate', safe_mean(self.ep_success_buffer))
# Pass the number of timesteps for tensorboard
logger.dump(step=self.num_timesteps)
def _on_step(self) -> None:
"""
Method called after each step in the environment.
It is meant to trigger DQN target network update
but can be used for other purposes
"""
pass
def collect_rollouts(self,
env: VecEnv, env: VecEnv,
# Type hint as string to avoid circular import callback: BaseCallback,
callback: 'BaseCallback',
n_episodes: int = 1, n_episodes: int = 1,
n_steps: int = -1, n_steps: int = -1,
action_noise: Optional[ActionNoise] = None, action_noise: Optional[ActionNoise] = None,
@ -156,17 +338,6 @@ class OffPolicyAlgorithm(BaseAlgorithm):
assert isinstance(env, VecEnv), "You must pass a VecEnv" assert isinstance(env, VecEnv), "You must pass a VecEnv"
assert env.num_envs == 1, "OffPolicyAlgorithm only support single environment" assert env.num_envs == 1, "OffPolicyAlgorithm only support single environment"
if n_episodes > 0 and n_steps > 0:
# Note we are refering to the constructor arguments
# that are named `train_freq` and `n_episodes_rollout`
# but correspond to `n_steps` and `n_episodes` here
warnings.warn("You passed a positive value for `train_freq` and `n_episodes_rollout`."
"Please make sure this is intended. "
"The agent will collect data by stepping in the environment "
"until both conditions are true: "
"`number of steps in the env` >= `train_freq` and "
"`number of episodes` > `n_episodes_rollout`")
if self.use_sde: if self.use_sde:
self.actor.reset_noise() self.actor.reset_noise()
@ -184,31 +355,7 @@ class OffPolicyAlgorithm(BaseAlgorithm):
self.actor.reset_noise() self.actor.reset_noise()
# Select action randomly or according to policy # Select action randomly or according to policy
if self.num_timesteps < learning_starts and not (self.use_sde and self.use_sde_at_warmup): action, buffer_action = self._sample_action(learning_starts, action_noise)
# Warmup phase
unscaled_action = np.array([self.action_space.sample()])
else:
# Note: we assume that the policy uses tanh to scale the action
# We use non-deterministic action in the case of SAC, for TD3, it does not matter
unscaled_action, _ = self.predict(self._last_obs, deterministic=False)
# Rescale the action from [low, high] to [-1, 1]
if isinstance(self.action_space, gym.spaces.Box):
scaled_action = self.policy.scale_action(unscaled_action)
# Add noise to the action (improve exploration)
if action_noise is not None:
# NOTE: in the original implementation of TD3, the noise was applied to the unscaled action
# Update(October 2019): Not anymore
scaled_action = np.clip(scaled_action + action_noise(), -1, 1)
# We store the scaled action in the buffer
buffer_action = scaled_action
action = self.policy.unscale_action(scaled_action)
else:
# Discrete case, no need to normalize or clip
buffer_action = unscaled_action
action = buffer_action
# Rescale and perform action # Rescale and perform action
new_obs, reward, done, infos = env.step(action) new_obs, reward, done, infos = env.step(action)
@ -242,6 +389,14 @@ class OffPolicyAlgorithm(BaseAlgorithm):
self.num_timesteps += 1 self.num_timesteps += 1
episode_timesteps += 1 episode_timesteps += 1
total_steps += 1 total_steps += 1
self._update_current_progress_remaining(self.num_timesteps, self._total_timesteps)
# For DQN, check if the target network should be updated
# and update the exploration schedule
# For SAC/TD3, the update is done as the same time as the gradient update
# see https://github.com/hill-a/stable-baselines/issues/900
self._on_step()
if 0 < n_steps <= total_steps: if 0 < n_steps <= total_steps:
break break
@ -256,21 +411,7 @@ class OffPolicyAlgorithm(BaseAlgorithm):
# Log training infos # Log training infos
if log_interval is not None and self._episode_num % log_interval == 0: if log_interval is not None and self._episode_num % log_interval == 0:
fps = int(self.num_timesteps / (time.time() - self.start_time)) self._dump_logs()
logger.record("time/episodes", self._episode_num, exclude="tensorboard")
if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0:
logger.record('rollout/ep_rew_mean', safe_mean([ep_info['r'] for ep_info in self.ep_info_buffer]))
logger.record('rollout/ep_len_mean', safe_mean([ep_info['l'] for ep_info in self.ep_info_buffer]))
logger.record("time/fps", fps)
logger.record('time/time_elapsed', int(time.time() - self.start_time), exclude="tensorboard")
logger.record("time/total timesteps", self.num_timesteps, exclude="tensorboard")
if self.use_sde:
logger.record("train/std", (self.actor.get_std()).mean().item())
if len(self.ep_success_buffer) > 0:
logger.record('rollout/success rate', safe_mean(self.ep_success_buffer))
# Pass the number of timesteps for tensorboard
logger.dump(step=self.num_timesteps)
mean_reward = np.mean(episode_rewards) if total_episodes > 0 else 0.0 mean_reward = np.mean(episode_rewards) if total_episodes > 0 else 0.0

View file

@ -89,6 +89,30 @@ def get_schedule_fn(value_schedule: Union[Callable, float]) -> Callable:
return value_schedule return value_schedule
def get_linear_fn(start: float, end: float, end_fraction: float) -> Callable:
"""
Create a function that interpolates linearly between start and end
between ``progress_remaining`` = 1 and ``progress_remaining`` = ``end_fraction``.
This is used in DQN for linearly annealing the exploration fraction
(epsilon for the epsilon-greedy strategy).
:params start: (float) value to start with if ``progress_remaining`` = 1
:params end: (float) value to end with if ``progress_remaining`` = 0
:params end_fraction: (float) fraction of ``progress_remaining``
where end is reached e.g 0.1 then end is reached after 10%
of the complete training process.
:return: (Callable)
"""
def func(progress_remaining: float) -> float:
if (1 - progress_remaining) > end_fraction:
return end
else:
return start + (1 - progress_remaining) * (end - start) / end_fraction
return func
def constant_fn(val: float) -> Callable: def constant_fn(val: float) -> Callable:
""" """
Create a function that returns a constant Create a function that returns a constant

View file

@ -0,0 +1,3 @@
from stable_baselines3.dqn.dqn import DQN
from stable_baselines3.dqn.policies import MlpPolicy
from stable_baselines3.dqn.policies import CnnPolicy

View file

@ -0,0 +1,219 @@
from typing import List, Tuple, Type, Union, Callable, Optional, Dict, Any
import numpy as np
import torch as th
import torch.nn.functional as F
from stable_baselines3.common import logger
from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm
from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback
from stable_baselines3.common.utils import get_linear_fn
from stable_baselines3.dqn.policies import DQNPolicy
class DQN(OffPolicyAlgorithm):
"""
Deep Q-Network (DQN)
Paper: https://arxiv.org/abs/1312.5602, https://www.nature.com/articles/nature14236
Default hyperparameters are taken from the nature paper,
except for the optimizer and learning rate that were taken from Stable Baselines defaults.
:param policy: (DQNPolicy or str) The policy model to use (MlpPolicy, CnnPolicy, ...)
:param env: (Gym environment or str) The environment to learn from (if registered in Gym, can be str)
:param learning_rate: (float or callable) The learning rate, it can be a function
of the current progress (from 1 to 0)
:param buffer_size: (int) size of the replay buffer
:param learning_starts: (int) how many steps of the model to collect transitions for before learning starts
:param batch_size: (int) Minibatch size for each gradient update
:param tau: (float) the soft update coefficient ("Polyak update", between 0 and 1) default 1 for hard update
:param gamma: (float) the discount factor
:param train_freq: (int) Update the model every ``train_freq`` steps.
:param gradient_steps: (int) How many gradient update after each step
:param n_episodes_rollout: (int) Update the model every ``n_episodes_rollout`` episodes.
Note that this cannot be used at the same time as ``train_freq``
:param optimize_memory_usage: (bool) Enable a memory efficient variant of the replay buffer
at a cost of more complexity.
See https://github.com/DLR-RM/stable-baselines3/issues/37#issuecomment-637501195
:param target_update_interval: (int) update the target network every ``target_update_interval``
environment steps.
:param exploration_fraction: (float) fraction of entire training period over which the exploration rate is reduced
:param exploration_initial_eps: (float) initial value of random action probability
:param exploration_final_eps: (float) final value of random action probability
:param max_grad_norm: (float) The maximum value for the gradient clipping
:param tensorboard_log: (str) the log location for tensorboard (if None, no logging)
:param create_eval_env: (bool) Whether to create a second environment that will be
used for evaluating the agent periodically. (Only available when passing string for the environment)
:param policy_kwargs: (dict) additional arguments to be passed to the policy on creation
:param verbose: (int) the verbosity level: 0 no output, 1 info, 2 debug
:param seed: (int) Seed for the pseudo random generators
:param device: (str or th.device) Device (cpu, cuda, ...) on which the code should be run.
Setting it to auto, the code will be run on the GPU if possible.
:param _init_setup_model: (bool) Whether or not to build the network at the creation of the instance
"""
def __init__(self, policy: Union[str, Type[DQNPolicy]],
env: Union[GymEnv, str],
learning_rate: Union[float, Callable] = 1e-4,
buffer_size: int = 1000000,
learning_starts: int = 50000,
batch_size: Optional[int] = 32,
tau: float = 1.0,
gamma: float = 0.99,
train_freq: int = 4,
gradient_steps: int = 1,
n_episodes_rollout: int = -1,
optimize_memory_usage: bool = False,
target_update_interval: int = 10000,
exploration_fraction: float = 0.1,
exploration_initial_eps: float = 1.0,
exploration_final_eps: float = 0.05,
max_grad_norm: float = 10,
tensorboard_log: Optional[str] = None,
create_eval_env: bool = False,
policy_kwargs: Optional[Dict[str, Any]] = None,
verbose: int = 0,
seed: Optional[int] = None,
device: Union[th.device, str] = 'auto',
_init_setup_model: bool = True):
super(DQN, self).__init__(policy, env, DQNPolicy, learning_rate,
buffer_size, learning_starts, batch_size,
tau, gamma, train_freq, gradient_steps,
n_episodes_rollout, action_noise=None, # No action noise
policy_kwargs=policy_kwargs,
tensorboard_log=tensorboard_log,
verbose=verbose, device=device,
create_eval_env=create_eval_env,
seed=seed, sde_support=False,
optimize_memory_usage=optimize_memory_usage)
self.exploration_initial_eps = exploration_initial_eps
self.exploration_final_eps = exploration_final_eps
self.exploration_fraction = exploration_fraction
self.target_update_interval = target_update_interval
self.max_grad_norm = max_grad_norm
# "epsilon" for the epsilon-greedy exploration
self.exploration_rate = 0.0
# Linear schedule will be defined in `_setup_model()`
self.exploration_schedule = None
self.q_net, self.q_net_target = None, None
if _init_setup_model:
self._setup_model()
def _setup_model(self) -> None:
super(DQN, self)._setup_model()
self._create_aliases()
self.exploration_schedule = get_linear_fn(self.exploration_initial_eps, self.exploration_final_eps,
self.exploration_fraction)
def _create_aliases(self) -> None:
self.q_net = self.policy.q_net
self.q_net_target = self.policy.q_net_target
def _on_step(self):
"""
Update the exploration rate and target network if needed.
This method is called in ``collect_rollout()`` after each step in the environment.
"""
if self.num_timesteps % self.target_update_interval == 0:
for param, target_param in zip(self.q_net.parameters(), self.q_net_target.parameters()):
target_param.data.copy_(self.tau * param.data + (1 - self.tau) * target_param.data)
self.exploration_rate = self.exploration_schedule(self._current_progress_remaining)
logger.record("rollout/exploration rate", self.exploration_rate)
def train(self, gradient_steps: int, batch_size: int = 100) -> None:
# Update learning rate according to schedule
self._update_learning_rate(self.policy.optimizer)
for gradient_step in range(gradient_steps):
# Sample replay buffer
replay_data = self.replay_buffer.sample(batch_size, env=self._vec_normalize_env)
with th.no_grad():
# Compute the target Q values
target_q = self.q_net_target(replay_data.next_observations)
# Follow greedy policy: use the one with the highest value
target_q, _ = target_q.max(dim=1)
# Avoid potential broadcast issue
target_q = target_q.reshape(-1, 1)
# 1-step TD target
target_q = replay_data.rewards + (1 - replay_data.dones) * self.gamma * target_q
# Get current Q estimates
current_q = self.q_net(replay_data.observations)
# Retrieve the q-values for the actions from the replay buffer
current_q = th.gather(current_q, dim=1, index=replay_data.actions.long())
# Compute Huber loss (less sensitive to outliers)
loss = F.smooth_l1_loss(current_q, target_q)
# Optimize the policy
self.policy.optimizer.zero_grad()
loss.backward()
# Clip gradient norm
th.nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
self.policy.optimizer.step()
# Increase update counter
self._n_updates += gradient_steps
logger.record("train/n_updates", self._n_updates, exclude='tensorboard')
def predict(self, observation: np.ndarray,
state: Optional[np.ndarray] = None,
mask: Optional[np.ndarray] = None,
deterministic: bool = False) -> Tuple[np.ndarray, Optional[np.ndarray]]:
"""
Overrides the base_class predict function to include epsilon-greedy exploration.
:param observation: (np.ndarray) the input observation
:param state: (Optional[np.ndarray]) The last states (can be None, used in recurrent policies)
:param mask: (Optional[np.ndarray]) The last masks (can be None, used in recurrent policies)
:param deterministic: (bool) Whether or not to return deterministic actions.
:return: (Tuple[np.ndarray, Optional[np.ndarray]]) the model's action and the next state
(used in recurrent policies)
"""
if not deterministic and np.random.rand() < self.exploration_rate:
n_batch = observation.shape[0]
action = np.array([self.action_space.sample() for _ in range(n_batch)])
else:
action, state = self.policy.predict(observation, state, mask, deterministic)
return action, state
def learn(self,
total_timesteps: int,
callback: MaybeCallback = None,
log_interval: int = 4,
eval_env: Optional[GymEnv] = None,
eval_freq: int = -1,
n_eval_episodes: int = 5,
tb_log_name: str = "DQN",
eval_log_path: Optional[str] = None,
reset_num_timesteps: bool = True) -> OffPolicyAlgorithm:
return super(DQN, self).learn(total_timesteps=total_timesteps, callback=callback, log_interval=log_interval,
eval_env=eval_env, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes,
tb_log_name=tb_log_name, eval_log_path=eval_log_path,
reset_num_timesteps=reset_num_timesteps)
def excluded_save_params(self) -> List[str]:
"""
Returns the names of the parameters that should be excluded by default
when saving the model.
:return: (List[str]) List of parameters that should be excluded from save
"""
# Exclude aliases
return super(DQN, self).excluded_save_params() + ["q_net", "q_net_target"]
def get_torch_variables(self) -> Tuple[List[str], List[str]]:
"""
cf base class
"""
state_dicts = ["policy", "policy.optimizer"]
return state_dicts, []

View file

@ -0,0 +1,229 @@
from typing import Optional, List, Callable, Union, Type, Any, Dict
import gym
import torch as th
import torch.nn as nn
from stable_baselines3.common.policies import BasePolicy, register_policy
from stable_baselines3.common.torch_layers import create_mlp, NatureCNN, BaseFeaturesExtractor, FlattenExtractor
class QNetwork(BasePolicy):
"""
Action-Value (Q-Value) network for DQN
:param observation_space: (gym.spaces.Space) Observation space
:param action_space: (gym.spaces.Space) Action space
:param net_arch: (Optional[List[int]]) The specification of the policy and value networks.
:param device: (str or th.device) Device on which the code should run.
:param activation_fn: (Type[nn.Module]) Activation function
:param normalize_images: (bool) Whether to normalize images or not,
dividing by 255.0 (True by default)
"""
def __init__(self, observation_space: gym.spaces.Space,
action_space: gym.spaces.Space,
features_extractor: nn.Module,
features_dim: int,
net_arch: Optional[List[int]] = None,
device: Union[th.device, str] = 'auto',
activation_fn: Type[nn.Module] = nn.ReLU,
normalize_images: bool = True):
super(QNetwork, self).__init__(observation_space, action_space,
features_extractor=features_extractor,
normalize_images=normalize_images,
device=device)
if net_arch is None:
net_arch = [64, 64]
self.net_arch = net_arch
self.activation_fn = activation_fn
self.features_extractor = features_extractor
self.features_dim = features_dim
self.normalize_images = normalize_images
action_dim = self.action_space.n # number of actions
q_net = create_mlp(self.features_dim, action_dim, self.net_arch, self.activation_fn)
self.q_net = nn.Sequential(*q_net)
def forward(self, obs: th.Tensor) -> th.Tensor:
"""
Predict the q-values.
:param obs: (th.Tensor) Observation
:return: (th.Tensor) The estimated Q-Value for each action.
"""
return self.q_net(self.extract_features(obs))
def _predict(self, observation: th.Tensor, deterministic: bool = True) -> th.Tensor:
q_values = self.forward(observation)
# Greedy action
action = q_values.argmax(dim=1).reshape(-1)
return action
def _get_data(self) -> Dict[str, Any]:
data = super()._get_data()
data.update(dict(
net_arch=self.net_arch,
features_dim=self.features_dim,
activation_fn=self.activation_fn,
features_extractor=self.features_extractor,
epsilon=self.epsilon,
))
return data
class DQNPolicy(BasePolicy):
"""
Policy class with Q-Value Net and target net for DQN
:param observation_space: (gym.spaces.Space) Observation space
:param action_space: (gym.spaces.Space) Action space
:param lr_schedule: (callable) Learning rate schedule (could be constant)
:param net_arch: (Optional[List[int]]) The specification of the policy and value networks.
:param device: (str or th.device) Device on which the code should run.
:param activation_fn: (Type[nn.Module]) Activation function
:param features_extractor_class: (Type[BaseFeaturesExtractor]) Features extractor to use.
:param features_extractor_kwargs: (Optional[Dict[str, Any]]) Keyword arguments
to pass to the feature extractor.
:param normalize_images: (bool) Whether to normalize images or not,
dividing by 255.0 (True by default)
:param optimizer_class: (Type[th.optim.Optimizer]) The optimizer to use,
``th.optim.Adam`` by default
:param optimizer_kwargs: (Optional[Dict[str, Any]]) Additional keyword arguments,
excluding the learning rate, to pass to the optimizer
"""
def __init__(self, observation_space: gym.spaces.Space,
action_space: gym.spaces.Space,
lr_schedule: Callable,
net_arch: Optional[List[int]] = None,
device: Union[th.device, str] = 'auto',
activation_fn: Type[nn.Module] = nn.ReLU,
features_extractor_class: Type[BaseFeaturesExtractor] = FlattenExtractor,
features_extractor_kwargs: Optional[Dict[str, Any]] = None,
normalize_images: bool = True,
optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam,
optimizer_kwargs: Optional[Dict[str, Any]] = None):
super(DQNPolicy, self).__init__(observation_space, action_space,
device,
features_extractor_class,
features_extractor_kwargs,
optimizer_class=optimizer_class,
optimizer_kwargs=optimizer_kwargs)
if net_arch is None:
if features_extractor_class == FlattenExtractor:
net_arch = [64, 64]
else:
net_arch = []
self.features_extractor = features_extractor_class(self.observation_space,
**self.features_extractor_kwargs)
self.features_dim = self.features_extractor.features_dim
self.net_arch = net_arch
self.activation_fn = activation_fn
self.normalize_images = normalize_images
self.net_args = {
'observation_space': self.observation_space,
'action_space': self.action_space,
'features_extractor': self.features_extractor,
'features_dim': self.features_dim,
'net_arch': self.net_arch,
'activation_fn': self.activation_fn,
'normalize_images': normalize_images,
'device': device
}
self.q_net, self.q_net_target = None, None
self._build(lr_schedule)
def _build(self, lr_schedule: Callable) -> None:
"""
Create the network and the optimizer.
:param lr_schedule: (Callable) Learning rate schedule
lr_schedule(1) is the initial learning rate
"""
self.q_net = self.make_q_net()
self.q_net_target = self.make_q_net()
self.q_net_target.load_state_dict(self.q_net.state_dict())
# Setup optimizer with initial learning rate
self.optimizer = self.optimizer_class(self.parameters(), lr=lr_schedule(1),
**self.optimizer_kwargs)
def make_q_net(self) -> QNetwork:
return QNetwork(**self.net_args).to(self.device)
def forward(self, obs: th.Tensor, deterministic: bool = True) -> th.Tensor:
return self._predict(obs, deterministic=deterministic)
def _predict(self, obs: th.Tensor, deterministic: bool = True) -> th.Tensor:
return self.q_net._predict(obs, deterministic=deterministic)
def _get_data(self) -> Dict[str, Any]:
data = super()._get_data()
data.update(dict(
net_arch=self.net_args['net_arch'],
activation_fn=self.net_args['activation_fn'],
lr_schedule=self._dummy_schedule, # dummy lr schedule, not needed for loading policy alone
optimizer_class=self.optimizer_class,
optimizer_kwargs=self.optimizer_kwargs,
features_extractor_class=self.features_extractor_class,
features_extractor_kwargs=self.features_extractor_kwargs
))
return data
MlpPolicy = DQNPolicy
class CnnPolicy(DQNPolicy):
"""
Policy class for DQN when using images as input.
:param observation_space: (gym.spaces.Space) Observation space
:param action_space: (gym.spaces.Space) Action space
:param lr_schedule: (callable) Learning rate schedule (could be constant)
:param net_arch: (Optional[List[int]]) The specification of the policy and value networks.
:param device: (str or th.device) Device on which the code should run.
:param activation_fn: (Type[nn.Module]) Activation function
:param features_extractor_class: (Type[BaseFeaturesExtractor]) Features extractor to use.
:param normalize_images: (bool) Whether to normalize images or not,
dividing by 255.0 (True by default)
:param optimizer_class: (Type[th.optim.Optimizer]) The optimizer to use,
``th.optim.Adam`` by default
:param optimizer_kwargs: (Optional[Dict[str, Any]]) Additional keyword arguments,
excluding the learning rate, to pass to the optimizer
"""
def __init__(self, observation_space: gym.spaces.Space,
action_space: gym.spaces.Space,
lr_schedule: Callable,
net_arch: Optional[List[int]] = None,
device: Union[th.device, str] = 'auto',
activation_fn: Type[nn.Module] = nn.ReLU,
features_extractor_class: Type[BaseFeaturesExtractor] = NatureCNN,
features_extractor_kwargs: Optional[Dict[str, Any]] = None,
normalize_images: bool = True,
optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam,
optimizer_kwargs: Optional[Dict[str, Any]] = None):
super(CnnPolicy, self).__init__(observation_space,
action_space,
lr_schedule,
net_arch,
device,
activation_fn,
features_extractor_class,
features_extractor_kwargs,
normalize_images,
optimizer_class,
optimizer_kwargs)
register_policy("MlpPolicy", MlpPolicy)
register_policy("CnnPolicy", CnnPolicy)

View file

@ -40,10 +40,14 @@ class SAC(OffPolicyAlgorithm):
Note that this cannot be used at the same time as ``train_freq`` Note that this cannot be used at the same time as ``train_freq``
:param action_noise: (ActionNoise) the action noise type (None by default), this can help :param action_noise: (ActionNoise) the action noise type (None by default), this can help
for hard exploration problem. Cf common.noise for the different action noise type. for hard exploration problem. Cf common.noise for the different action noise type.
:param optimize_memory_usage: (bool) Enable a memory efficient variant of the replay buffer
at a cost of more complexity.
See https://github.com/DLR-RM/stable-baselines3/issues/37#issuecomment-637501195
:param ent_coef: (str or float) Entropy regularization coefficient. (Equivalent to :param ent_coef: (str or float) Entropy regularization coefficient. (Equivalent to
inverse of reward scale in the original SAC paper.) Controlling exploration/exploitation trade-off. inverse of reward scale in the original SAC paper.) Controlling exploration/exploitation trade-off.
Set it to 'auto' to learn it automatically (and 'auto_0.1' for using 0.1 as initial value) Set it to 'auto' to learn it automatically (and 'auto_0.1' for using 0.1 as initial value)
:param target_update_interval: (int) update the target network every ``target_network_update_freq`` steps. :param target_update_interval: (int) update the target network every ``target_network_update_freq``
gradient steps.
:param target_entropy: (str or float) target entropy when learning ``ent_coef`` (``ent_coef = 'auto'``) :param target_entropy: (str or float) target entropy when learning ``ent_coef`` (``ent_coef = 'auto'``)
:param use_sde: (bool) Whether to use generalized State Dependent Exploration (gSDE) :param use_sde: (bool) Whether to use generalized State Dependent Exploration (gSDE)
instead of action noise exploration (default: False) instead of action noise exploration (default: False)
@ -73,6 +77,7 @@ class SAC(OffPolicyAlgorithm):
gradient_steps: int = 1, gradient_steps: int = 1,
n_episodes_rollout: int = -1, n_episodes_rollout: int = -1,
action_noise: Optional[ActionNoise] = None, action_noise: Optional[ActionNoise] = None,
optimize_memory_usage: bool = False,
ent_coef: Union[str, float] = 'auto', ent_coef: Union[str, float] = 'auto',
target_update_interval: int = 1, target_update_interval: int = 1,
target_entropy: Union[str, float] = 'auto', target_entropy: Union[str, float] = 'auto',
@ -89,24 +94,22 @@ class SAC(OffPolicyAlgorithm):
super(SAC, self).__init__(policy, env, SACPolicy, learning_rate, super(SAC, self).__init__(policy, env, SACPolicy, learning_rate,
buffer_size, learning_starts, batch_size, buffer_size, learning_starts, batch_size,
policy_kwargs, tensorboard_log, verbose, device, tau, gamma, train_freq, gradient_steps,
n_episodes_rollout, action_noise,
policy_kwargs=policy_kwargs,
tensorboard_log=tensorboard_log,
verbose=verbose, device=device,
create_eval_env=create_eval_env, seed=seed, create_eval_env=create_eval_env, seed=seed,
use_sde=use_sde, sde_sample_freq=sde_sample_freq, use_sde=use_sde, sde_sample_freq=sde_sample_freq,
use_sde_at_warmup=use_sde_at_warmup) use_sde_at_warmup=use_sde_at_warmup,
optimize_memory_usage=optimize_memory_usage)
self.target_entropy = target_entropy self.target_entropy = target_entropy
self.log_ent_coef = None # type: Optional[th.Tensor] self.log_ent_coef = None # type: Optional[th.Tensor]
self.target_update_interval = target_update_interval
self.tau = tau
# Entropy coefficient / Entropy temperature # Entropy coefficient / Entropy temperature
# Inverse of the reward scale # Inverse of the reward scale
self.ent_coef = ent_coef self.ent_coef = ent_coef
self.target_update_interval = target_update_interval self.target_update_interval = target_update_interval
self.train_freq = train_freq
self.gradient_steps = gradient_steps
self.n_episodes_rollout = n_episodes_rollout
self.action_noise = action_noise
self.gamma = gamma
self.ent_coef_optimizer = None self.ent_coef_optimizer = None
if _init_setup_model: if _init_setup_model:
@ -254,30 +257,10 @@ class SAC(OffPolicyAlgorithm):
eval_log_path: Optional[str] = None, eval_log_path: Optional[str] = None,
reset_num_timesteps: bool = True) -> OffPolicyAlgorithm: reset_num_timesteps: bool = True) -> OffPolicyAlgorithm:
total_timesteps, callback = self._setup_learn(total_timesteps, eval_env, callback, eval_freq, return super(SAC, self).learn(total_timesteps=total_timesteps, callback=callback, log_interval=log_interval,
n_eval_episodes, eval_log_path, reset_num_timesteps, eval_env=eval_env, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes,
tb_log_name) tb_log_name=tb_log_name, eval_log_path=eval_log_path,
callback.on_training_start(locals(), globals()) reset_num_timesteps=reset_num_timesteps)
while self.num_timesteps < total_timesteps:
rollout = self.collect_rollouts(self.env, n_episodes=self.n_episodes_rollout,
n_steps=self.train_freq, action_noise=self.action_noise,
callback=callback,
learning_starts=self.learning_starts,
replay_buffer=self.replay_buffer,
log_interval=log_interval)
if rollout.continue_training is False:
break
self._update_current_progress_remaining(self.num_timesteps, total_timesteps)
if self.num_timesteps > 0 and self.num_timesteps > self.learning_starts:
gradient_steps = self.gradient_steps if self.gradient_steps > 0 else rollout.episode_timesteps
self.train(gradient_steps, batch_size=self.batch_size)
callback.on_training_end()
return self
def excluded_save_params(self) -> List[str]: def excluded_save_params(self) -> List[str]:
""" """

View file

@ -34,6 +34,9 @@ class TD3(OffPolicyAlgorithm):
Note that this cannot be used at the same time as ``train_freq`` Note that this cannot be used at the same time as ``train_freq``
:param action_noise: (ActionNoise) the action noise type (None by default), this can help :param action_noise: (ActionNoise) the action noise type (None by default), this can help
for hard exploration problem. Cf common.noise for the different action noise type. for hard exploration problem. Cf common.noise for the different action noise type.
:param optimize_memory_usage: (bool) Enable a memory efficient variant of the replay buffer
at a cost of more complexity.
See https://github.com/DLR-RM/stable-baselines3/issues/37#issuecomment-637501195
:param policy_delay: (int) Policy and target networks will only be updated once every policy_delay steps :param policy_delay: (int) Policy and target networks will only be updated once every policy_delay steps
per training steps. The Q values will be updated policy_delay more often (update every training step). per training steps. The Q values will be updated policy_delay more often (update every training step).
:param target_policy_noise: (float) Standard deviation of Gaussian noise added to target policy :param target_policy_noise: (float) Standard deviation of Gaussian noise added to target policy
@ -61,6 +64,7 @@ class TD3(OffPolicyAlgorithm):
gradient_steps: int = -1, gradient_steps: int = -1,
n_episodes_rollout: int = 1, n_episodes_rollout: int = 1,
action_noise: Optional[ActionNoise] = None, action_noise: Optional[ActionNoise] = None,
optimize_memory_usage: bool = False,
policy_delay: int = 2, policy_delay: int = 2,
target_policy_noise: float = 0.2, target_policy_noise: float = 0.2,
target_noise_clip: float = 0.5, target_noise_clip: float = 0.5,
@ -74,16 +78,14 @@ class TD3(OffPolicyAlgorithm):
super(TD3, self).__init__(policy, env, TD3Policy, learning_rate, super(TD3, self).__init__(policy, env, TD3Policy, learning_rate,
buffer_size, learning_starts, batch_size, buffer_size, learning_starts, batch_size,
policy_kwargs, tensorboard_log, verbose, device, tau, gamma, train_freq, gradient_steps,
n_episodes_rollout, action_noise=action_noise,
policy_kwargs=policy_kwargs,
tensorboard_log=tensorboard_log,
verbose=verbose, device=device,
create_eval_env=create_eval_env, seed=seed, create_eval_env=create_eval_env, seed=seed,
sde_support=False) sde_support=False, optimize_memory_usage=optimize_memory_usage)
self.train_freq = train_freq
self.gradient_steps = gradient_steps
self.n_episodes_rollout = n_episodes_rollout
self.tau = tau
self.gamma = gamma
self.action_noise = action_noise
self.policy_delay = policy_delay self.policy_delay = policy_delay
self.target_noise_clip = target_noise_clip self.target_noise_clip = target_noise_clip
self.target_policy_noise = target_policy_noise self.target_policy_noise = target_policy_noise
@ -101,7 +103,7 @@ class TD3(OffPolicyAlgorithm):
self.critic = self.policy.critic self.critic = self.policy.critic
self.critic_target = self.policy.critic_target self.critic_target = self.policy.critic_target
def train(self, gradient_steps: int, batch_size: int = 100, policy_delay: int = 2) -> None: def train(self, gradient_steps: int, batch_size: int = 100) -> None:
# Update learning rate according to lr schedule # Update learning rate according to lr schedule
self._update_learning_rate([self.actor.optimizer, self.critic.optimizer]) self._update_learning_rate([self.actor.optimizer, self.critic.optimizer])
@ -134,7 +136,7 @@ class TD3(OffPolicyAlgorithm):
self.critic.optimizer.step() self.critic.optimizer.step()
# Delayed policy updates # Delayed policy updates
if gradient_step % policy_delay == 0: if gradient_step % self.policy_delay == 0:
# Compute actor loss # Compute actor loss
actor_loss = -self.critic.q1_forward(replay_data.observations, actor_loss = -self.critic.q1_forward(replay_data.observations,
self.actor(replay_data.observations)).mean() self.actor(replay_data.observations)).mean()
@ -165,32 +167,10 @@ class TD3(OffPolicyAlgorithm):
eval_log_path: Optional[str] = None, eval_log_path: Optional[str] = None,
reset_num_timesteps: bool = True) -> OffPolicyAlgorithm: reset_num_timesteps: bool = True) -> OffPolicyAlgorithm:
total_timesteps, callback = self._setup_learn(total_timesteps, eval_env, callback, eval_freq, return super(TD3, self).learn(total_timesteps=total_timesteps, callback=callback, log_interval=log_interval,
n_eval_episodes, eval_log_path, reset_num_timesteps, eval_env=eval_env, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes,
tb_log_name) tb_log_name=tb_log_name, eval_log_path=eval_log_path,
callback.on_training_start(locals(), globals()) reset_num_timesteps=reset_num_timesteps)
while self.num_timesteps < total_timesteps:
rollout = self.collect_rollouts(self.env, n_episodes=self.n_episodes_rollout,
n_steps=self.train_freq, action_noise=self.action_noise,
callback=callback,
learning_starts=self.learning_starts,
replay_buffer=self.replay_buffer,
log_interval=log_interval)
if rollout.continue_training is False:
break
self._update_current_progress_remaining(self.num_timesteps, total_timesteps)
if self.num_timesteps > 0 and self.num_timesteps > self.learning_starts:
gradient_steps = self.gradient_steps if self.gradient_steps > 0 else rollout.episode_timesteps
self.train(gradient_steps, batch_size=self.batch_size, policy_delay=self.policy_delay)
callback.on_training_end()
return self
def excluded_save_params(self) -> List[str]: def excluded_save_params(self) -> List[str]:
""" """

View file

@ -1 +1 @@
0.8.0a0 0.8.0a1

View file

@ -4,21 +4,24 @@ import shutil
import pytest import pytest
import gym import gym
from stable_baselines3 import A2C, PPO, SAC, TD3 from stable_baselines3 import A2C, PPO, SAC, TD3, DQN
from stable_baselines3.common.callbacks import (CallbackList, CheckpointCallback, EvalCallback, from stable_baselines3.common.callbacks import (CallbackList, CheckpointCallback, EvalCallback,
EveryNTimesteps, StopTrainingOnRewardThreshold) EveryNTimesteps, StopTrainingOnRewardThreshold)
@pytest.mark.parametrize("model_class", [A2C, PPO, SAC, TD3]) @pytest.mark.parametrize("model_class", [A2C, PPO, SAC, TD3, DQN])
def test_callbacks(model_class): def test_callbacks(tmp_path, model_class):
log_folder = './logs/callbacks/' log_folder = tmp_path / 'logs/callbacks/'
# Dyn only support discrete actions
env_name = select_env(model_class)
# Create RL model # Create RL model
# Small network for fast test # Small network for fast test
model = model_class('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[32])) model = model_class('MlpPolicy', env_name, policy_kwargs=dict(net_arch=[32]))
checkpoint_callback = CheckpointCallback(save_freq=1000, save_path=log_folder) checkpoint_callback = CheckpointCallback(save_freq=1000, save_path=log_folder)
eval_env = gym.make('Pendulum-v0') eval_env = gym.make(env_name)
# Stop training if the performance is good enough # Stop training if the performance is good enough
callback_on_best = StopTrainingOnRewardThreshold(reward_threshold=-1200, verbose=1) callback_on_best = StopTrainingOnRewardThreshold(reward_threshold=-1200, verbose=1)
@ -42,3 +45,10 @@ def test_callbacks(model_class):
model.learn(500, callback=lambda _locals, _globals: True) model.learn(500, callback=lambda _locals, _globals: True)
if os.path.exists(log_folder): if os.path.exists(log_folder):
shutil.rmtree(log_folder) shutil.rmtree(log_folder)
def select_env(model_class) -> str:
if model_class is DQN:
return 'CartPole-v0'
else:
return 'Pendulum-v0'

View file

@ -3,14 +3,13 @@ import os
import numpy as np import numpy as np
import pytest import pytest
from stable_baselines3 import A2C, PPO, SAC, TD3 from stable_baselines3 import A2C, PPO, SAC, TD3, DQN
from stable_baselines3.common.identity_env import FakeImageEnv from stable_baselines3.common.identity_env import FakeImageEnv
SAVE_PATH = './cnn_model.zip'
@pytest.mark.parametrize('model_class', [A2C, PPO, SAC, TD3, DQN])
@pytest.mark.parametrize('model_class', [A2C, PPO, SAC, TD3]) def test_cnn(tmp_path, model_class):
def test_cnn(model_class): SAVE_NAME = 'cnn_model.zip'
# Fake grayscale with frameskip # Fake grayscale with frameskip
# Atari after preprocessing: 84x84x1, here we are using lower resolution # Atari after preprocessing: 84x84x1, here we are using lower resolution
# to check that the network handle it automatically # to check that the network handle it automatically
@ -29,12 +28,12 @@ def test_cnn(model_class):
action, _ = model.predict(obs, deterministic=True) action, _ = model.predict(obs, deterministic=True)
model.save(SAVE_PATH) model.save(tmp_path / SAVE_NAME)
del model del model
model = model_class.load(SAVE_PATH) model = model_class.load(tmp_path / SAVE_NAME)
# Check that the prediction is the same # Check that the prediction is the same
assert np.allclose(action, model.predict(obs, deterministic=True)[0]) assert np.allclose(action, model.predict(obs, deterministic=True)[0])
os.remove(SAVE_PATH) os.remove(str(tmp_path / SAVE_NAME))

View file

@ -1,13 +1,13 @@
import pytest import pytest
from stable_baselines3 import A2C, PPO, SAC, TD3 from stable_baselines3 import A2C, PPO, SAC, TD3, DQN
from stable_baselines3.common.noise import NormalActionNoise from stable_baselines3.common.noise import NormalActionNoise
N_STEPS_TRAINING = 3000 N_STEPS_TRAINING = 3000
SEED = 0 SEED = 0
@pytest.mark.parametrize("algo", [A2C, PPO, SAC, TD3]) @pytest.mark.parametrize("algo", [A2C, DQN, PPO, SAC, TD3])
def test_deterministic_training_common(algo): def test_deterministic_training_common(algo):
results = [[], []] results = [[], []]
rewards = [[], []] rewards = [[], []]
@ -19,8 +19,8 @@ def test_deterministic_training_common(algo):
'learning_starts': 100}) 'learning_starts': 100})
else: else:
env_id = 'CartPole-v1' env_id = 'CartPole-v1'
# if algo == DQN: if algo == DQN:
# kwargs.update({'learning_starts': 100}) kwargs.update({'learning_starts': 100})
for i in range(2): for i in range(2):
model = algo('MlpPolicy', env_id, seed=SEED, **kwargs) model = algo('MlpPolicy', env_id, seed=SEED, **kwargs)

View file

@ -1,7 +1,7 @@
import numpy as np import numpy as np
import pytest import pytest
from stable_baselines3 import A2C, PPO, SAC, TD3 from stable_baselines3 import A2C, PPO, SAC, TD3, DQN
from stable_baselines3.common.identity_env import (IdentityEnvBox, IdentityEnv, from stable_baselines3.common.identity_env import (IdentityEnvBox, IdentityEnv,
IdentityEnvMultiBinary, IdentityEnvMultiDiscrete) IdentityEnvMultiBinary, IdentityEnvMultiDiscrete)
@ -9,17 +9,25 @@ from stable_baselines3.common.vec_env import DummyVecEnv
from stable_baselines3.common.evaluation import evaluate_policy from stable_baselines3.common.evaluation import evaluate_policy
from stable_baselines3.common.noise import NormalActionNoise from stable_baselines3.common.noise import NormalActionNoise
DIM = 4 DIM = 4
@pytest.mark.parametrize("model_class", [A2C, PPO]) @pytest.mark.parametrize("model_class", [A2C, PPO, DQN])
@pytest.mark.parametrize("env", [IdentityEnv(DIM), IdentityEnvMultiDiscrete(DIM), IdentityEnvMultiBinary(DIM)]) @pytest.mark.parametrize("env", [IdentityEnv(DIM), IdentityEnvMultiDiscrete(DIM), IdentityEnvMultiBinary(DIM)])
def test_discrete(model_class, env): def test_discrete(model_class, env):
env = DummyVecEnv([lambda: env]) env_ = DummyVecEnv([lambda: env])
model = model_class('MlpPolicy', env, gamma=0.5, seed=1).learn(3000) kwargs = {}
n_steps = 3000
if model_class == DQN:
kwargs = dict(learning_starts=0)
n_steps = 4000
# DQN only support discrete actions
if isinstance(env, (IdentityEnvMultiDiscrete, IdentityEnvMultiBinary)):
return
evaluate_policy(model, env, n_eval_episodes=20, reward_threshold=90) model = model_class('MlpPolicy', env_, gamma=0.5, seed=1, **kwargs).learn(n_steps)
evaluate_policy(model, env_, n_eval_episodes=20, reward_threshold=90)
obs = env.reset() obs = env.reset()
assert np.shape(model.predict(obs)[0]) == np.shape(obs) assert np.shape(model.predict(obs)[0]) == np.shape(obs)

View file

@ -1,7 +1,7 @@
import gym import gym
import pytest import pytest
from stable_baselines3 import A2C, PPO, SAC, TD3 from stable_baselines3 import A2C, PPO, SAC, TD3, DQN
from stable_baselines3.common.vec_env import DummyVecEnv from stable_baselines3.common.vec_env import DummyVecEnv
MODEL_LIST = [ MODEL_LIST = [
@ -9,14 +9,21 @@ MODEL_LIST = [
A2C, A2C,
TD3, TD3,
SAC, SAC,
DQN,
] ]
@pytest.mark.parametrize("model_class", MODEL_LIST) @pytest.mark.parametrize("model_class", MODEL_LIST)
def test_auto_wrap(model_class): def test_auto_wrap(model_class):
# test auto wrapping of env into a VecEnv # test auto wrapping of env into a VecEnv
env = gym.make('Pendulum-v0')
eval_env = gym.make('Pendulum-v0') # Use different environment for DQN
if model_class is DQN:
env_name = 'CartPole-v0'
else:
env_name = 'Pendulum-v0'
env = gym.make(env_name)
eval_env = gym.make(env_name)
model = model_class('MlpPolicy', env) model = model_class('MlpPolicy', env)
model.learn(100, eval_env=eval_env) model.learn(100, eval_env=eval_env)
@ -24,7 +31,10 @@ def test_auto_wrap(model_class):
@pytest.mark.parametrize("model_class", MODEL_LIST) @pytest.mark.parametrize("model_class", MODEL_LIST)
@pytest.mark.parametrize("env_id", ['Pendulum-v0', 'CartPole-v1']) @pytest.mark.parametrize("env_id", ['Pendulum-v0', 'CartPole-v1'])
def test_predict(model_class, env_id): def test_predict(model_class, env_id):
if env_id == 'CartPole-v1' and model_class not in [PPO, A2C]: if env_id == 'CartPole-v1':
if model_class in [SAC, TD3]:
return
elif model_class in [DQN]:
return return
# test detection of different shapes by the predict method # test detection of different shapes by the predict method

View file

@ -1,7 +1,7 @@
import numpy as np import numpy as np
import pytest import pytest
from stable_baselines3 import A2C, PPO, SAC, TD3 from stable_baselines3 import A2C, PPO, SAC, TD3, DQN
from stable_baselines3.common.noise import NormalActionNoise, OrnsteinUhlenbeckActionNoise from stable_baselines3.common.noise import NormalActionNoise, OrnsteinUhlenbeckActionNoise
normal_action_noise = NormalActionNoise(np.zeros(1), 0.1 * np.ones(1)) normal_action_noise = NormalActionNoise(np.zeros(1), 0.1 * np.ones(1))
@ -40,3 +40,9 @@ def test_sac(ent_coef):
learning_starts=100, verbose=1, create_eval_env=True, ent_coef=ent_coef, learning_starts=100, verbose=1, create_eval_env=True, ent_coef=ent_coef,
action_noise=NormalActionNoise(np.zeros(1), np.zeros(1))) action_noise=NormalActionNoise(np.zeros(1), np.zeros(1)))
model.learn(total_timesteps=1000, eval_freq=500) model.learn(total_timesteps=1000, eval_freq=500)
def test_dqn():
model = DQN('MlpPolicy', 'CartPole-v1', policy_kwargs=dict(net_arch=[64, 64]),
learning_starts=500, buffer_size=500, learning_rate=3e-4, verbose=1, create_eval_env=True)
model.learn(total_timesteps=1000, eval_freq=500)

View file

@ -1,26 +1,39 @@
import os import os
import warnings
from copy import deepcopy from copy import deepcopy
import pytest import pytest
import gym
import numpy as np import numpy as np
import torch as th import torch as th
from stable_baselines3 import A2C, PPO, SAC, TD3 from stable_baselines3 import A2C, PPO, SAC, TD3, DQN
from stable_baselines3.common.identity_env import IdentityEnvBox from stable_baselines3.common.base_class import BaseAlgorithm
from stable_baselines3.common.identity_env import IdentityEnvBox, IdentityEnv
from stable_baselines3.common.vec_env import DummyVecEnv from stable_baselines3.common.vec_env import DummyVecEnv
from stable_baselines3.common.identity_env import FakeImageEnv from stable_baselines3.common.identity_env import FakeImageEnv
MODEL_LIST = [ MODEL_LIST = [
PPO, PPO,
A2C, A2C,
TD3, TD3,
SAC, SAC,
DQN,
] ]
def select_env(model_class: BaseAlgorithm) -> gym.Env:
"""
Selects an environment with the correct action space as DQN only supports discrete action space
"""
if model_class == DQN:
return IdentityEnv(10)
else:
return IdentityEnvBox(10)
@pytest.mark.parametrize("model_class", MODEL_LIST) @pytest.mark.parametrize("model_class", MODEL_LIST)
def test_save_load(model_class): def test_save_load(tmp_path, model_class):
""" """
Test if 'save' and 'load' saves and loads model correctly Test if 'save' and 'load' saves and loads model correctly
and if 'load_parameters' and 'get_policy_parameters' work correctly and if 'load_parameters' and 'get_policy_parameters' work correctly
@ -29,14 +42,15 @@ def test_save_load(model_class):
:param model_class: (BaseAlgorithm) A RL model :param model_class: (BaseAlgorithm) A RL model
""" """
env = DummyVecEnv([lambda: IdentityEnvBox(10)])
env = DummyVecEnv([lambda: select_env(model_class)])
# create model # create model
model = model_class('MlpPolicy', env, policy_kwargs=dict(net_arch=[16]), verbose=1) model = model_class('MlpPolicy', env, policy_kwargs=dict(net_arch=[16]), verbose=1)
model.learn(total_timesteps=500, eval_freq=250) model.learn(total_timesteps=500, eval_freq=250)
env.reset() env.reset()
observations = np.concatenate([env.step(env.action_space.sample())[0] for _ in range(10)], axis=0) observations = np.concatenate([env.step([env.action_space.sample()])[0] for _ in range(10)], axis=0)
# Get dictionary of current parameters # Get dictionary of current parameters
params = deepcopy(model.policy.state_dict()) params = deepcopy(model.policy.state_dict())
@ -58,9 +72,9 @@ def test_save_load(model_class):
selected_actions, _ = model.predict(observations, deterministic=True) selected_actions, _ = model.predict(observations, deterministic=True)
# Check # Check
model.save("test_save.zip") model.save(tmp_path / "test_save.zip")
del model del model
model = model_class.load("test_save", env=env) model = model_class.load(str(tmp_path / "test_save"), env=env)
# check if params are still the same after load # check if params are still the same after load
new_params = model.policy.state_dict() new_params = model.policy.state_dict()
@ -77,7 +91,7 @@ def test_save_load(model_class):
model.learn(total_timesteps=1000, eval_freq=500) model.learn(total_timesteps=1000, eval_freq=500)
# clear file from os # clear file from os
os.remove("test_save.zip") os.remove(tmp_path / "test_save.zip")
@pytest.mark.parametrize("model_class", MODEL_LIST) @pytest.mark.parametrize("model_class", MODEL_LIST)
@ -86,9 +100,11 @@ def test_set_env(model_class):
Test if set_env function does work correct Test if set_env function does work correct
:param model_class: (BaseAlgorithm) A RL model :param model_class: (BaseAlgorithm) A RL model
""" """
env = DummyVecEnv([lambda: IdentityEnvBox(10)])
env2 = DummyVecEnv([lambda: IdentityEnvBox(10)]) # use discrete for DQN
env3 = IdentityEnvBox(10) env = DummyVecEnv([lambda: select_env(model_class)])
env2 = DummyVecEnv([lambda: select_env(model_class)])
env3 = select_env(model_class)
# create model # create model
model = model_class('MlpPolicy', env, policy_kwargs=dict(net_arch=[16])) model = model_class('MlpPolicy', env, policy_kwargs=dict(net_arch=[16]))
@ -107,42 +123,40 @@ def test_set_env(model_class):
@pytest.mark.parametrize("model_class", MODEL_LIST) @pytest.mark.parametrize("model_class", MODEL_LIST)
def test_exclude_include_saved_params(model_class): def test_exclude_include_saved_params(tmp_path, model_class):
""" """
Test if exclude and include parameters of save() work Test if exclude and include parameters of save() work
:param model_class: (BaseAlgorithm) A RL model :param model_class: (BaseAlgorithm) A RL model
""" """
env = DummyVecEnv([lambda: IdentityEnvBox(10)]) env = DummyVecEnv([lambda: select_env(model_class)])
# create model, set verbose as 2, which is not standard # create model, set verbose as 2, which is not standard
model = model_class('MlpPolicy', env, policy_kwargs=dict(net_arch=[16]), verbose=2) model = model_class('MlpPolicy', env, policy_kwargs=dict(net_arch=[16]), verbose=2)
# Check if exclude works # Check if exclude works
model.save("test_save.zip", exclude=["verbose"]) model.save(tmp_path / "test_save.zip", exclude=["verbose"])
del model del model
model = model_class.load("test_save") model = model_class.load(str(tmp_path / "test_save"))
# check if verbose was not saved # check if verbose was not saved
assert model.verbose != 2 assert model.verbose != 2
# set verbose as something different then standard settings # set verbose as something different then standard settings
model.verbose = 2 model.verbose = 2
# Check if include works # Check if include works
model.save("test_save.zip", exclude=["verbose"], include=["verbose"]) model.save(tmp_path / "test_save.zip", exclude=["verbose"], include=["verbose"])
del model del model
model = model_class.load("test_save") model = model_class.load(str(tmp_path / "test_save"))
assert model.verbose == 2 assert model.verbose == 2
# clear file from os # clear file from os
os.remove("test_save.zip") os.remove(tmp_path / "test_save.zip")
@pytest.mark.parametrize("model_class", [SAC, TD3]) @pytest.mark.parametrize("model_class", [SAC, TD3, DQN])
def test_save_load_replay_buffer(model_class): def test_save_load_replay_buffer(tmp_path, model_class):
log_folder = 'logs' replay_path = tmp_path / 'replay_buffer.pkl'
replay_path = os.path.join(log_folder, 'replay_buffer.pkl') model = model_class('MlpPolicy', select_env(model_class), buffer_size=1000)
os.makedirs(log_folder, exist_ok=True)
model = model_class('MlpPolicy', 'Pendulum-v0', buffer_size=1000)
model.learn(500) model.learn(500)
old_replay_buffer = deepcopy(model.replay_buffer) old_replay_buffer = deepcopy(model.replay_buffer)
model.save_replay_buffer(replay_path) model.save_replay_buffer(replay_path)
@ -151,21 +165,54 @@ def test_save_load_replay_buffer(model_class):
assert np.allclose(old_replay_buffer.observations, model.replay_buffer.observations) assert np.allclose(old_replay_buffer.observations, model.replay_buffer.observations)
assert np.allclose(old_replay_buffer.actions, model.replay_buffer.actions) assert np.allclose(old_replay_buffer.actions, model.replay_buffer.actions)
assert np.allclose(old_replay_buffer.next_observations, model.replay_buffer.next_observations)
assert np.allclose(old_replay_buffer.rewards, model.replay_buffer.rewards) assert np.allclose(old_replay_buffer.rewards, model.replay_buffer.rewards)
assert np.allclose(old_replay_buffer.dones, model.replay_buffer.dones) assert np.allclose(old_replay_buffer.dones, model.replay_buffer.dones)
# test extending replay buffer # test extending replay buffer
model.replay_buffer.extend(old_replay_buffer.observations, old_replay_buffer.next_observations, model.replay_buffer.extend(old_replay_buffer.observations, old_replay_buffer.observations,
old_replay_buffer.actions, old_replay_buffer.rewards, old_replay_buffer.dones) old_replay_buffer.actions, old_replay_buffer.rewards, old_replay_buffer.dones)
# clear file from os # clear file from os
os.remove(replay_path) os.remove(replay_path)
@pytest.mark.parametrize("model_class", [DQN, SAC, TD3])
@pytest.mark.parametrize("optimize_memory_usage", [False, True])
def test_warn_buffer(recwarn, model_class, optimize_memory_usage):
"""
When using memory efficient replay buffer,
a warning must be emitted when calling `.learn()`
multiple times.
See https://github.com/DLR-RM/stable-baselines3/issues/46
"""
# remove gym warnings
warnings.filterwarnings(action='ignore', category=DeprecationWarning)
warnings.filterwarnings(action='ignore', category=UserWarning, module='gym')
model = model_class('MlpPolicy', select_env(model_class), buffer_size=100,
optimize_memory_usage=optimize_memory_usage, policy_kwargs=dict(net_arch=[64]),
learning_starts=10)
model.learn(150)
model.learn(150, reset_num_timesteps=False)
# Check that there is no warning
assert len(recwarn) == 0
model.learn(150)
if optimize_memory_usage:
assert len(recwarn) == 1
warning = recwarn.pop(UserWarning)
assert "The last trajectory in the replay buffer will be truncated" in str(warning.message)
else:
assert len(recwarn) == 0
@pytest.mark.parametrize("model_class", MODEL_LIST) @pytest.mark.parametrize("model_class", MODEL_LIST)
@pytest.mark.parametrize("policy_str", ['MlpPolicy', 'CnnPolicy']) @pytest.mark.parametrize("policy_str", ['MlpPolicy', 'CnnPolicy'])
def test_save_load_policy(model_class, policy_str): def test_save_load_policy(tmp_path, model_class, policy_str):
""" """
Test saving and loading policy only. Test saving and loading policy only.
@ -174,14 +221,14 @@ def test_save_load_policy(model_class, policy_str):
""" """
kwargs = {} kwargs = {}
if policy_str == 'MlpPolicy': if policy_str == 'MlpPolicy':
env = IdentityEnvBox(10) env = select_env(model_class)
else: else:
if model_class in [SAC, TD3]: if model_class in [SAC, TD3, DQN]:
# Avoid memory error when using replay buffer # Avoid memory error when using replay buffer
# Reduce the size of the features # Reduce the size of the features
kwargs = dict(buffer_size=250) kwargs = dict(buffer_size=250)
env = FakeImageEnv(screen_height=40, screen_width=40, n_channels=2, env = FakeImageEnv(screen_height=40, screen_width=40, n_channels=2,
discrete=False) discrete=model_class == DQN)
env = DummyVecEnv([lambda: env]) env = DummyVecEnv([lambda: env])
@ -191,7 +238,7 @@ def test_save_load_policy(model_class, policy_str):
model.learn(total_timesteps=500, eval_freq=250) model.learn(total_timesteps=500, eval_freq=250)
env.reset() env.reset()
observations = np.concatenate([env.step(env.action_space.sample())[0] for _ in range(10)], axis=0) observations = np.concatenate([env.step([env.action_space.sample()])[0] for _ in range(10)], axis=0)
policy = model.policy policy = model.policy
policy_class = policy.__class__ policy_class = policy.__class__
@ -223,16 +270,16 @@ def test_save_load_policy(model_class, policy_str):
selected_actions_actor, _ = actor.predict(observations, deterministic=True) selected_actions_actor, _ = actor.predict(observations, deterministic=True)
# Save and load policy # Save and load policy
policy.save("./logs/policy.pkl") policy.save(tmp_path / "policy.pkl")
# Save and load actor # Save and load actor
if actor is not None: if actor is not None:
actor.save("./logs/actor.pkl") actor.save(tmp_path / "actor.pkl")
del policy, actor del policy, actor
policy = policy_class.load("./logs/policy.pkl") policy = policy_class.load(tmp_path / "policy.pkl")
if actor_class is not None: if actor_class is not None:
actor = actor_class.load("./logs/actor.pkl") actor = actor_class.load(tmp_path / "actor.pkl")
# check if params are still the same after load # check if params are still the same after load
new_params = policy.state_dict() new_params = policy.state_dict()
@ -251,6 +298,6 @@ def test_save_load_policy(model_class, policy_str):
assert np.allclose(selected_actions_actor, new_selected_actions, 1e-4) assert np.allclose(selected_actions_actor, new_selected_actions, 1e-4)
# clear file from os # clear file from os
os.remove("./logs/policy.pkl") os.remove(tmp_path / "policy.pkl")
if actor_class is not None: if actor_class is not None:
os.remove("./logs/actor.pkl") os.remove(tmp_path / "actor.pkl")

View file

@ -2,7 +2,7 @@ import numpy as np
import pytest import pytest
import gym import gym
from stable_baselines3 import SAC, TD3 from stable_baselines3 import DQN, SAC, TD3
from stable_baselines3.common.evaluation import evaluate_policy from stable_baselines3.common.evaluation import evaluate_policy
@ -32,13 +32,17 @@ class DummyMultiBinary(gym.Env):
return self.observation_space.sample(), 0.0, False, {} return self.observation_space.sample(), 0.0, False, {}
@pytest.mark.parametrize("model_class", [SAC, TD3]) @pytest.mark.parametrize("model_class", [SAC, TD3, DQN])
@pytest.mark.parametrize("env", [DummyMultiDiscreteSpace([4, 3]), DummyMultiBinary(8)]) @pytest.mark.parametrize("env", [DummyMultiDiscreteSpace([4, 3]), DummyMultiBinary(8)])
def test_identity_spaces(model_class, env): def test_identity_spaces(model_class, env):
""" """
Additional tests for SAC/TD3 to check observation space support Additional tests for DQ/SAC/TD3 to check observation space support
for MultiDiscrete and MultiBinary. for MultiDiscrete and MultiBinary.
""" """
# DQN only support discrete actions
if model_class == DQN:
env.action_space = gym.spaces.Discrete(4)
env = gym.wrappers.TimeLimit(env, max_episode_steps=100) env = gym.wrappers.TimeLimit(env, max_episode_steps=100)
model = model_class("MlpPolicy", env, gamma=0.5, seed=1, policy_kwargs=dict(net_arch=[64])) model = model_class("MlpPolicy", env, gamma=0.5, seed=1, policy_kwargs=dict(net_arch=[64]))

View file

@ -7,11 +7,12 @@ import numpy as np
from stable_baselines3 import A2C from stable_baselines3 import A2C
from stable_baselines3.common.monitor import Monitor from stable_baselines3.common.monitor import Monitor
from stable_baselines3.common.atari_wrappers import ClipRewardEnv
from stable_baselines3.common.evaluation import evaluate_policy from stable_baselines3.common.evaluation import evaluate_policy
from stable_baselines3.common.cmd_util import make_vec_env, make_atari_env from stable_baselines3.common.cmd_util import make_vec_env, make_atari_env
from stable_baselines3.common.vec_env import DummyVecEnv, SubprocVecEnv from stable_baselines3.common.vec_env import DummyVecEnv, SubprocVecEnv
from stable_baselines3.common.noise import ( from stable_baselines3.common.noise import (VectorizedActionNoise,
VectorizedActionNoise, OrnsteinUhlenbeckActionNoise, ActionNoise) OrnsteinUhlenbeckActionNoise, ActionNoise)
@pytest.mark.parametrize("env_id", ['CartPole-v1', lambda: gym.make('CartPole-v1')]) @pytest.mark.parametrize("env_id", ['CartPole-v1', lambda: gym.make('CartPole-v1')])
@ -57,11 +58,11 @@ def test_make_atari_env(env_id, n_envs, wrapper_kwargs):
if wrapper_kwargs is not None: if wrapper_kwargs is not None:
assert obs.shape == (n_envs, 60, 60, 1) assert obs.shape == (n_envs, 60, 60, 1)
assert wrapped_atari_env.observation_space.shape == (60, 60, 1) assert wrapped_atari_env.observation_space.shape == (60, 60, 1)
assert wrapped_atari_env.clip_reward is False assert not isinstance(wrapped_atari_env.env, ClipRewardEnv)
else: else:
assert obs.shape == (n_envs, 84, 84, 1) assert obs.shape == (n_envs, 84, 84, 1)
assert wrapped_atari_env.observation_space.shape == (84, 84, 1) assert wrapped_atari_env.observation_space.shape == (84, 84, 1)
assert wrapped_atari_env.clip_reward is True assert isinstance(wrapped_atari_env.env, ClipRewardEnv)
assert np.max(np.abs(reward)) < 1.0 assert np.max(np.abs(reward)) < 1.0

View file

@ -68,7 +68,7 @@ def test_runningmeanstd():
assert np.allclose(moments_1, moments_2) assert np.allclose(moments_1, moments_2)
def test_vec_env(tmpdir): def test_vec_env(tmp_path):
"""Test VecNormalize Object""" """Test VecNormalize Object"""
clip_obs = 0.5 clip_obs = 0.5
clip_reward = 5.0 clip_reward = 5.0
@ -82,7 +82,7 @@ def test_vec_env(tmpdir):
assert np.max(np.abs(obs)) <= clip_obs assert np.max(np.abs(obs)) <= clip_obs
assert np.max(np.abs(rew)) <= clip_reward assert np.max(np.abs(rew)) <= clip_reward
path = str(tmpdir.join("vec_normalize")) path = tmp_path / "vec_normalize"
norm_venv.save(path) norm_venv.save(path)
deserialized = VecNormalize.load(path, venv=orig_venv) deserialized = VecNormalize.load(path, venv=orig_venv)
check_vec_norm_equal(norm_venv, deserialized) check_vec_norm_equal(norm_venv, deserialized)
@ -125,7 +125,7 @@ def test_offpolicy_normalization(model_class):
eval_env = DummyVecEnv([make_env]) eval_env = DummyVecEnv([make_env])
eval_env = VecNormalize(eval_env, training=False, norm_obs=True, norm_reward=False, clip_obs=10., clip_reward=10.) eval_env = VecNormalize(eval_env, training=False, norm_obs=True, norm_reward=False, clip_obs=10., clip_reward=10.)
model = model_class('MlpPolicy', env, verbose=1) model = model_class('MlpPolicy', env, verbose=1, policy_kwargs=dict(net_arch=[64]))
model.learn(total_timesteps=1000, eval_env=eval_env, eval_freq=500) model.learn(total_timesteps=1000, eval_env=eval_env, eval_freq=500)
# Check getter # Check getter
assert isinstance(model.get_vec_normalize_env(), VecNormalize) assert isinstance(model.get_vec_normalize_env(), VecNormalize)