diff --git a/docs/misc/changelog.md b/docs/misc/changelog.md index 06eb1902b..bf028bedb 100644 --- a/docs/misc/changelog.md +++ b/docs/misc/changelog.md @@ -11,6 +11,7 @@ ### Bug Fixes: - Fixed file descriptor leak in `SubprocVecEnv.close()` by closing parent-side pipes to prevent "too many open files" errors in long-running loops +- Fixed `VecNormalize` writing the normalized image bounds into the `Dict` observation space of the environment it wraps, which changed how a model built on that environment afterwards treated the image ### [SB3-Contrib] diff --git a/stable_baselines3/common/vec_env/vec_normalize.py b/stable_baselines3/common/vec_env/vec_normalize.py index dea4f4c0b..54ee57270 100644 --- a/stable_baselines3/common/vec_env/vec_normalize.py +++ b/stable_baselines3/common/vec_env/vec_normalize.py @@ -54,6 +54,8 @@ def __init__( self._sanity_checks() if isinstance(self.observation_space, spaces.Dict): + # Copy first: the image bounds below are written into this space, which the wrapped env owns + self.observation_space = deepcopy(self.observation_space) self.obs_spaces = self.observation_space.spaces self.obs_rms = {key: RunningMeanStd(shape=self.obs_spaces[key].shape) for key in self.norm_obs_keys} # type: ignore[arg-type, union-attr] # Update observation space when using image diff --git a/tests/test_vec_normalize.py b/tests/test_vec_normalize.py index af7f48a4d..bb963c41b 100644 --- a/tests/test_vec_normalize.py +++ b/tests/test_vec_normalize.py @@ -1,4 +1,5 @@ import operator +from copy import deepcopy from typing import Any import gymnasium as gym @@ -7,7 +8,7 @@ from gymnasium import spaces from stable_baselines3 import SAC, TD3, HerReplayBuffer -from stable_baselines3.common.envs import FakeImageEnv +from stable_baselines3.common.envs import FakeImageEnv, SimpleMultiObsEnv from stable_baselines3.common.monitor import Monitor from stable_baselines3.common.running_mean_std import RunningMeanStd from stable_baselines3.common.vec_env import ( @@ -496,3 +497,13 @@ def test_non_dict_obs_keys(): # Test dict obs with norm_obs set to False _make_warmstart(lambda: DummyMixedDictEnv(), norm_obs=False) + + +def test_vec_normalize_keeps_the_wrapped_obs_space(): + venv = DummyVecEnv([lambda: SimpleMultiObsEnv()]) + original_space = deepcopy(venv.observation_space) + + vec_normalize = VecNormalize(venv) + + assert venv.observation_space == original_space + assert vec_normalize.observation_space["img"].dtype == np.float32