Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions docs/misc/changelog.md
Original file line number Diff line number Diff line change
Expand Up @@ -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]

Expand Down
2 changes: 2 additions & 0 deletions stable_baselines3/common/vec_env/vec_normalize.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
13 changes: 12 additions & 1 deletion tests/test_vec_normalize.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import operator
from copy import deepcopy
from typing import Any

import gymnasium as gym
Expand All @@ -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 (
Expand Down Expand Up @@ -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