diff --git a/docs/misc/changelog.rst b/docs/misc/changelog.rst index 69fce7d..608e882 100644 --- a/docs/misc/changelog.rst +++ b/docs/misc/changelog.rst @@ -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: ^^^^^^^^^^^^^^^^^ - SB3 now requires PyTorch >= 1.8.1 +- ``VecNormalize`` ``ret`` attribute was renamed to ``returns`` New Features: ^^^^^^^^^^^^^ 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 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 diff --git a/stable_baselines3/common/vec_env/vec_normalize.py b/stable_baselines3/common/vec_env/vec_normalize.py index f1feeee..ad7d87c 100644 --- a/stable_baselines3/common/vec_env/vec_normalize.py +++ b/stable_baselines3/common/vec_env/vec_normalize.py @@ -1,4 +1,5 @@ import pickle +import warnings from copy import deepcopy from typing import Any, Dict, Union @@ -54,7 +55,7 @@ class VecNormalize(VecEnvWrapper): self.clip_obs = clip_obs self.clip_reward = clip_reward # Returns: discounted rewards - self.ret = np.zeros(self.num_envs) + self.returns = np.zeros(self.num_envs) self.gamma = gamma self.epsilon = epsilon self.training = training @@ -73,7 +74,7 @@ class VecNormalize(VecEnvWrapper): del state["venv"] del state["class_attributes"] # these attributes depend on the above and so we would prefer not to pickle - del state["ret"] + del state["returns"] return state def __setstate__(self, state: Dict[str, Any]) -> None: @@ -101,7 +102,7 @@ class VecNormalize(VecEnvWrapper): # Check only that the observation_space match 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: """ @@ -134,13 +135,13 @@ class VecNormalize(VecEnvWrapper): if "terminal_observation" in infos[idx]: infos[idx]["terminal_observation"] = self.normalize_obs(infos[idx]["terminal_observation"]) - self.ret[dones] = 0 + self.returns[dones] = 0 return obs, rewards, dones, infos def _update_reward(self, reward: np.ndarray) -> None: """Update reward normalization statistics.""" - self.ret = self.ret * self.gamma + reward - self.ret_rms.update(self.ret) + self.returns = self.returns * self.gamma + reward + self.ret_rms.update(self.returns) def _normalize_obs(self, obs: np.ndarray, obs_rms: RunningMeanStd) -> np.ndarray: """ @@ -220,9 +221,13 @@ class VecNormalize(VecEnvWrapper): """ obs = self.venv.reset() self.old_obs = obs - self.ret = np.zeros(self.num_envs) + self.returns = np.zeros(self.num_envs) 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) @staticmethod @@ -248,3 +253,8 @@ class VecNormalize(VecEnvWrapper): """ with open(save_path, "wb") as 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 diff --git a/stable_baselines3/version.txt b/stable_baselines3/version.txt index 9b27bf0..26aaba0 100644 --- a/stable_baselines3/version.txt +++ b/stable_baselines3/version.txt @@ -1 +1 @@ -1.2.0a3 +1.2.0 diff --git a/tests/test_vec_normalize.py b/tests/test_vec_normalize.py index 63d4dbf..cce63f9 100644 --- a/tests/test_vec_normalize.py +++ b/tests/test_vec_normalize.py @@ -17,6 +17,27 @@ from stable_baselines3.common.vec_env import ( 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): """ Dummy gym goal env for testing purposes @@ -69,6 +90,15 @@ def make_dict_env(): 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): if isinstance(rmsa, dict): 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_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.epsilon == normb.epsilon assert norma.training == normb.training @@ -143,6 +173,24 @@ def test_runningmeanstd(): 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]) def test_vec_env(tmp_path, make_env): """Test VecNormalize Object"""