mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-29 20:14:17 +00:00
Update VecNormalize (pickling) and improve tests
This commit is contained in:
parent
89db65b1fb
commit
e5c6601726
2 changed files with 205 additions and 50 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Reference in a new issue