Introduce norm_obs_keys param for VecNormalize environment wrapper (#631)

* Implement new norm_obs_keys param for VecNormalize environment wrapper

* Simplified doc string to avoid issues with lint and doc

* Updated changelog

* Update changelog.rst

* Update test_vec_normalize.py

* Update sanity checks

* Fix backward compat

* Update doc

* Update changelog

* Fix lint warnings

* Fix tests

* Minor edit

* observation_space sanity check was applied twice

Co-authored-by: Oleksii Kachaiev <okachaiev@riotgames.com>
Co-authored-by: Anssi <kaneran21@hotmail.com>
Co-authored-by: Antonin Raffin <antonin.raffin@dlr.de>
This commit is contained in:
Oleksii Kachaiev 2021-10-28 10:18:39 -07:00 committed by GitHub
parent 7b977d7b03
commit 0503e694b2
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
5 changed files with 131 additions and 19 deletions

View file

@ -4,6 +4,29 @@ Changelog
========== ==========
Release 1.3.1a0 (WIP)
---------------------------
Breaking Changes:
^^^^^^^^^^^^^^^^^
New Features:
^^^^^^^^^^^^^
- Added ``norm_obs_keys`` param for ``VecNormalize`` wrapper to configure which observation keys to normalize (@kachayev)
Bug Fixes:
^^^^^^^^^^
Deprecations:
^^^^^^^^^^^^^
Others:
^^^^^^^
Documentation:
^^^^^^^^^^^^^^
Release 1.3.0 (2021-10-23) Release 1.3.0 (2021-10-23)
--------------------------- ---------------------------
@ -30,13 +53,13 @@ New Features:
- Added methods ``get_distribution`` and ``predict_values`` for ``ActorCriticPolicy`` for A2C/PPO/TRPO (@cyprienc) - Added methods ``get_distribution`` and ``predict_values`` for ``ActorCriticPolicy`` for A2C/PPO/TRPO (@cyprienc)
- Added methods ``forward_actor`` and ``forward_critic`` for ``MlpExtractor`` - Added methods ``forward_actor`` and ``forward_critic`` for ``MlpExtractor``
- Added ``sb3.get_system_info()`` helper function to gather version information relevant to SB3 (e.g., Python and PyTorch version) - Added ``sb3.get_system_info()`` helper function to gather version information relevant to SB3 (e.g., Python and PyTorch version)
- Saved models now store system information where agent was trained, and load functions have ``print_system_info`` parameter to help debugging load issues. - Saved models now store system information where agent was trained, and load functions have ``print_system_info`` parameter to help debugging load issues
Bug Fixes: Bug Fixes:
^^^^^^^^^^ ^^^^^^^^^^
- Fixed ``dtype`` of observations for ``SimpleMultiObsEnv`` - Fixed ``dtype`` of observations for ``SimpleMultiObsEnv``
- Allow `VecNormalize` to wrap discrete-observation environments to normalize reward - Allow `VecNormalize` to wrap discrete-observation environments to normalize reward
when observation normalization is disabled. when observation normalization is disabled
- Fixed a bug where ``DQN`` would throw an error when using ``Discrete`` observation and stochastic actions - Fixed a bug where ``DQN`` would throw an error when using ``Discrete`` observation and stochastic actions
- Fixed a bug where sub-classed observation spaces could not be used - Fixed a bug where sub-classed observation spaces could not be used
- Added ``force_reset`` argument to ``load()`` and ``set_env()`` in order to be able to call ``learn(reset_num_timesteps=False)`` with a new environment - Added ``force_reset`` argument to ``load()`` and ``set_env()`` in order to be able to call ``learn(reset_num_timesteps=False)`` with a new environment
@ -804,4 +827,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 @JadenTravnik @decodyng @ardabbour @lorenz-h @mschweizer @lorepieri8 @vwxyzjn @tfederico @hn2 @LucasAlegre @AptX395 @zampanteymedio @JadenTravnik @decodyng @ardabbour @lorenz-h @mschweizer @lorepieri8 @vwxyzjn
@ShangqunYu @PierreExeter @JacopoPan @ltbd78 @tom-doerr @Atlis @liusida @09tangriro @amy12xx @juancroldan @benblack769 @bstee615 @c-rizz @skandermoalla @MihaiAnca13 @davidblom603 @ayeright @cyprienc @wkirgsn @AechPro @CUN-bjy @batu @IljaAvadiev @timokau @ShangqunYu @PierreExeter @JacopoPan @ltbd78 @tom-doerr @Atlis @liusida @09tangriro @amy12xx @juancroldan @benblack769 @bstee615 @c-rizz @skandermoalla @MihaiAnca13 @davidblom603 @ayeright @cyprienc @wkirgsn @AechPro @CUN-bjy @batu @IljaAvadiev @timokau @kachayev

View file

@ -1,7 +1,7 @@
import pickle import pickle
import warnings import warnings
from copy import deepcopy from copy import deepcopy
from typing import Any, Dict, Union from typing import Any, Dict, List, Optional, Union
import gym import gym
import numpy as np import numpy as np
@ -24,6 +24,8 @@ class VecNormalize(VecEnvWrapper):
:param clip_reward: Max value absolute for discounted reward :param clip_reward: Max value absolute for discounted reward
:param gamma: discount factor :param gamma: discount factor
:param epsilon: To avoid division by zero :param epsilon: To avoid division by zero
:param norm_obs_keys: Which keys from observation dict to normalize.
If not specified, all keys will be normalized.
""" """
def __init__( def __init__(
@ -36,19 +38,21 @@ class VecNormalize(VecEnvWrapper):
clip_reward: float = 10.0, clip_reward: float = 10.0,
gamma: float = 0.99, gamma: float = 0.99,
epsilon: float = 1e-8, epsilon: float = 1e-8,
norm_obs_keys: Optional[List[str]] = None,
): ):
VecEnvWrapper.__init__(self, venv) VecEnvWrapper.__init__(self, venv)
if norm_obs: self.norm_obs = norm_obs
if not isinstance(self.observation_space, (gym.spaces.Box, gym.spaces.Dict)): self.norm_obs_keys = norm_obs_keys
raise ValueError("VecNormalize only supports `gym.spaces.Box` and `gym.spaces.Dict` observation spaces") # Check observation spaces
if self.norm_obs:
self._sanity_checks()
if isinstance(self.observation_space, gym.spaces.Dict): if isinstance(self.observation_space, gym.spaces.Dict):
self.obs_keys = set(self.observation_space.spaces.keys())
self.obs_spaces = self.observation_space.spaces self.obs_spaces = self.observation_space.spaces
self.obs_rms = {key: RunningMeanStd(shape=space.shape) for key, space in self.obs_spaces.items()} self.obs_rms = {key: RunningMeanStd(shape=self.obs_spaces[key].shape) for key in self.norm_obs_keys}
else: else:
self.obs_keys, self.obs_spaces = None, None self.obs_spaces = None
self.obs_rms = RunningMeanStd(shape=self.observation_space.shape) self.obs_rms = RunningMeanStd(shape=self.observation_space.shape)
self.ret_rms = RunningMeanStd(shape=()) self.ret_rms = RunningMeanStd(shape=())
@ -64,6 +68,34 @@ class VecNormalize(VecEnvWrapper):
self.old_obs = np.array([]) self.old_obs = np.array([])
self.old_reward = np.array([]) self.old_reward = np.array([])
def _sanity_checks(self) -> None:
"""
Check the observations that are going to be normalized are of the correct type (spaces.Box).
"""
if isinstance(self.observation_space, gym.spaces.Dict):
# By default, we normalize all keys
if self.norm_obs_keys is None:
self.norm_obs_keys = list(self.observation_space.spaces.keys())
# Check that all keys are of type Box
for obs_key in self.norm_obs_keys:
if not isinstance(self.observation_space.spaces[obs_key], gym.spaces.Box):
raise ValueError(
f"VecNormalize only supports `gym.spaces.Box` observation spaces but {obs_key} "
f"is of type {self.observation_space.spaces[obs_key]}. "
"You should probably explicitely pass the observation keys "
" that should be normalized via the `norm_obs_keys` parameter."
)
elif isinstance(self.observation_space, gym.spaces.Box):
if self.norm_obs_keys is not None:
raise ValueError("`norm_obs_keys` param is applicable only with `gym.spaces.Dict` observation spaces")
else:
raise ValueError(
"VecNormalize only supports `gym.spaces.Box` and `gym.spaces.Dict` observation spaces, "
f"not {self.observation_space}"
)
def __getstate__(self) -> Dict[str, Any]: def __getstate__(self) -> Dict[str, Any]:
""" """
Gets state for pickling. Gets state for pickling.
@ -84,6 +116,9 @@ class VecNormalize(VecEnvWrapper):
User must call set_venv() after unpickling before using. User must call set_venv() after unpickling before using.
:param state:""" :param state:"""
# Backward compatibility
if "norm_obs_keys" not in state:
state["norm_obs_keys"] = list(state["observation_space"].spaces.keys())
self.__dict__.update(state) self.__dict__.update(state)
assert "venv" not in state assert "venv" not in state
self.venv = None self.venv = None
@ -170,7 +205,8 @@ class VecNormalize(VecEnvWrapper):
obs_ = deepcopy(obs) obs_ = deepcopy(obs)
if self.norm_obs: if self.norm_obs:
if isinstance(obs, dict) and isinstance(self.obs_rms, dict): if isinstance(obs, dict) and isinstance(self.obs_rms, dict):
for key in self.obs_rms.keys(): # Only normalize the specified keys
for key in self.norm_obs_keys:
obs_[key] = self._normalize_obs(obs[key], self.obs_rms[key]).astype(np.float32) obs_[key] = self._normalize_obs(obs[key], self.obs_rms[key]).astype(np.float32)
else: else:
obs_ = self._normalize_obs(obs, self.obs_rms).astype(np.float32) obs_ = self._normalize_obs(obs, self.obs_rms).astype(np.float32)
@ -190,7 +226,7 @@ class VecNormalize(VecEnvWrapper):
obs_ = deepcopy(obs) obs_ = deepcopy(obs)
if self.norm_obs: if self.norm_obs:
if isinstance(obs, dict) and isinstance(self.obs_rms, dict): if isinstance(obs, dict) and isinstance(self.obs_rms, dict):
for key in self.obs_rms.keys(): for key in self.norm_obs_keys:
obs_[key] = self._unnormalize_obs(obs[key], self.obs_rms[key]) obs_[key] = self._unnormalize_obs(obs[key], self.obs_rms[key])
else: else:
obs_ = self._unnormalize_obs(obs, self.obs_rms) obs_ = self._unnormalize_obs(obs, self.obs_rms)

View file

@ -1 +1 @@
1.3.0 1.3.1a0

View file

@ -66,7 +66,7 @@ class DummyDictEnv(gym.Env):
def step(self, action): def step(self, action):
reward = 0.0 reward = 0.0
done = False done = np.random.rand() > 0.8
return self.observation_space.sample(), reward, done, {} return self.observation_space.sample(), reward, done, {}
def compute_reward(self, achieved_goal, desired_goal, info): def compute_reward(self, achieved_goal, desired_goal, info):
@ -266,8 +266,8 @@ def test_vec_normalize(model_class):
Additional tests for PPO/A2C/SAC/DDPG/TD3/DQN to check observation space support Additional tests for PPO/A2C/SAC/DDPG/TD3/DQN to check observation space support
for GoalEnv and VecNormalize using MultiInputPolicy. for GoalEnv and VecNormalize using MultiInputPolicy.
""" """
env = DummyVecEnv([lambda: BitFlippingEnv(n_bits=4, continuous=not (model_class == DQN))]) env = DummyVecEnv([lambda: DummyDictEnv(use_discrete_actions=model_class == DQN)])
env = VecNormalize(env) env = VecNormalize(env, norm_obs_keys=["vec"])
kwargs = {} kwargs = {}
n_steps = 256 n_steps = 256

View file

@ -1,3 +1,5 @@
import operator
import gym import gym
import numpy as np import numpy as np
import pytest import pytest
@ -68,6 +70,31 @@ class DummyDictEnv(gym.GoalEnv):
return -(distance > 0).astype(np.float32) return -(distance > 0).astype(np.float32)
class DummyMixedDictEnv(gym.Env):
"""
Dummy mixed gym env for testing purposes
"""
def __init__(self):
super().__init__()
self.observation_space = spaces.Dict(
{
"obs1": spaces.Box(low=-20.0, high=20.0, shape=(4,), dtype=np.float32),
"obs2": spaces.Discrete(1),
"obs3": spaces.Box(low=-20.0, high=20.0, shape=(4,), dtype=np.float32),
}
)
self.action_space = spaces.Box(low=-1, high=1, shape=(3,), dtype=np.float32)
def reset(self):
return self.observation_space.sample()
def step(self, action):
obs = self.observation_space.sample()
done = np.random.rand() > 0.8
return obs, 0.0, done, {}
def allclose(obs_1, obs_2): def allclose(obs_1, obs_2):
""" """
Generalized np.allclose() to work with dict spaces. Generalized np.allclose() to work with dict spaces.
@ -152,9 +179,9 @@ def _make_warmstart_cartpole():
return _make_warmstart(lambda: gym.make("CartPole-v1")) return _make_warmstart(lambda: gym.make("CartPole-v1"))
def _make_warmstart_dict_env(): def _make_warmstart_dict_env(**kwargs):
"""Warm-start VecNormalize by stepping through BitFlippingEnv""" """Warm-start VecNormalize by stepping through DummyDictEnv"""
return _make_warmstart(make_dict_env) return _make_warmstart(make_dict_env, **kwargs)
def test_runningmeanstd(): def test_runningmeanstd():
@ -267,6 +294,21 @@ def test_normalize_external():
assert np.all(norm_rewards < 1) assert np.all(norm_rewards < 1)
def test_normalize_dict_selected_keys():
venv = _make_warmstart_dict_env(norm_obs=True, norm_obs_keys=["observation"])
for _ in range(3):
actions = [venv.action_space.sample()]
obs, rewards, _, _ = venv.step(actions)
orig_obs = venv.get_original_obs()
# "observation" is expected to be normalized
np.testing.assert_array_compare(operator.__ne__, obs["observation"], orig_obs["observation"])
assert allclose(venv.normalize_obs(orig_obs), obs)
# other keys are expected to be presented "as is"
np.testing.assert_array_equal(obs["achieved_goal"], orig_obs["achieved_goal"])
@pytest.mark.parametrize("model_class", [SAC, TD3, HerReplayBuffer]) @pytest.mark.parametrize("model_class", [SAC, TD3, HerReplayBuffer])
@pytest.mark.parametrize("online_sampling", [False, True]) @pytest.mark.parametrize("online_sampling", [False, True])
def test_offpolicy_normalization(model_class, online_sampling): def test_offpolicy_normalization(model_class, online_sampling):
@ -358,3 +400,14 @@ def test_discrete_obs():
# Smoke test that it runs with norm_obs False # Smoke test that it runs with norm_obs False
_make_warmstart_cliffwalking(norm_obs=False) _make_warmstart_cliffwalking(norm_obs=False)
def test_non_dict_obs_keys():
with pytest.raises(ValueError, match=".*is applicable only.*"):
_make_warmstart(lambda: DummyRewardEnv(), norm_obs_keys=["key"])
with pytest.raises(ValueError, match=".* explicitely pass the observation keys.*"):
_make_warmstart(lambda: DummyMixedDictEnv())
# Ignore Discrete observation key
_make_warmstart(lambda: DummyMixedDictEnv(), norm_obs_keys=["obs1", "obs3"])