diff --git a/tests/test_vec_normalize.py b/tests/test_vec_normalize.py index 181c42d..ab0f7f1 100644 --- a/tests/test_vec_normalize.py +++ b/tests/test_vec_normalize.py @@ -3,12 +3,49 @@ import pytest import numpy as np from torchy_baselines.common.running_mean_std import RunningMeanStd -from torchy_baselines.common.vec_env.dummy_vec_env import DummyVecEnv -from torchy_baselines.common.vec_env.vec_normalize import VecNormalize +from torchy_baselines.common.vec_env import DummyVecEnv, VecNormalize, VecFrameStack, sync_envs_normalization from torchy_baselines import CEMRL, SAC, TD3 ENV_ID = 'Pendulum-v0' +def make_env(): + return gym.make(ENV_ID) + +def check_rms_equal(rmsa, rmsb): + assert np.all(rmsa.mean == rmsb.mean) + assert np.all(rmsa.var == rmsb.var) + assert np.all(rmsa.count == rmsb.count) + + +def check_vec_norm_equal(norma, normb): + assert norma.observation_space == normb.observation_space + assert norma.action_space == normb.action_space + assert norma.num_envs == normb.num_envs + + check_rms_equal(norma.obs_rms, normb.obs_rms) + check_rms_equal(norma.ret_rms, normb.ret_rms) + assert norma.clip_obs == normb.clip_obs + assert norma.clip_reward == normb.clip_reward + assert norma.norm_obs == normb.norm_obs + assert norma.norm_reward == normb.norm_reward + + assert np.all(norma.ret == normb.ret) + assert norma.gamma == normb.gamma + assert norma.epsilon == normb.epsilon + assert norma.training == normb.training + +def _make_warmstart_cartpole(): + """Warm-start VecNormalize by stepping through CartPole""" + venv = DummyVecEnv([lambda: gym.make("CartPole-v1")]) + venv = VecNormalize(venv) + venv.reset() + venv.get_original_obs() + + for _ in range(100): + actions = [venv.action_space.sample()] + venv.step(actions) + return venv + def test_runningmeanstd(): """Test RunningMeanStd object""" @@ -27,29 +64,87 @@ def test_runningmeanstd(): assert np.allclose(moments_1, moments_2) -def test_vec_env(): +def test_vec_env(tmpdir): """Test VecNormalize Object""" + clip_obs = 0.5 + clip_reward = 5.0 - def make_env(): - return gym.make(ENV_ID) - - env = DummyVecEnv([make_env]) - env = VecNormalize(env, norm_obs=True, norm_reward=True, clip_obs=10., clip_reward=10.) - _, done = env.reset(), [False] - obs = None + orig_venv = DummyVecEnv([make_env]) + norm_venv = VecNormalize(orig_venv, norm_obs=True, norm_reward=True, clip_obs=clip_obs, clip_reward=clip_reward) + _, done = norm_venv.reset(), [False] while not done[0]: - actions = [env.action_space.sample()] - obs, _, done, _ = env.step(actions) - assert np.max(obs) <= 10 + actions = [norm_venv.action_space.sample()] + obs, rew, done, _ = norm_venv.step(actions) + assert np.max(np.abs(obs)) <= clip_obs + assert np.max(np.abs(rew)) <= clip_reward + + path = str(tmpdir.join("vec_normalize")) + norm_venv.save(path) + deserialized = VecNormalize.load(path, venv=orig_venv) + check_vec_norm_equal(norm_venv, deserialized) + + +def test_get_original(): + venv = _make_warmstart_cartpole() + for _ in range(3): + actions = [venv.action_space.sample()] + obs, rewards, _, _ = venv.step(actions) + obs = obs[0] + orig_obs = venv.get_original_obs()[0] + rewards = rewards[0] + orig_rewards = venv.get_original_reward()[0] + + assert np.all(orig_rewards == 1) + assert orig_obs.shape == obs.shape + assert orig_rewards.dtype == rewards.dtype + assert not np.array_equal(orig_obs, obs) + assert not np.array_equal(orig_rewards, rewards) + np.testing.assert_allclose(venv.normalize_obs(orig_obs), obs) + np.testing.assert_allclose(venv.normalize_reward(orig_rewards), rewards) + + +def test_normalize_external(): + venv = _make_warmstart_cartpole() + + rewards = np.array([1, 1]) + norm_rewards = venv.normalize_reward(rewards) + assert norm_rewards.shape == rewards.shape + # Episode return is almost always >= 1 in CartPole. So reward should shrink. + assert np.all(norm_rewards < 1) @pytest.mark.parametrize("model_class", [SAC, TD3, CEMRL]) def test_offpolicy_normalization(model_class): - env = DummyVecEnv([lambda: gym.make(ENV_ID)]) + env = DummyVecEnv([make_env]) env = VecNormalize(env, norm_obs=True, norm_reward=True, clip_obs=10., clip_reward=10.) - eval_env = DummyVecEnv([lambda: gym.make(ENV_ID)]) - eval_env = VecNormalize(eval_env, norm_obs=True, norm_reward=False, clip_obs=10., clip_reward=10.) + eval_env = DummyVecEnv([make_env]) + eval_env = VecNormalize(eval_env, training=False, norm_obs=True, norm_reward=False, clip_obs=10., clip_reward=10.) model = model_class('MlpPolicy', env, verbose=1) model.learn(total_timesteps=1000, eval_env=eval_env, eval_freq=500) + + +def test_sync_vec_normalize(): + env = DummyVecEnv([make_env]) + env = VecNormalize(env, norm_obs=True, norm_reward=True, clip_obs=10., clip_reward=10.) + env = VecFrameStack(env, 1) + + eval_env = DummyVecEnv([make_env]) + eval_env = VecNormalize(eval_env, training=False, norm_obs=True, norm_reward=True, clip_obs=10., clip_reward=10.) + eval_env = VecFrameStack(eval_env, 1) + + env.reset() + # Initialize running mean + for _ in range(100): + env.step([env.action_space.sample()]) + + obs = env.reset() + original_obs = env.get_original_obs() + # Normalization must be different + assert not np.allclose(obs, eval_env.normalize_obs(original_obs)) + + sync_envs_normalization(env, eval_env) + + # Now they must be synced + assert np.allclose(obs, eval_env.normalize_obs(original_obs)) diff --git a/torchy_baselines/common/vec_env/vec_normalize.py b/torchy_baselines/common/vec_env/vec_normalize.py index 3cb3c67..0824f6a 100644 --- a/torchy_baselines/common/vec_env/vec_normalize.py +++ b/torchy_baselines/common/vec_env/vec_normalize.py @@ -38,6 +38,45 @@ class VecNormalize(VecEnvWrapper): self.old_obs = np.array([]) self.old_reward = np.array([]) + def __getstate__(self): + """ + Gets state for pickling. + + Excludes self.venv, as in general VecEnv's may not be pickleable.""" + state = self.__dict__.copy() + # these attributes are not pickleable + del state['venv'] + del state['class_attributes'] + # these attributes depend on the above and so we would prefer not to pickle + del state['ret'] + return state + + def __setstate__(self, state): + """ + Restores pickled state. + + User must call set_venv() after unpickling before using. + + :param state: (dict)""" + self.__dict__.update(state) + assert 'venv' not in state + self.venv = None + + def set_venv(self, venv): + """ + Sets the vector environment to wrap to venv. + + Also sets attributes derived from this such as `num_env`. + + :param venv: (VecEnv) + """ + if self.venv is not None: + raise ValueError("Trying to set venv of already initialized VecNormalize wrapper.") + VecEnvWrapper.__init__(self, venv) + if self.obs_rms.mean.shape != self.observation_space.shape: + raise ValueError("venv is incompatible with current statistics.") + self.ret = np.zeros(self.num_envs) + def step_wait(self): """ Apply sequence of actions to sequence of environments @@ -46,37 +85,44 @@ class VecNormalize(VecEnvWrapper): where 'news' is a boolean vector indicating whether each element is new. """ obs, rews, news, infos = self.venv.step_wait() - self.ret = self.ret * self.gamma + rews - self.old_obs = obs.copy() - self.old_reward = rews.copy() - obs = self._normalize_observation(obs) - if self.norm_reward: - if self.training: - self.ret_rms.update(self.ret) - rews = self.normalize_reward(rews) + self.old_obs = obs + self.old_rews = rews + + if self.training: + self.obs_rms.update(obs) + obs = self.normalize_obs(obs) + + if self.training: + self._update_reward(rews) + rews = self.normalize_reward(rews) + self.ret[news] = 0 return obs, rews, news, infos - def _normalize_observation(self, obs): - """ - :param obs: (numpy tensor) - """ - if self.norm_obs: - if self.training: - self.obs_rms.update(obs) - return self.normalize_obs(obs) - else: - return obs + def _update_reward(self, reward): + """Update reward normalization statistics.""" + self.ret = self.ret * self.gamma + reward + self.ret_rms.update(self.ret) def normalize_obs(self, obs): + """ + Normalize observations using this VecNormalize's observations statistics. + Calling this method does not update statistics. + """ if self.norm_obs: - return np.clip((obs - self.obs_rms.mean) / np.sqrt(self.obs_rms.var + self.epsilon), -self.clip_obs, - self.clip_obs) + obs = np.clip((obs - self.obs_rms.mean) / np.sqrt(self.obs_rms.var + self.epsilon), + -self.clip_obs, + self.clip_obs) return obs def normalize_reward(self, reward): + """ + Normalize rewards using this VecNormalize's rewards statistics. + Calling this method does not update statistics. + """ if self.norm_reward: - return np.clip(reward / np.sqrt(self.ret_rms.var + self.epsilon), -self.clip_reward, self.clip_reward) + reward = np.clip(reward / np.sqrt(self.ret_rms.var + self.epsilon), + -self.clip_reward, self.clip_reward) return reward def unnormalize_obs(self, obs): @@ -91,31 +137,45 @@ class VecNormalize(VecEnvWrapper): def get_original_obs(self): """ - returns the unnormalized observation - - :return: (numpy float) + Returns an unnormalized version of the observations from the most recent + step or reset. """ - return self.old_obs + return self.old_obs.copy() def get_original_reward(self): """ - returns the unnormalized observation - - :return: (numpy float) + Returns an unnormalized version of the rewards from the most recent step. """ - return self.old_reward + return self.old_rews.copy() def reset(self): """ Reset all environments """ obs = self.venv.reset() - if len(np.array(obs).shape) == 1: # for when num_cpu is 1 - self.old_obs = [obs] - else: - self.old_obs = obs + self.old_obs = obs self.ret = np.zeros(self.num_envs) - return self._normalize_observation(obs) + if self.training: + self._update_reward(self.ret) + return self.normalize_obs(obs) + + @staticmethod + def load(load_path, venv): + """ + Loads a saved VecNormalize object. + + :param load_path: the path to load from. + :param venv: the VecEnv to wrap. + :return: (VecNormalize) + """ + with open(load_path, "rb") as file_handler: + vec_normalize = pickle.load(file_handler) + vec_normalize.set_venv(venv) + return vec_normalize + + def save(self, save_path): + with open(save_path, "wb") as file_handler: + pickle.dump(self, file_handler) def save_running_average(self, path): """