TD3 Code review (#245)

* Removed unneeded overrides of feature_extractor and normalize_images in the TD3 Actor.

* Add learning rate schedule example (#248)

* Add learning rate schedule example

* Update docs/guide/examples.rst

Co-authored-by: Adam Gleave <adam@gleave.me>

* Address comments

Co-authored-by: Adam Gleave <adam@gleave.me>

* Add supported action spaces checks (#254)

* Add supported action spaces checks

* Address comment

* Use `pass` in an abstractmethod instead of deleting the arguments.

* Remove the "deterministic" keyword from the forward method of the TD3 Actor since it always is deterministic anyways.

* Rename _get_data to _get_data_to_reconstruct_model.

_get_data was too generic and could have meant anything.

* Remove the n_episodes_rollout parameter and allow passing tuples as train_freq instead.

* Fix docstring of `train_freq` parameter.

* Black fixes.

* Fix TD3 delayed update + rename `_get_data()`

* Fix TD3 test

* Normalize `train_freq` to a tuple in the constructor and turn the warning into an assert.

* Make one step the default train frequency.

* Black fixes.

* Change np.bool to bool.

* Use the tuple format to specify an amount of steps in terms of steps or episodes in the collect_collouts of the off policy algorithm.

* Use the tuple format to specify an amount of steps in terms of steps or episodes in the collect_collouts of HER.

* Use named tuple for train freq

* Rename train_freq to train_every and TrainFreq to ExperienceDuration. Also add some type annotations and documentation.

* Black fixes.

* Revert to train_freq

* Fix terminal observation issues

* Typo

* Fix action noise bug in HER

* Add assert when loading HER models

* Update version

Co-authored-by: Antonin RAFFIN <antonin.raffin@ensta.org>
Co-authored-by: Adam Gleave <adam@gleave.me>
This commit is contained in:
M. Ernestus 2021-02-27 17:33:50 +01:00 committed by GitHub
parent ce3f9e3302
commit 0c50d75ecb
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
21 changed files with 275 additions and 167 deletions

View file

@ -405,6 +405,8 @@ The parking env is a goal-conditioned continuous control task, in which the vehi
model.save("her_sac_highway") model.save("her_sac_highway")
# Load saved model # Load saved model
# Because it needs access to `env.compute_reward()`
# HER must be loaded with the env
model = HER.load("her_sac_highway", env=env) model = HER.load("her_sac_highway", env=env)
obs = env.reset() obs = env.reset()

View file

@ -3,7 +3,7 @@
Changelog Changelog
========== ==========
Pre-Release 0.11.0a7 (WIP) Pre-Release 0.11.0 (2021-02-27)
------------------------------- -------------------------------
Breaking Changes: Breaking Changes:
@ -12,6 +12,18 @@ Breaking Changes:
this allows to return the unnormalized reward in the case of Atari games for instance. this allows to return the unnormalized reward in the case of Atari games for instance.
- Renamed ``common.vec_env.is_wrapped`` to ``common.vec_env.is_vecenv_wrapped`` to avoid confusion - Renamed ``common.vec_env.is_wrapped`` to ``common.vec_env.is_vecenv_wrapped`` to avoid confusion
with the new ``is_wrapped()`` helper with the new ``is_wrapped()`` helper
- Renamed ``_get_data()`` to ``_get_constructor_parameters()`` for policies (this affects independent saving/loading of policies)
- Removed ``n_episodes_rollout`` and merged it with ``train_freq``, which now accepts a tuple ``(frequency, unit)``:
- ``replay_buffer`` in ``collect_rollout`` is no more optional
.. code-block:: python
# SB3 < 0.11.0
# model = SAC("MlpPolicy", env, n_episodes_rollout=1, train_freq=-1)
# SB3 >= 0.11.0:
model = SAC("MlpPolicy", env, train_freq=(1, "episode"))
New Features: New Features:
^^^^^^^^^^^^^ ^^^^^^^^^^^^^
@ -39,7 +51,12 @@ Bug Fixes:
- Added informative ``PPO`` construction error in edge-case scenario where ``n_steps * n_envs = 1`` (size of rollout buffer), - Added informative ``PPO`` construction error in edge-case scenario where ``n_steps * n_envs = 1`` (size of rollout buffer),
which otherwise causes downstream breaking errors in training (@decodyng) which otherwise causes downstream breaking errors in training (@decodyng)
- Fixed discrete observation space support when using multiple envs with A2C/PPO (thanks @ardabbour) - Fixed discrete observation space support when using multiple envs with A2C/PPO (thanks @ardabbour)
- Fixed a bug for TD3 delayed update (the update was off-by-one and not delayed when ``train_freq=1``)
- Fixed numpy warning (replaced ``np.bool`` with ``bool``) - Fixed numpy warning (replaced ``np.bool`` with ``bool``)
- Fixed a bug where ``VecNormalize`` was not normalizing the terminal observation
- Fixed a bug where ``VecTranspose`` was not transposing the terminal observation
- Fixed a bug where the terminal observation stored in the replay buffer was not the right one for off-policy algorithms
- Fixed a bug where ``action_noise`` was not used when using ``HER`` (thanks @ShangqunYu)
Deprecations: Deprecations:
^^^^^^^^^^^^^ ^^^^^^^^^^^^^
@ -550,3 +567,4 @@ And all the contributors:
@tirafesi @blurLake @koulakis @joeljosephjin @shwang @rk37 @andyshih12 @RaphaelWag @xicocaio @tirafesi @blurLake @koulakis @joeljosephjin @shwang @rk37 @andyshih12 @RaphaelWag @xicocaio
@diditforlulz273 @liorcohen5 @ManifoldFR @mloo3 @SwamyDev @wmmc88 @megan-klaiber @thisray @diditforlulz273 @liorcohen5 @ManifoldFR @mloo3 @SwamyDev @wmmc88 @megan-klaiber @thisray
@tfederico @hn2 @LucasAlegre @AptX395 @zampanteymedio @decodyng @ardabbour @lorenz-h @mschweizer @lorepieri8 @tfederico @hn2 @LucasAlegre @AptX395 @zampanteymedio @decodyng @ardabbour @lorenz-h @mschweizer @lorepieri8
@ShangqunYu

View file

@ -32,6 +32,13 @@ It creates "virtual" transitions by relabeling transitions (changing the desired
``HER`` supports ``VecNormalize`` wrapper but only when ``online_sampling=True`` ``HER`` supports ``VecNormalize`` wrapper but only when ``online_sampling=True``
.. warning::
Because it needs access to ``env.compute_reward()``
``HER`` must be loaded with the env. If you just want to use the trained policy
without instantiating the environment, we recommend saving the policy only.
Notes Notes
----- -----
@ -78,11 +85,13 @@ Example
model.learn(1000) model.learn(1000)
model.save("./her_bit_env") model.save("./her_bit_env")
# Because it needs access to `env.compute_reward()`
# HER must be loaded with the env
model = HER.load('./her_bit_env', env=env) model = HER.load('./her_bit_env', env=env)
obs = env.reset() obs = env.reset()
for _ in range(100): for _ in range(100):
action, _ = model.model.predict(obs, deterministic=True) action, _ = model.predict(obs, deterministic=True)
obs, reward, done, _ = env.step(action) obs, reward, done, _ = env.step(action)
if done: if done:

View file

@ -143,7 +143,7 @@ class BaseAlgorithm(ABC):
# Buffers for logging # Buffers for logging
self.ep_info_buffer = None # type: Optional[deque] self.ep_info_buffer = None # type: Optional[deque]
self.ep_success_buffer = None # type: Optional[deque] self.ep_success_buffer = None # type: Optional[deque]
# For logging # For logging (and TD3 delayed updates)
self._n_updates = 0 # type: int self._n_updates = 0 # type: int
# Create and wrap the env if needed # Create and wrap the env if needed
@ -397,10 +397,11 @@ class BaseAlgorithm(ABC):
def _update_info_buffer(self, infos: List[Dict[str, Any]], dones: Optional[np.ndarray] = None) -> None: def _update_info_buffer(self, infos: List[Dict[str, Any]], dones: Optional[np.ndarray] = None) -> None:
""" """
Retrieve reward and episode length and update the buffer Retrieve reward, episode length, episode success and update the buffer
if using Monitor wrapper. if using Monitor wrapper or a GoalEnv.
:param infos: :param infos: List of additional information about the transition.
:param dones: Termination signals
""" """
if dones is None: if dones is None:
dones = np.array([False] * len(infos)) dones = np.array([False] * len(infos))

View file

@ -2,7 +2,7 @@ import io
import pathlib import pathlib
import time import time
import warnings import warnings
from typing import Any, Dict, Optional, Tuple, Type, Union from typing import Any, Dict, List, Optional, Tuple, Type, Union
import gym import gym
import numpy as np import numpy as np
@ -15,8 +15,8 @@ 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.policies import BasePolicy from stable_baselines3.common.policies import BasePolicy
from stable_baselines3.common.save_util import load_from_pkl, save_to_pkl from stable_baselines3.common.save_util import load_from_pkl, save_to_pkl
from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, RolloutReturn, Schedule from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, RolloutReturn, Schedule, TrainFreq, TrainFrequencyUnit
from stable_baselines3.common.utils import safe_mean from stable_baselines3.common.utils import safe_mean, should_collect_more_steps
from stable_baselines3.common.vec_env import VecEnv from stable_baselines3.common.vec_env import VecEnv
@ -35,13 +35,11 @@ class OffPolicyAlgorithm(BaseAlgorithm):
:param batch_size: Minibatch size for each gradient update :param batch_size: Minibatch size for each gradient update
:param tau: the soft update coefficient ("Polyak update", between 0 and 1) :param tau: the soft update coefficient ("Polyak update", between 0 and 1)
:param gamma: the discount factor :param gamma: the discount factor
:param train_freq: Update the model every ``train_freq`` steps. Set to `-1` to disable. :param train_freq: Update the model every ``train_freq`` steps. Alternatively pass a tuple of frequency and unit
:param gradient_steps: How many gradient steps to do after each rollout like ``(5, "step")`` or ``(2, "episode")``.
(see ``train_freq`` and ``n_episodes_rollout``) :param gradient_steps: How many gradient steps to do after each rollout (see ``train_freq``)
Set to ``-1`` means to do as many gradient steps as steps done in the environment Set to ``-1`` means to do as many gradient steps as steps done in the environment
during the rollout. during the rollout.
:param n_episodes_rollout: Update the model every ``n_episodes_rollout`` episodes.
Note that this cannot be used at the same time as ``train_freq``. Set to `-1` to disable.
:param action_noise: the action noise type (None by default), this can help :param action_noise: 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: Enable a memory efficient variant of the replay buffer :param optimize_memory_usage: Enable a memory efficient variant of the replay buffer
@ -83,9 +81,8 @@ class OffPolicyAlgorithm(BaseAlgorithm):
batch_size: int = 256, batch_size: int = 256,
tau: float = 0.005, tau: float = 0.005,
gamma: float = 0.99, gamma: float = 0.99,
train_freq: int = 1, train_freq: Union[int, Tuple[int, str]] = (1, "step"),
gradient_steps: int = 1, gradient_steps: int = 1,
n_episodes_rollout: int = -1,
action_noise: Optional[ActionNoise] = None, action_noise: Optional[ActionNoise] = None,
optimize_memory_usage: bool = False, optimize_memory_usage: bool = False,
policy_kwargs: Dict[str, Any] = None, policy_kwargs: Dict[str, Any] = None,
@ -126,9 +123,7 @@ class OffPolicyAlgorithm(BaseAlgorithm):
self.learning_starts = learning_starts self.learning_starts = learning_starts
self.tau = tau self.tau = tau
self.gamma = gamma self.gamma = gamma
self.train_freq = train_freq
self.gradient_steps = gradient_steps self.gradient_steps = gradient_steps
self.n_episodes_rollout = n_episodes_rollout
self.action_noise = action_noise self.action_noise = action_noise
self.optimize_memory_usage = optimize_memory_usage self.optimize_memory_usage = optimize_memory_usage
@ -136,15 +131,15 @@ class OffPolicyAlgorithm(BaseAlgorithm):
# see https://github.com/hill-a/stable-baselines/issues/863 # see https://github.com/hill-a/stable-baselines/issues/863
self.remove_time_limit_termination = remove_time_limit_termination self.remove_time_limit_termination = remove_time_limit_termination
if train_freq > 0 and n_episodes_rollout > 0: if isinstance(train_freq, int):
warnings.warn( train_freq = (train_freq, "step")
"You passed a positive value for `train_freq` and `n_episodes_rollout`."
"Please make sure this is intended. " try:
"The agent will collect data by stepping in the environment " train_freq = (train_freq[0], TrainFrequencyUnit(train_freq[1]))
"until both conditions are true: " except ValueError:
"`number of steps in the env` >= `train_freq` and " raise ValueError(f"The unit of the `train_freq` must be either 'step' or 'episode' not '{train_freq[1]}'!")
"`number of episodes` > `n_episodes_rollout`"
) self.train_freq = TrainFreq(*train_freq)
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]
@ -168,7 +163,7 @@ class OffPolicyAlgorithm(BaseAlgorithm):
self.observation_space, self.observation_space,
self.action_space, self.action_space,
self.lr_schedule, self.lr_schedule,
**self.policy_kwargs # pytype:disable=not-instantiable **self.policy_kwargs, # pytype:disable=not-instantiable
) )
self.policy = self.policy.to(self.device) self.policy = self.policy.to(self.device)
@ -250,11 +245,9 @@ class OffPolicyAlgorithm(BaseAlgorithm):
callback.on_training_start(locals(), globals()) callback.on_training_start(locals(), globals())
while self.num_timesteps < total_timesteps: while self.num_timesteps < total_timesteps:
rollout = self.collect_rollouts( rollout = self.collect_rollouts(
self.env, self.env,
n_episodes=self.n_episodes_rollout, train_freq=self.train_freq,
n_steps=self.train_freq,
action_noise=self.action_noise, action_noise=self.action_noise,
callback=callback, callback=callback,
learning_starts=self.learning_starts, learning_starts=self.learning_starts,
@ -354,15 +347,62 @@ class OffPolicyAlgorithm(BaseAlgorithm):
""" """
pass pass
def _store_transition(
self,
replay_buffer: ReplayBuffer,
buffer_action: np.ndarray,
new_obs: np.ndarray,
reward: np.ndarray,
done: np.ndarray,
infos: List[Dict[str, Any]],
) -> None:
"""
Store transition in the replay buffer.
We store the normalized action and the unnormalized observation.
It also handles terminal observations (because VecEnv resets automatically).
:param replay_buffer: Replay buffer object where to store the transition.
:param buffer_action: normalized action
:param new_obs: next observation in the current episode
or first observation of the episode (when done is True)
:param reward: reward for the current transition
:param done: Termination signal
:param infos: List of additional information about the transition.
It contains the terminal observations.
"""
# Store only the unnormalized version
if self._vec_normalize_env is not None:
new_obs_ = self._vec_normalize_env.get_original_obs()
reward_ = self._vec_normalize_env.get_original_reward()
else:
# Avoid changing the original ones
self._last_original_obs, new_obs_, reward_ = self._last_obs, new_obs, reward
# As the VecEnv resets automatically, new_obs is already the
# first observation of the next episode
if done and infos[0].get("terminal_observation") is not None:
next_obs = infos[0]["terminal_observation"]
# VecNormalize normalizes the terminal observation
if self._vec_normalize_env is not None:
next_obs = self._vec_normalize_env.unnormalize_obs(next_obs)
else:
next_obs = new_obs_
replay_buffer.add(self._last_original_obs, next_obs, buffer_action, reward_, done)
self._last_obs = new_obs
# Save the unnormalized observation
if self._vec_normalize_env is not None:
self._last_original_obs = new_obs_
def collect_rollouts( def collect_rollouts(
self, self,
env: VecEnv, env: VecEnv,
callback: BaseCallback, callback: BaseCallback,
n_episodes: int = 1, train_freq: TrainFreq,
n_steps: int = -1, replay_buffer: ReplayBuffer,
action_noise: Optional[ActionNoise] = None, action_noise: Optional[ActionNoise] = None,
learning_starts: int = 0, learning_starts: int = 0,
replay_buffer: Optional[ReplayBuffer] = None,
log_interval: Optional[int] = None, log_interval: Optional[int] = None,
) -> RolloutReturn: ) -> RolloutReturn:
""" """
@ -371,10 +411,11 @@ class OffPolicyAlgorithm(BaseAlgorithm):
:param env: The training environment :param env: The training environment
:param callback: Callback that will be called at each step :param callback: Callback that will be called at each step
(and at the beginning and end of the rollout) (and at the beginning and end of the rollout)
:param n_episodes: Number of episodes to use to collect rollout data :param train_freq: How much experience to collect
You can also specify a ``n_steps`` instead by doing rollouts of current policy.
:param n_steps: Number of steps to use to collect rollout data Either ``TrainFreq(<n>, TrainFrequencyUnit.STEP)``
You can also specify a ``n_episodes`` instead. or ``TrainFreq(<n>, TrainFrequencyUnit.EPISODE)``
with ``<n>`` being an integer greater than 0.
:param action_noise: Action noise that will be used for exploration :param action_noise: Action noise that will be used for exploration
Required for deterministic policy (e.g. TD3). This can also be used Required for deterministic policy (e.g. TD3). This can also be used
in addition to the stochastic policy for SAC. in addition to the stochastic policy for SAC.
@ -384,10 +425,11 @@ class OffPolicyAlgorithm(BaseAlgorithm):
:return: :return:
""" """
episode_rewards, total_timesteps = [], [] episode_rewards, total_timesteps = [], []
total_steps, total_episodes = 0, 0 num_collected_steps, num_collected_episodes = 0, 0
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"
assert train_freq.frequency > 0, "Should at least collect one step or episode."
if self.use_sde: if self.use_sde:
self.actor.reset_noise() self.actor.reset_noise()
@ -395,13 +437,13 @@ class OffPolicyAlgorithm(BaseAlgorithm):
callback.on_rollout_start() callback.on_rollout_start()
continue_training = True continue_training = True
while total_steps < n_steps or total_episodes < n_episodes: while should_collect_more_steps(train_freq, num_collected_steps, num_collected_episodes):
done = False done = False
episode_reward, episode_timesteps = 0.0, 0 episode_reward, episode_timesteps = 0.0, 0
while not done: while not done:
if self.use_sde and self.sde_sample_freq > 0 and total_steps % self.sde_sample_freq == 0: if self.use_sde and self.sde_sample_freq > 0 and num_collected_steps % self.sde_sample_freq == 0:
# Sample a new noise matrix # Sample a new noise matrix
self.actor.reset_noise() self.actor.reset_noise()
@ -413,35 +455,21 @@ class OffPolicyAlgorithm(BaseAlgorithm):
self.num_timesteps += 1 self.num_timesteps += 1
episode_timesteps += 1 episode_timesteps += 1
total_steps += 1 num_collected_steps += 1
# Give access to local variables # Give access to local variables
callback.update_locals(locals()) callback.update_locals(locals())
# Only stop training if return value is False, not when it is None. # Only stop training if return value is False, not when it is None.
if callback.on_step() is False: if callback.on_step() is False:
return RolloutReturn(0.0, total_steps, total_episodes, continue_training=False) return RolloutReturn(0.0, num_collected_steps, num_collected_episodes, continue_training=False)
episode_reward += reward episode_reward += reward
# Retrieve reward and episode length if using Monitor wrapper # Retrieve reward and episode length if using Monitor wrapper
self._update_info_buffer(infos, done) self._update_info_buffer(infos, done)
# Store data in replay buffer # Store data in replay buffer (normalized action and unnormalized observation)
if replay_buffer is not None: self._store_transition(replay_buffer, buffer_action, new_obs, reward, done, infos)
# Store only the unnormalized version
if self._vec_normalize_env is not None:
new_obs_ = self._vec_normalize_env.get_original_obs()
reward_ = self._vec_normalize_env.get_original_reward()
else:
# Avoid changing the original ones
self._last_original_obs, new_obs_, reward_ = self._last_obs, new_obs, reward
replay_buffer.add(self._last_original_obs, new_obs_, buffer_action, reward_, done)
self._last_obs = new_obs
# Save the unnormalized observation
if self._vec_normalize_env is not None:
self._last_original_obs = new_obs_
self._update_current_progress_remaining(self.num_timesteps, self._total_timesteps) self._update_current_progress_remaining(self.num_timesteps, self._total_timesteps)
@ -451,11 +479,11 @@ class OffPolicyAlgorithm(BaseAlgorithm):
# see https://github.com/hill-a/stable-baselines/issues/900 # see https://github.com/hill-a/stable-baselines/issues/900
self._on_step() self._on_step()
if 0 < n_steps <= total_steps: if not should_collect_more_steps(train_freq, num_collected_steps, num_collected_episodes):
break break
if done: if done:
total_episodes += 1 num_collected_episodes += 1
self._episode_num += 1 self._episode_num += 1
episode_rewards.append(episode_reward) episode_rewards.append(episode_reward)
total_timesteps.append(episode_timesteps) total_timesteps.append(episode_timesteps)
@ -467,8 +495,8 @@ class OffPolicyAlgorithm(BaseAlgorithm):
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:
self._dump_logs() self._dump_logs()
mean_reward = np.mean(episode_rewards) if total_episodes > 0 else 0.0 mean_reward = np.mean(episode_rewards) if num_collected_episodes > 0 else 0.0
callback.on_rollout_end() callback.on_rollout_end()
return RolloutReturn(mean_reward, total_steps, total_episodes, continue_training) return RolloutReturn(mean_reward, num_collected_steps, num_collected_episodes, continue_training)

View file

@ -82,7 +82,7 @@ class BaseModel(nn.Module, ABC):
@abstractmethod @abstractmethod
def forward(self, *args, **kwargs): def forward(self, *args, **kwargs):
del args, kwargs pass
def _update_features_extractor( def _update_features_extractor(
self, net_kwargs: Dict[str, Any], features_extractor: Optional[BaseFeaturesExtractor] = None self, net_kwargs: Dict[str, Any], features_extractor: Optional[BaseFeaturesExtractor] = None
@ -119,12 +119,11 @@ class BaseModel(nn.Module, ABC):
preprocessed_obs = preprocess_obs(obs, self.observation_space, normalize_images=self.normalize_images) preprocessed_obs = preprocess_obs(obs, self.observation_space, normalize_images=self.normalize_images)
return self.features_extractor(preprocessed_obs) return self.features_extractor(preprocessed_obs)
def _get_data(self) -> Dict[str, Any]: def _get_constructor_parameters(self) -> Dict[str, Any]:
""" """
Get data that need to be saved in order to re-create the model. Get data that need to be saved in order to re-create the model when loading it from disk.
This corresponds to the arguments of the constructor.
:return: :return: The dictionary to pass to the as kwargs constructor when reconstruction this model.
""" """
return dict( return dict(
observation_space=self.observation_space, observation_space=self.observation_space,
@ -151,7 +150,7 @@ class BaseModel(nn.Module, ABC):
:param path: :param path:
""" """
th.save({"state_dict": self.state_dict(), "data": self._get_data()}, path) th.save({"state_dict": self.state_dict(), "data": self._get_constructor_parameters()}, path)
@classmethod @classmethod
def load(cls, path: str, device: Union[th.device, str] = "auto") -> "BaseModel": def load(cls, path: str, device: Union[th.device, str] = "auto") -> "BaseModel":
@ -434,8 +433,8 @@ class ActorCriticPolicy(BasePolicy):
self._build(lr_schedule) self._build(lr_schedule)
def _get_data(self) -> Dict[str, Any]: def _get_constructor_parameters(self) -> Dict[str, Any]:
data = super()._get_data() data = super()._get_constructor_parameters()
default_none_kwargs = self.dist_kwargs or collections.defaultdict(lambda: None) default_none_kwargs = self.dist_kwargs or collections.defaultdict(lambda: None)

View file

@ -1,5 +1,6 @@
"""Common aliases for type hints""" """Common aliases for type hints"""
from enum import Enum
from typing import Any, Callable, Dict, List, NamedTuple, Tuple, Union from typing import Any, Callable, Dict, List, NamedTuple, Tuple, Union
import gym import gym
@ -41,3 +42,13 @@ class RolloutReturn(NamedTuple):
episode_timesteps: int episode_timesteps: int
n_episodes: int n_episodes: int
continue_training: bool continue_training: bool
class TrainFrequencyUnit(Enum):
STEP = "step"
EPISODE = "episode"
class TrainFreq(NamedTuple):
frequency: int
unit: TrainFrequencyUnit # either "step" or "episode"

View file

@ -16,7 +16,7 @@ except ImportError:
SummaryWriter = None SummaryWriter = None
from stable_baselines3.common import logger from stable_baselines3.common import logger
from stable_baselines3.common.type_aliases import GymEnv, Schedule from stable_baselines3.common.type_aliases import GymEnv, Schedule, TrainFreq, TrainFrequencyUnit
def set_random_seed(seed: int, using_cuda: bool = False) -> None: def set_random_seed(seed: int, using_cuda: bool = False) -> None:
@ -317,3 +317,31 @@ def polyak_update(params: Iterable[th.nn.Parameter], target_params: Iterable[th.
for param, target_param in zip_strict(params, target_params): for param, target_param in zip_strict(params, target_params):
target_param.data.mul_(1 - tau) target_param.data.mul_(1 - tau)
th.add(target_param.data, param.data, alpha=tau, out=target_param.data) th.add(target_param.data, param.data, alpha=tau, out=target_param.data)
def should_collect_more_steps(
train_freq: TrainFreq,
num_collected_steps: int,
num_collected_episodes: int,
) -> bool:
"""
Helper used in ``collect_rollouts()`` of off-policy algorithms
to determine the termination condition.
:param train_freq: How much experience should be collected before updating the policy.
:param num_collected_steps: The number of already collected steps.
:param num_collected_episodes: The number of already collected episodes.
:return: Whether to continue or not collecting experience
by doing rollouts of the current policy.
"""
if train_freq.unit == TrainFrequencyUnit.STEP:
return num_collected_steps < train_freq.frequency
elif train_freq.unit == TrainFrequencyUnit.EPISODE:
return num_collected_episodes < train_freq.frequency
else:
raise ValueError(
"The unit of the `train_freq` must be either TrainFrequencyUnit.STEP "
f"or TrainFrequencyUnit.EPISODE not '{train_freq.unit}'!"
)

View file

@ -106,13 +106,13 @@ class VecNormalize(VecEnvWrapper):
def step_wait(self) -> VecEnvStepReturn: def step_wait(self) -> VecEnvStepReturn:
""" """
Apply sequence of actions to sequence of environments Apply sequence of actions to sequence of environments
actions -> (observations, rewards, news) actions -> (observations, rewards, dones)
where 'news' is a boolean vector indicating whether each element is new. where ``dones`` is a boolean vector indicating whether each element is new.
""" """
obs, rews, news, infos = self.venv.step_wait() obs, rewards, dones, infos = self.venv.step_wait()
self.old_obs = obs self.old_obs = obs
self.old_reward = rews self.old_reward = rewards
if self.training: if self.training:
if isinstance(obs, dict) and isinstance(self.obs_rms, dict): if isinstance(obs, dict) and isinstance(self.obs_rms, dict):
@ -124,11 +124,17 @@ class VecNormalize(VecEnvWrapper):
obs = self.normalize_obs(obs) obs = self.normalize_obs(obs)
if self.training: if self.training:
self._update_reward(rews) self._update_reward(rewards)
rews = self.normalize_reward(rews) rewards = self.normalize_reward(rewards)
self.ret[news] = 0 # Normalize the terminal observations
return obs, rews, news, infos for idx, done in enumerate(dones):
if not done:
continue
infos[idx]["terminal_observation"] = self.normalize_obs(infos[idx]["terminal_observation"])
self.ret[dones] = 0
return obs, rewards, dones, infos
def _update_reward(self, reward: np.ndarray) -> None: def _update_reward(self, reward: np.ndarray) -> None:
"""Update reward normalization statistics.""" """Update reward normalization statistics."""

View file

@ -46,6 +46,13 @@ class VecTransposeImage(VecEnvWrapper):
def step_wait(self) -> VecEnvStepReturn: def step_wait(self) -> VecEnvStepReturn:
observations, rewards, dones, infos = self.venv.step_wait() observations, rewards, dones, infos = self.venv.step_wait()
# Transpose the terminal observations
for idx, done in enumerate(dones):
if not done:
continue
infos[idx]["terminal_observation"] = self.transpose_image(infos[idx]["terminal_observation"])
return self.transpose_image(observations), rewards, dones, infos return self.transpose_image(observations), rewards, dones, infos
def reset(self) -> np.ndarray: def reset(self) -> np.ndarray:

View file

@ -1,4 +1,4 @@
from typing import Any, Dict, Optional, Type, Union from typing import Any, Dict, Optional, Tuple, Type, Union
import torch as th import torch as th
@ -29,13 +29,11 @@ class DDPG(TD3):
:param batch_size: Minibatch size for each gradient update :param batch_size: Minibatch size for each gradient update
:param tau: the soft update coefficient ("Polyak update", between 0 and 1) :param tau: the soft update coefficient ("Polyak update", between 0 and 1)
:param gamma: the discount factor :param gamma: the discount factor
:param train_freq: Update the model every ``train_freq`` steps. Set to `-1` to disable. :param train_freq: Update the model every ``train_freq`` steps. Alternatively pass a tuple of frequency and unit
:param gradient_steps: How many gradient steps to do after each rollout like ``(5, "step")`` or ``(2, "episode")``.
(see ``train_freq`` and ``n_episodes_rollout``) :param gradient_steps: How many gradient steps to do after each rollout (see ``train_freq``)
Set to ``-1`` means to do as many gradient steps as steps done in the environment Set to ``-1`` means to do as many gradient steps as steps done in the environment
during the rollout. during the rollout.
:param n_episodes_rollout: Update the model every ``n_episodes_rollout`` episodes.
Note that this cannot be used at the same time as ``train_freq``. Set to `-1` to disable.
:param action_noise: the action noise type (None by default), this can help :param action_noise: 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: Enable a memory efficient variant of the replay buffer :param optimize_memory_usage: Enable a memory efficient variant of the replay buffer
@ -61,9 +59,8 @@ class DDPG(TD3):
batch_size: int = 100, batch_size: int = 100,
tau: float = 0.005, tau: float = 0.005,
gamma: float = 0.99, gamma: float = 0.99,
train_freq: int = -1, train_freq: Union[int, Tuple[int, str]] = (1, "episode"),
gradient_steps: int = -1, gradient_steps: int = -1,
n_episodes_rollout: int = 1,
action_noise: Optional[ActionNoise] = None, action_noise: Optional[ActionNoise] = None,
optimize_memory_usage: bool = False, optimize_memory_usage: bool = False,
tensorboard_log: Optional[str] = None, tensorboard_log: Optional[str] = None,
@ -86,7 +83,6 @@ class DDPG(TD3):
gamma=gamma, gamma=gamma,
train_freq=train_freq, train_freq=train_freq,
gradient_steps=gradient_steps, gradient_steps=gradient_steps,
n_episodes_rollout=n_episodes_rollout,
action_noise=action_noise, action_noise=action_noise,
policy_kwargs=policy_kwargs, policy_kwargs=policy_kwargs,
tensorboard_log=tensorboard_log, tensorboard_log=tensorboard_log,

View file

@ -29,13 +29,11 @@ class DQN(OffPolicyAlgorithm):
:param batch_size: Minibatch size for each gradient update :param batch_size: Minibatch size for each gradient update
:param tau: the soft update coefficient ("Polyak update", between 0 and 1) default 1 for hard update :param tau: the soft update coefficient ("Polyak update", between 0 and 1) default 1 for hard update
:param gamma: the discount factor :param gamma: the discount factor
:param train_freq: Update the model every ``train_freq`` steps. Set to `-1` to disable. :param train_freq: Update the model every ``train_freq`` steps. Alternatively pass a tuple of frequency and unit
:param gradient_steps: How many gradient steps to do after each rollout like ``(5, "step")`` or ``(2, "episode")``.
(see ``train_freq`` and ``n_episodes_rollout``) :param gradient_steps: How many gradient steps to do after each rollout (see ``train_freq``)
Set to ``-1`` means to do as many gradient steps as steps done in the environment Set to ``-1`` means to do as many gradient steps as steps done in the environment
during the rollout. during the rollout.
:param n_episodes_rollout: Update the model every ``n_episodes_rollout`` episodes.
Note that this cannot be used at the same time as ``train_freq``. Set to `-1` to disable.
:param optimize_memory_usage: Enable a memory efficient variant of the replay buffer :param optimize_memory_usage: Enable a memory efficient variant of the replay buffer
at a cost of more complexity. at a cost of more complexity.
See https://github.com/DLR-RM/stable-baselines3/issues/37#issuecomment-637501195 See https://github.com/DLR-RM/stable-baselines3/issues/37#issuecomment-637501195
@ -66,9 +64,8 @@ class DQN(OffPolicyAlgorithm):
batch_size: Optional[int] = 32, batch_size: Optional[int] = 32,
tau: float = 1.0, tau: float = 1.0,
gamma: float = 0.99, gamma: float = 0.99,
train_freq: int = 4, train_freq: Union[int, Tuple[int, str]] = 4,
gradient_steps: int = 1, gradient_steps: int = 1,
n_episodes_rollout: int = -1,
optimize_memory_usage: bool = False, optimize_memory_usage: bool = False,
target_update_interval: int = 10000, target_update_interval: int = 10000,
exploration_fraction: float = 0.1, exploration_fraction: float = 0.1,
@ -96,7 +93,6 @@ class DQN(OffPolicyAlgorithm):
gamma, gamma,
train_freq, train_freq,
gradient_steps, gradient_steps,
n_episodes_rollout,
action_noise=None, # No action noise action_noise=None, # No action noise
policy_kwargs=policy_kwargs, policy_kwargs=policy_kwargs,
tensorboard_log=tensorboard_log, tensorboard_log=tensorboard_log,

View file

@ -65,8 +65,8 @@ class QNetwork(BasePolicy):
action = q_values.argmax(dim=1).reshape(-1) action = q_values.argmax(dim=1).reshape(-1)
return action return action
def _get_data(self) -> Dict[str, Any]: def _get_constructor_parameters(self) -> Dict[str, Any]:
data = super()._get_data() data = super()._get_constructor_parameters()
data.update( data.update(
dict( dict(
@ -168,8 +168,8 @@ class DQNPolicy(BasePolicy):
def _predict(self, obs: th.Tensor, deterministic: bool = True) -> th.Tensor: def _predict(self, obs: th.Tensor, deterministic: bool = True) -> th.Tensor:
return self.q_net._predict(obs, deterministic=deterministic) return self.q_net._predict(obs, deterministic=deterministic)
def _get_data(self) -> Dict[str, Any]: def _get_constructor_parameters(self) -> Dict[str, Any]:
data = super()._get_data() data = super()._get_constructor_parameters()
data.update( data.update(
dict( dict(

View file

@ -12,8 +12,8 @@ from stable_baselines3.common.noise import ActionNoise
from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm
from stable_baselines3.common.policies import BasePolicy from stable_baselines3.common.policies import BasePolicy
from stable_baselines3.common.save_util import load_from_zip_file, recursive_setattr from stable_baselines3.common.save_util import load_from_zip_file, recursive_setattr
from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, RolloutReturn from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback, RolloutReturn, TrainFreq
from stable_baselines3.common.utils import check_for_correct_spaces from stable_baselines3.common.utils import check_for_correct_spaces, should_collect_more_steps
from stable_baselines3.common.vec_env import VecEnv from stable_baselines3.common.vec_env import VecEnv
from stable_baselines3.common.vec_env.obs_dict_wrapper import ObsDictWrapper from stable_baselines3.common.vec_env.obs_dict_wrapper import ObsDictWrapper
from stable_baselines3.her.goal_selection_strategy import KEY_TO_GOAL_STRATEGY, GoalSelectionStrategy from stable_baselines3.her.goal_selection_strategy import KEY_TO_GOAL_STRATEGY, GoalSelectionStrategy
@ -108,6 +108,8 @@ class HER(BaseAlgorithm):
**kwargs, # pytype: disable=wrong-keyword-args **kwargs, # pytype: disable=wrong-keyword-args
) )
# Make HER use self.model.action_noise
del self.action_noise
self.verbose = self.model.verbose self.verbose = self.model.verbose
self.tensorboard_log = self.model.tensorboard_log self.tensorboard_log = self.model.tensorboard_log
@ -132,6 +134,9 @@ class HER(BaseAlgorithm):
# storage for transitions of current episode for offline sampling # storage for transitions of current episode for offline sampling
# for online sampling, it replaces the "classic" replay buffer completely # for online sampling, it replaces the "classic" replay buffer completely
her_buffer_size = self.buffer_size if online_sampling else self.max_episode_length her_buffer_size = self.buffer_size if online_sampling else self.max_episode_length
assert self.env is not None, "Because it needs access to `env.compute_reward()` HER you must provide the env."
self._episode_storage = HerReplayBuffer( self._episode_storage = HerReplayBuffer(
self.env, self.env,
her_buffer_size, her_buffer_size,
@ -193,11 +198,9 @@ class HER(BaseAlgorithm):
callback.on_training_start(locals(), globals()) callback.on_training_start(locals(), globals())
while self.num_timesteps < total_timesteps: while self.num_timesteps < total_timesteps:
rollout = self.collect_rollouts( rollout = self.collect_rollouts(
self.env, self.env,
n_episodes=self.n_episodes_rollout, train_freq=self.train_freq,
n_steps=self.train_freq,
action_noise=self.action_noise, action_noise=self.action_noise,
callback=callback, callback=callback,
learning_starts=self.learning_starts, learning_starts=self.learning_starts,
@ -221,8 +224,7 @@ class HER(BaseAlgorithm):
self, self,
env: VecEnv, env: VecEnv,
callback: BaseCallback, callback: BaseCallback,
n_episodes: int = 1, train_freq: TrainFreq,
n_steps: int = -1,
action_noise: Optional[ActionNoise] = None, action_noise: Optional[ActionNoise] = None,
learning_starts: int = 0, learning_starts: int = 0,
log_interval: Optional[int] = None, log_interval: Optional[int] = None,
@ -233,10 +235,11 @@ class HER(BaseAlgorithm):
:param env: The training environment :param env: The training environment
:param callback: Callback that will be called at each step :param callback: Callback that will be called at each step
(and at the beginning and end of the rollout) (and at the beginning and end of the rollout)
:param n_episodes: Number of episodes to use to collect rollout data :param train_freq: How much experience to collect
You can also specify a ``n_steps`` instead by doing rollouts of current policy.
:param n_steps: Number of steps to use to collect rollout data Either ``TrainFreq(<n>, TrainFrequencyUnit.STEP)``
You can also specify a ``n_episodes`` instead. or ``TrainFreq(<n>, TrainFrequencyUnit.EPISODE)``
with ``<n>`` being an integer greater than 0.
:param action_noise: Action noise that will be used for exploration :param action_noise: Action noise that will be used for exploration
Required for deterministic policy (e.g. TD3). This can also be used Required for deterministic policy (e.g. TD3). This can also be used
in addition to the stochastic policy for SAC. in addition to the stochastic policy for SAC.
@ -246,10 +249,11 @@ class HER(BaseAlgorithm):
""" """
episode_rewards, total_timesteps = [], [] episode_rewards, total_timesteps = [], []
total_steps, total_episodes = 0, 0 num_collected_steps, num_collected_episodes = 0, 0
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"
assert train_freq.frequency > 0, "Should at least collect one step or episode."
if self.model.use_sde: if self.model.use_sde:
self.actor.reset_noise() self.actor.reset_noise()
@ -257,7 +261,7 @@ class HER(BaseAlgorithm):
callback.on_rollout_start() callback.on_rollout_start()
continue_training = True continue_training = True
while total_steps < n_steps or total_episodes < n_episodes: while should_collect_more_steps(train_freq, num_collected_steps, num_collected_episodes):
done = False done = False
episode_reward, episode_timesteps = 0.0, 0 episode_reward, episode_timesteps = 0.0, 0
@ -266,7 +270,11 @@ class HER(BaseAlgorithm):
observation = self._last_obs observation = self._last_obs
self._last_obs = ObsDictWrapper.convert_dict(observation) self._last_obs = ObsDictWrapper.convert_dict(observation)
if self.model.use_sde and self.model.sde_sample_freq > 0 and total_steps % self.model.sde_sample_freq == 0: if (
self.model.use_sde
and self.model.sde_sample_freq > 0
and num_collected_steps % self.model.sde_sample_freq == 0
):
# Sample a new noise matrix # Sample a new noise matrix
self.actor.reset_noise() self.actor.reset_noise()
@ -280,11 +288,11 @@ class HER(BaseAlgorithm):
self.num_timesteps += 1 self.num_timesteps += 1
self.model.num_timesteps = self.num_timesteps self.model.num_timesteps = self.num_timesteps
episode_timesteps += 1 episode_timesteps += 1
total_steps += 1 num_collected_steps += 1
# Only stop training if return value is False, not when it is None. # Only stop training if return value is False, not when it is None.
if callback.on_step() is False: if callback.on_step() is False:
return RolloutReturn(0.0, total_steps, total_episodes, continue_training=False) return RolloutReturn(0.0, num_collected_steps, num_collected_episodes, continue_training=False)
episode_reward += reward episode_reward += reward
@ -307,10 +315,10 @@ class HER(BaseAlgorithm):
# As the VecEnv resets automatically, new_obs is already the # As the VecEnv resets automatically, new_obs is already the
# first observation of the next episode # first observation of the next episode
if done and infos[0].get("terminal_observation") is not None: if done and infos[0].get("terminal_observation") is not None:
# The saved terminal_observation is not passed through other
# VecEnvWrapper, so no need to unnormalize
# NOTE: this may be an issue when using other wrappers
next_obs = infos[0]["terminal_observation"] next_obs = infos[0]["terminal_observation"]
# VecNormalize normalizes the terminal observation
if self._vec_normalize_env is not None:
next_obs = self._vec_normalize_env.unnormalize_obs(next_obs)
else: else:
next_obs = new_obs_ next_obs = new_obs_
@ -343,7 +351,7 @@ class HER(BaseAlgorithm):
self.episode_steps += 1 self.episode_steps += 1
if 0 < n_steps <= total_steps: if not should_collect_more_steps(train_freq, num_collected_steps, num_collected_episodes):
break break
if done or self.episode_steps >= self.max_episode_length: if done or self.episode_steps >= self.max_episode_length:
@ -356,7 +364,7 @@ class HER(BaseAlgorithm):
# clear storage for current episode # clear storage for current episode
self._episode_storage.reset() self._episode_storage.reset()
total_episodes += 1 num_collected_episodes += 1
self._episode_num += 1 self._episode_num += 1
self.model._episode_num = self._episode_num self.model._episode_num = self._episode_num
episode_rewards.append(episode_reward) episode_rewards.append(episode_reward)
@ -371,11 +379,11 @@ class HER(BaseAlgorithm):
self.episode_steps = 0 self.episode_steps = 0
mean_reward = np.mean(episode_rewards) if total_episodes > 0 else 0.0 mean_reward = np.mean(episode_rewards) if num_collected_episodes > 0 else 0.0
callback.on_rollout_end() callback.on_rollout_end()
return RolloutReturn(mean_reward, total_steps, total_episodes, continue_training) return RolloutReturn(mean_reward, num_collected_steps, num_collected_episodes, continue_training)
def _sample_her_transitions(self) -> None: def _sample_her_transitions(self) -> None:
""" """

View file

@ -112,8 +112,8 @@ class Actor(BasePolicy):
self.mu = nn.Linear(last_layer_dim, action_dim) self.mu = nn.Linear(last_layer_dim, action_dim)
self.log_std = nn.Linear(last_layer_dim, action_dim) self.log_std = nn.Linear(last_layer_dim, action_dim)
def _get_data(self) -> Dict[str, Any]: def _get_constructor_parameters(self) -> Dict[str, Any]:
data = super()._get_data() data = super()._get_constructor_parameters()
data.update( data.update(
dict( dict(
@ -316,8 +316,8 @@ class SACPolicy(BasePolicy):
self.critic.optimizer = self.optimizer_class(critic_parameters, lr=lr_schedule(1), **self.optimizer_kwargs) self.critic.optimizer = self.optimizer_class(critic_parameters, lr=lr_schedule(1), **self.optimizer_kwargs)
def _get_data(self) -> Dict[str, Any]: def _get_constructor_parameters(self) -> Dict[str, Any]:
data = super()._get_data() data = super()._get_constructor_parameters()
data.update( data.update(
dict( dict(

View file

@ -37,13 +37,11 @@ class SAC(OffPolicyAlgorithm):
:param batch_size: Minibatch size for each gradient update :param batch_size: Minibatch size for each gradient update
:param tau: the soft update coefficient ("Polyak update", between 0 and 1) :param tau: the soft update coefficient ("Polyak update", between 0 and 1)
:param gamma: the discount factor :param gamma: the discount factor
:param train_freq: Update the model every ``train_freq`` steps. Set to `-1` to disable. :param train_freq: Update the model every ``train_freq`` steps. Alternatively pass a tuple of frequency and unit
:param gradient_steps: How many gradient steps to do after each rollout like ``(5, "step")`` or ``(2, "episode")``.
(see ``train_freq`` and ``n_episodes_rollout``) :param gradient_steps: How many gradient steps to do after each rollout (see ``train_freq``)
Set to ``-1`` means to do as many gradient steps as steps done in the environment Set to ``-1`` means to do as many gradient steps as steps done in the environment
during the rollout. during the rollout.
:param n_episodes_rollout: Update the model every ``n_episodes_rollout`` episodes.
Note that this cannot be used at the same time as ``train_freq``. Set to `-1` to disable.
:param action_noise: the action noise type (None by default), this can help :param action_noise: 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: Enable a memory efficient variant of the replay buffer :param optimize_memory_usage: Enable a memory efficient variant of the replay buffer
@ -81,9 +79,8 @@ class SAC(OffPolicyAlgorithm):
batch_size: int = 256, batch_size: int = 256,
tau: float = 0.005, tau: float = 0.005,
gamma: float = 0.99, gamma: float = 0.99,
train_freq: int = 1, train_freq: Union[int, Tuple[int, str]] = 1,
gradient_steps: int = 1, gradient_steps: int = 1,
n_episodes_rollout: int = -1,
action_noise: Optional[ActionNoise] = None, action_noise: Optional[ActionNoise] = None,
optimize_memory_usage: bool = False, optimize_memory_usage: bool = False,
ent_coef: Union[str, float] = "auto", ent_coef: Union[str, float] = "auto",
@ -113,7 +110,6 @@ class SAC(OffPolicyAlgorithm):
gamma, gamma,
train_freq, train_freq,
gradient_steps, gradient_steps,
n_episodes_rollout,
action_noise, action_noise,
policy_kwargs=policy_kwargs, policy_kwargs=policy_kwargs,
tensorboard_log=tensorboard_log, tensorboard_log=tensorboard_log,

View file

@ -49,8 +49,6 @@ class Actor(BasePolicy):
squash_output=True, squash_output=True,
) )
self.features_extractor = features_extractor
self.normalize_images = normalize_images
self.net_arch = net_arch self.net_arch = net_arch
self.features_dim = features_dim self.features_dim = features_dim
self.activation_fn = activation_fn self.activation_fn = activation_fn
@ -60,8 +58,8 @@ class Actor(BasePolicy):
# Deterministic action # Deterministic action
self.mu = nn.Sequential(*actor_net) self.mu = nn.Sequential(*actor_net)
def _get_data(self) -> Dict[str, Any]: def _get_constructor_parameters(self) -> Dict[str, Any]:
data = super()._get_data() data = super()._get_constructor_parameters()
data.update( data.update(
dict( dict(
@ -73,13 +71,15 @@ class Actor(BasePolicy):
) )
return data return data
def forward(self, obs: th.Tensor, deterministic: bool = True) -> th.Tensor: def forward(self, obs: th.Tensor) -> th.Tensor:
# assert deterministic, 'The TD3 actor only outputs deterministic actions' # assert deterministic, 'The TD3 actor only outputs deterministic actions'
features = self.extract_features(obs) features = self.extract_features(obs)
return self.mu(features) return self.mu(features)
def _predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor: def _predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
return self.forward(observation, deterministic=deterministic) # Note: the deterministic deterministic parameter is ignored in the case of TD3.
# Predictions are always deterministic.
return self.forward(observation)
class TD3Policy(BasePolicy): class TD3Policy(BasePolicy):
@ -190,8 +190,8 @@ class TD3Policy(BasePolicy):
self.critic_target.load_state_dict(self.critic.state_dict()) self.critic_target.load_state_dict(self.critic.state_dict())
self.critic.optimizer = self.optimizer_class(self.critic.parameters(), lr=lr_schedule(1), **self.optimizer_kwargs) self.critic.optimizer = self.optimizer_class(self.critic.parameters(), lr=lr_schedule(1), **self.optimizer_kwargs)
def _get_data(self) -> Dict[str, Any]: def _get_constructor_parameters(self) -> Dict[str, Any]:
data = super()._get_data() data = super()._get_constructor_parameters()
data.update( data.update(
dict( dict(
@ -220,7 +220,9 @@ class TD3Policy(BasePolicy):
return self._predict(observation, deterministic=deterministic) return self._predict(observation, deterministic=deterministic)
def _predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor: def _predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
return self.actor(observation, deterministic=deterministic) # Note: the deterministic deterministic parameter is ignored in the case of TD3.
# Predictions are always deterministic.
return self.actor(observation)
MlpPolicy = TD3Policy MlpPolicy = TD3Policy

View file

@ -32,13 +32,11 @@ class TD3(OffPolicyAlgorithm):
:param batch_size: Minibatch size for each gradient update :param batch_size: Minibatch size for each gradient update
:param tau: the soft update coefficient ("Polyak update", between 0 and 1) :param tau: the soft update coefficient ("Polyak update", between 0 and 1)
:param gamma: the discount factor :param gamma: the discount factor
:param train_freq: Update the model every ``train_freq`` steps. Set to `-1` to disable. :param train_freq: Update the model every ``train_freq`` steps. Alternatively pass a tuple of frequency and unit
:param gradient_steps: How many gradient steps to do after each rollout like ``(5, "step")`` or ``(2, "episode")``.
(see ``train_freq`` and ``n_episodes_rollout``) :param gradient_steps: How many gradient steps to do after each rollout (see ``train_freq``)
Set to ``-1`` means to do as many gradient steps as steps done in the environment Set to ``-1`` means to do as many gradient steps as steps done in the environment
during the rollout. during the rollout.
:param n_episodes_rollout: Update the model every ``n_episodes_rollout`` episodes.
Note that this cannot be used at the same time as ``train_freq``. Set to `-1` to disable.
:param action_noise: the action noise type (None by default), this can help :param action_noise: 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: Enable a memory efficient variant of the replay buffer :param optimize_memory_usage: Enable a memory efficient variant of the replay buffer
@ -69,9 +67,8 @@ class TD3(OffPolicyAlgorithm):
batch_size: int = 100, batch_size: int = 100,
tau: float = 0.005, tau: float = 0.005,
gamma: float = 0.99, gamma: float = 0.99,
train_freq: int = -1, train_freq: Union[int, Tuple[int, str]] = (1, "episode"),
gradient_steps: int = -1, gradient_steps: int = -1,
n_episodes_rollout: int = 1,
action_noise: Optional[ActionNoise] = None, action_noise: Optional[ActionNoise] = None,
optimize_memory_usage: bool = False, optimize_memory_usage: bool = False,
policy_delay: int = 2, policy_delay: int = 2,
@ -98,7 +95,6 @@ class TD3(OffPolicyAlgorithm):
gamma, gamma,
train_freq, train_freq,
gradient_steps, gradient_steps,
n_episodes_rollout,
action_noise=action_noise, action_noise=action_noise,
policy_kwargs=policy_kwargs, policy_kwargs=policy_kwargs,
tensorboard_log=tensorboard_log, tensorboard_log=tensorboard_log,
@ -137,6 +133,7 @@ class TD3(OffPolicyAlgorithm):
for gradient_step in range(gradient_steps): for gradient_step in range(gradient_steps):
self._n_updates += 1
# Sample replay buffer # Sample replay buffer
replay_data = self.replay_buffer.sample(batch_size, env=self._vec_normalize_env) replay_data = self.replay_buffer.sample(batch_size, env=self._vec_normalize_env)
@ -164,7 +161,7 @@ class TD3(OffPolicyAlgorithm):
self.critic.optimizer.step() self.critic.optimizer.step()
# Delayed policy updates # Delayed policy updates
if gradient_step % self.policy_delay == 0: if self._n_updates % self.policy_delay == 0:
# Compute actor loss # Compute actor loss
actor_loss = -self.critic.q1_forward(replay_data.observations, self.actor(replay_data.observations)).mean() actor_loss = -self.critic.q1_forward(replay_data.observations, self.actor(replay_data.observations)).mean()
actor_losses.append(actor_loss.item()) actor_losses.append(actor_loss.item())
@ -177,9 +174,9 @@ class TD3(OffPolicyAlgorithm):
polyak_update(self.critic.parameters(), self.critic_target.parameters(), self.tau) polyak_update(self.critic.parameters(), self.critic_target.parameters(), self.tau)
polyak_update(self.actor.parameters(), self.actor_target.parameters(), self.tau) polyak_update(self.actor.parameters(), self.actor_target.parameters(), self.tau)
self._n_updates += gradient_steps
logger.record("train/n_updates", self._n_updates, exclude="tensorboard") logger.record("train/n_updates", self._n_updates, exclude="tensorboard")
logger.record("train/actor_loss", np.mean(actor_losses)) if len(actor_losses) > 0:
logger.record("train/actor_loss", np.mean(actor_losses))
logger.record("train/critic_loss", np.mean(critic_losses)) logger.record("train/critic_loss", np.mean(critic_losses))
def learn( def learn(

View file

@ -1 +1 @@
0.11.0a7 0.11.0

View file

@ -88,6 +88,10 @@ def test_features_extractor_target_net(model_class, share_features_extractor):
if model_class != DQN: if model_class != DQN:
kwargs["policy_kwargs"]["share_features_extractor"] = share_features_extractor kwargs["policy_kwargs"]["share_features_extractor"] = share_features_extractor
# No delay for TD3 (changes when the actor and polyak update take place)
if model_class == TD3:
kwargs["policy_delay"] = 1
model = model_class("CnnPolicy", env, seed=0, **kwargs) model = model_class("CnnPolicy", env, seed=0, **kwargs)
patch_dqn_names_(model) patch_dqn_names_(model)

View file

@ -10,6 +10,7 @@ import torch as th
from stable_baselines3 import DDPG, DQN, HER, SAC, TD3 from stable_baselines3 import DDPG, DQN, HER, SAC, TD3
from stable_baselines3.common.bit_flipping_env import BitFlippingEnv from stable_baselines3.common.bit_flipping_env import BitFlippingEnv
from stable_baselines3.common.noise import NormalActionNoise
from stable_baselines3.common.vec_env import DummyVecEnv from stable_baselines3.common.vec_env import DummyVecEnv
from stable_baselines3.common.vec_env.obs_dict_wrapper import ObsDictWrapper from stable_baselines3.common.vec_env.obs_dict_wrapper import ObsDictWrapper
from stable_baselines3.her.goal_selection_strategy import GoalSelectionStrategy from stable_baselines3.her.goal_selection_strategy import GoalSelectionStrategy
@ -32,8 +33,7 @@ def test_her(model_class, online_sampling):
goal_selection_strategy="future", goal_selection_strategy="future",
online_sampling=online_sampling, online_sampling=online_sampling,
gradient_steps=1, gradient_steps=1,
train_freq=1, train_freq=4,
n_episodes_rollout=-1,
max_episode_length=n_bits, max_episode_length=n_bits,
policy_kwargs=dict(net_arch=[64]), policy_kwargs=dict(net_arch=[64]),
learning_starts=100, learning_starts=100,
@ -60,6 +60,8 @@ def test_goal_selection_strategy(goal_selection_strategy, online_sampling):
""" """
env = BitFlippingEnv(continuous=True) env = BitFlippingEnv(continuous=True)
normal_action_noise = NormalActionNoise(np.zeros(1), 0.1 * np.ones(1))
model = HER( model = HER(
"MlpPolicy", "MlpPolicy",
env, env,
@ -67,12 +69,13 @@ def test_goal_selection_strategy(goal_selection_strategy, online_sampling):
goal_selection_strategy=goal_selection_strategy, goal_selection_strategy=goal_selection_strategy,
online_sampling=online_sampling, online_sampling=online_sampling,
gradient_steps=1, gradient_steps=1,
train_freq=1, train_freq=4,
n_episodes_rollout=-1,
max_episode_length=10, max_episode_length=10,
policy_kwargs=dict(net_arch=[64]), policy_kwargs=dict(net_arch=[64]),
learning_starts=100, learning_starts=100,
action_noise=normal_action_noise,
) )
assert model.action_noise is not None
model.learn(total_timesteps=300) model.learn(total_timesteps=300)
@ -109,7 +112,6 @@ def test_save_load(tmp_path, model_class, use_sde, online_sampling):
gradient_steps=1, gradient_steps=1,
train_freq=4, train_freq=4,
learning_starts=100, learning_starts=100,
n_episodes_rollout=-1,
max_episode_length=n_bits, max_episode_length=n_bits,
**kwargs **kwargs
) )
@ -172,7 +174,7 @@ def test_save_load(tmp_path, model_class, use_sde, online_sampling):
os.remove(tmp_path / "test_save.zip") os.remove(tmp_path / "test_save.zip")
@pytest.mark.parametrize("online_sampling, truncate_last_trajectory", [(False, None), (True, True), (True, False)]) @pytest.mark.parametrize("online_sampling, truncate_last_trajectory", [(False, False), (True, True), (True, False)])
def test_save_load_replay_buffer(tmp_path, recwarn, online_sampling, truncate_last_trajectory): def test_save_load_replay_buffer(tmp_path, recwarn, online_sampling, truncate_last_trajectory):
""" """
Test if 'save_replay_buffer' and 'load_replay_buffer' works correctly Test if 'save_replay_buffer' and 'load_replay_buffer' works correctly
@ -191,8 +193,7 @@ def test_save_load_replay_buffer(tmp_path, recwarn, online_sampling, truncate_la
goal_selection_strategy="future", goal_selection_strategy="future",
online_sampling=online_sampling, online_sampling=online_sampling,
gradient_steps=1, gradient_steps=1,
train_freq=1, train_freq=4,
n_episodes_rollout=-1,
max_episode_length=4, max_episode_length=4,
buffer_size=int(2e4), buffer_size=int(2e4),
policy_kwargs=dict(net_arch=[64]), policy_kwargs=dict(net_arch=[64]),
@ -266,8 +267,7 @@ def test_full_replay_buffer():
goal_selection_strategy="future", goal_selection_strategy="future",
online_sampling=True, online_sampling=True,
gradient_steps=1, gradient_steps=1,
train_freq=1, train_freq=4,
n_episodes_rollout=-1,
max_episode_length=n_bits, max_episode_length=n_bits,
policy_kwargs=dict(net_arch=[64]), policy_kwargs=dict(net_arch=[64]),
learning_starts=1, learning_starts=1,