mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-15 22:10:25 +00:00
Hotfix for Vecnormalize (#558)
* Hotfix for Vecnormalize * Rename `ret` to `returns`
This commit is contained in:
parent
f9e5753acd
commit
f8a0869073
4 changed files with 73 additions and 11 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1 +1 @@
|
|||
1.2.0a3
|
||||
1.2.0
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Reference in a new issue