Hotfix for Vecnormalize (#558)

* Hotfix for Vecnormalize

* Rename `ret` to `returns`
This commit is contained in:
Antonin RAFFIN 2021-09-08 12:30:20 +02:00 committed by GitHub
parent f9e5753acd
commit f8a0869073
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
4 changed files with 73 additions and 11 deletions

View file

@ -4,18 +4,22 @@ Changelog
========== ==========
Release 1.2.0a3 (WIP) Release 1.2.0 (2021-09-03)
--------------------------- ---------------------------
**Hotfix for VecNormalize, training/eval mode support**
Breaking Changes: Breaking Changes:
^^^^^^^^^^^^^^^^^ ^^^^^^^^^^^^^^^^^
- SB3 now requires PyTorch >= 1.8.1 - SB3 now requires PyTorch >= 1.8.1
- ``VecNormalize`` ``ret`` attribute was renamed to ``returns``
New Features: New Features:
^^^^^^^^^^^^^ ^^^^^^^^^^^^^
Bug Fixes: Bug Fixes:
^^^^^^^^^^ ^^^^^^^^^^
- Hotfix for ``VecNormalize`` where the observation filter was not updated at reset (thanks @vwxyzjn)
- Fixed model predictions when using batch normalization and dropout layers by calling ``train()`` and ``eval()`` (@davidblom603) - Fixed model predictions when using batch normalization and dropout layers by calling ``train()`` and ``eval()`` (@davidblom603)
- Fixed model training for DQN, TD3 and SAC so that their target nets always remain in evaluation mode (@ayeright) - Fixed model training for DQN, TD3 and SAC so that their target nets always remain in evaluation mode (@ayeright)
- Passing ``gradient_steps=0`` to an off-policy algorithm will result in no gradient steps being taken (vs as many gradient steps as steps done in the environment - Passing ``gradient_steps=0`` to an off-policy algorithm will result in no gradient steps being taken (vs as many gradient steps as steps done in the environment

View file

@ -1,4 +1,5 @@
import pickle import pickle
import warnings
from copy import deepcopy from copy import deepcopy
from typing import Any, Dict, Union from typing import Any, Dict, Union
@ -54,7 +55,7 @@ class VecNormalize(VecEnvWrapper):
self.clip_obs = clip_obs self.clip_obs = clip_obs
self.clip_reward = clip_reward self.clip_reward = clip_reward
# Returns: discounted rewards # Returns: discounted rewards
self.ret = np.zeros(self.num_envs) self.returns = np.zeros(self.num_envs)
self.gamma = gamma self.gamma = gamma
self.epsilon = epsilon self.epsilon = epsilon
self.training = training self.training = training
@ -73,7 +74,7 @@ class VecNormalize(VecEnvWrapper):
del state["venv"] del state["venv"]
del state["class_attributes"] del state["class_attributes"]
# these attributes depend on the above and so we would prefer not to pickle # these attributes depend on the above and so we would prefer not to pickle
del state["ret"] del state["returns"]
return state return state
def __setstate__(self, state: Dict[str, Any]) -> None: def __setstate__(self, state: Dict[str, Any]) -> None:
@ -101,7 +102,7 @@ class VecNormalize(VecEnvWrapper):
# Check only that the observation_space match # Check only that the observation_space match
utils.check_for_correct_spaces(venv, self.observation_space, venv.action_space) utils.check_for_correct_spaces(venv, self.observation_space, venv.action_space)
self.ret = np.zeros(self.num_envs) self.returns = np.zeros(self.num_envs)
def step_wait(self) -> VecEnvStepReturn: def step_wait(self) -> VecEnvStepReturn:
""" """
@ -134,13 +135,13 @@ class VecNormalize(VecEnvWrapper):
if "terminal_observation" in infos[idx]: if "terminal_observation" in infos[idx]:
infos[idx]["terminal_observation"] = self.normalize_obs(infos[idx]["terminal_observation"]) infos[idx]["terminal_observation"] = self.normalize_obs(infos[idx]["terminal_observation"])
self.ret[dones] = 0 self.returns[dones] = 0
return obs, rewards, dones, infos 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."""
self.ret = self.ret * self.gamma + reward self.returns = self.returns * self.gamma + reward
self.ret_rms.update(self.ret) self.ret_rms.update(self.returns)
def _normalize_obs(self, obs: np.ndarray, obs_rms: RunningMeanStd) -> np.ndarray: def _normalize_obs(self, obs: np.ndarray, obs_rms: RunningMeanStd) -> np.ndarray:
""" """
@ -220,9 +221,13 @@ class VecNormalize(VecEnvWrapper):
""" """
obs = self.venv.reset() obs = self.venv.reset()
self.old_obs = obs self.old_obs = obs
self.ret = np.zeros(self.num_envs) self.returns = np.zeros(self.num_envs)
if self.training: if self.training:
self._update_reward(self.ret) if isinstance(obs, dict) and isinstance(self.obs_rms, dict):
for key in self.obs_rms.keys():
self.obs_rms[key].update(obs[key])
else:
self.obs_rms.update(obs)
return self.normalize_obs(obs) return self.normalize_obs(obs)
@staticmethod @staticmethod
@ -248,3 +253,8 @@ class VecNormalize(VecEnvWrapper):
""" """
with open(save_path, "wb") as file_handler: with open(save_path, "wb") as file_handler:
pickle.dump(self, file_handler) pickle.dump(self, file_handler)
@property
def ret(self) -> np.ndarray:
warnings.warn("`VecNormalize` `ret` attribute is deprecated. Please use `returns` instead.", DeprecationWarning)
return self.returns

View file

@ -1 +1 @@
1.2.0a3 1.2.0

View file

@ -17,6 +17,27 @@ from stable_baselines3.common.vec_env import (
ENV_ID = "Pendulum-v0" ENV_ID = "Pendulum-v0"
class DummyRewardEnv(gym.Env):
metadata = {}
def __init__(self, return_reward_idx=0):
self.action_space = gym.spaces.Discrete(2)
self.observation_space = gym.spaces.Box(low=np.array([-1.0]), high=np.array([1.0]))
self.returned_rewards = [0, 1, 3, 4]
self.return_reward_idx = return_reward_idx
self.t = self.return_reward_idx
def step(self, action):
self.t += 1
index = (self.t + self.return_reward_idx) % len(self.returned_rewards)
returned_value = self.returned_rewards[index]
return np.array([returned_value]), returned_value, self.t == len(self.returned_rewards), {}
def reset(self):
self.t = 0
return np.array([self.returned_rewards[self.return_reward_idx]])
class DummyDictEnv(gym.GoalEnv): class DummyDictEnv(gym.GoalEnv):
""" """
Dummy gym goal env for testing purposes Dummy gym goal env for testing purposes
@ -69,6 +90,15 @@ def make_dict_env():
return Monitor(DummyDictEnv()) return Monitor(DummyDictEnv())
def test_deprecation():
venv = DummyVecEnv([lambda: gym.make("CartPole-v1")])
venv = VecNormalize(venv)
with pytest.warns(None) as record:
assert np.allclose(venv.ret, venv.returns)
# Deprecation warning when using .ret
assert len(record) == 1
def check_rms_equal(rmsa, rmsb): def check_rms_equal(rmsa, rmsb):
if isinstance(rmsa, dict): if isinstance(rmsa, dict):
for key in rmsa.keys(): for key in rmsa.keys():
@ -93,7 +123,7 @@ def check_vec_norm_equal(norma, normb):
assert norma.norm_obs == normb.norm_obs assert norma.norm_obs == normb.norm_obs
assert norma.norm_reward == normb.norm_reward assert norma.norm_reward == normb.norm_reward
assert np.all(norma.ret == normb.ret) assert np.all(norma.returns == normb.returns)
assert norma.gamma == normb.gamma assert norma.gamma == normb.gamma
assert norma.epsilon == normb.epsilon assert norma.epsilon == normb.epsilon
assert norma.training == normb.training assert norma.training == normb.training
@ -143,6 +173,24 @@ def test_runningmeanstd():
assert np.allclose(moments_1, moments_2) assert np.allclose(moments_1, moments_2)
def test_obs_rms_vec_normalize():
env_fns = [lambda: DummyRewardEnv(0), lambda: DummyRewardEnv(1)]
env = DummyVecEnv(env_fns)
env = VecNormalize(env)
env.reset()
assert np.allclose(env.obs_rms.mean, 0.5, atol=1e-4)
assert np.allclose(env.ret_rms.mean, 0.0, atol=1e-4)
env.step([env.action_space.sample() for _ in range(len(env_fns))])
assert np.allclose(env.obs_rms.mean, 1.25, atol=1e-4)
assert np.allclose(env.ret_rms.mean, 2, atol=1e-4)
# Check convergence to true mean
for _ in range(3000):
env.step([env.action_space.sample() for _ in range(len(env_fns))])
assert np.allclose(env.obs_rms.mean, 2.0, atol=1e-3)
assert np.allclose(env.ret_rms.mean, 5.688, atol=1e-3)
@pytest.mark.parametrize("make_env", [make_env, make_dict_env]) @pytest.mark.parametrize("make_env", [make_env, make_dict_env])
def test_vec_env(tmp_path, make_env): def test_vec_env(tmp_path, make_env):
"""Test VecNormalize Object""" """Test VecNormalize Object"""