diff --git a/README.md b/README.md index 74b55c6..c13574d 100644 --- a/README.md +++ b/README.md @@ -24,6 +24,7 @@ TODO: - SDE: reduce the number of parameters (only n_features instead of n_features x n_actions) for A2C (done for TD3) - SDE: learn the feature extractor? +- DEBUG normalization with replay buffer (apparently pb with observation normalization) Later: - get_parameters / set_parameters diff --git a/tests/test_vec_normalize.py b/tests/test_vec_normalize.py index 78437a3..7353de0 100644 --- a/tests/test_vec_normalize.py +++ b/tests/test_vec_normalize.py @@ -1,9 +1,11 @@ import gym +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 import CEMRL, SAC, TD3 ENV_ID = 'Pendulum-v0' @@ -39,3 +41,15 @@ def test_vec_env(): actions = [env.action_space.sample()] obs, _, done, _ = env.step(actions) assert np.max(obs) <= 10 + + +@pytest.mark.parametrize("model_class", [SAC, TD3, CEMRL]) +def test_offpolicy_normalization(model_class): + env = DummyVecEnv([lambda: gym.make(ENV_ID)]) + env = VecNormalize(env, norm_obs=False, norm_reward=False, 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.) + + model = model_class('MlpPolicy', env, verbose=1) + model.learn(total_timesteps=10000, eval_env=eval_env, eval_freq=1000) diff --git a/torchy_baselines/cem_rl/cem_rl.py b/torchy_baselines/cem_rl/cem_rl.py index ad798ba..ff7100d 100644 --- a/torchy_baselines/cem_rl/cem_rl.py +++ b/torchy_baselines/cem_rl/cem_rl.py @@ -5,6 +5,7 @@ import torch as th from torchy_baselines.cem_rl.cem import CEM from torchy_baselines.common.evaluation import evaluate_policy from torchy_baselines.td3.td3 import TD3 +from torchy_baselines.common.vec_env import sync_envs_normalization class CEMRL(TD3): @@ -102,7 +103,7 @@ class CEMRL(TD3): n_training_steps = 2 * (actor_steps // self.n_grad) for it in range(n_training_steps): # Sample replay buffer - replay_data = self.replay_buffer.sample(self.batch_size) + replay_data = self.replay_buffer.sample(self.batch_size, env=self._vec_normalize_env) self.train_critic(replay_data=replay_data) # Delayed policy updates @@ -117,6 +118,7 @@ class CEMRL(TD3): timesteps_since_eval %= eval_freq self.actor.load_from_vector(self.es.mu) + sync_envs_normalization(self.env, eval_env) mean_reward, _ = evaluate_policy(self, eval_env, n_eval_episodes) evaluations.append(mean_reward) diff --git a/torchy_baselines/common/base_class.py b/torchy_baselines/common/base_class.py index bf71232..fd82905 100644 --- a/torchy_baselines/common/base_class.py +++ b/torchy_baselines/common/base_class.py @@ -8,7 +8,7 @@ import numpy as np from torchy_baselines.common.policies import get_policy_from_name from torchy_baselines.common.utils import set_random_seed, get_schedule_fn, update_learning_rate -from torchy_baselines.common.vec_env import DummyVecEnv, VecEnv +from torchy_baselines.common.vec_env import DummyVecEnv, VecEnv, unwrap_vec_normalize from torchy_baselines.common.monitor import Monitor from torchy_baselines.common import logger @@ -46,6 +46,8 @@ class BaseRLModel(object): print("Using {} device".format(self.device)) self.env = env + # get VecNormalize object if needed + self._vec_normalize_env = unwrap_vec_normalize(env) self.verbose = verbose self.policy_kwargs = {} if policy_kwargs is None else policy_kwargs self.observation_space = None @@ -373,6 +375,13 @@ class BaseRLModel(object): # Store data in replay buffer if replay_buffer is not None: + # Store only the unnormalized version + if self._vec_normalize_env is not None: + # TODO: save it instead of unnormalizing + obs = self._vec_normalize_env.unnormalize_obs(obs) + new_obs = self._vec_normalize_env.get_original_obs() + reward = self._vec_normalize_env.get_original_reward() + replay_buffer.add(obs, new_obs, action, reward, done_bool) if self.rollout_data is not None: diff --git a/torchy_baselines/common/buffers.py b/torchy_baselines/common/buffers.py index 7ca61b8..2c3d1fb 100644 --- a/torchy_baselines/common/buffers.py +++ b/torchy_baselines/common/buffers.py @@ -1,6 +1,8 @@ import numpy as np import torch as th +from torchy_baselines.common.vec_env import unwrap_vec_normalize + class BaseBuffer(object): def __init__(self, buffer_size, obs_dim, action_dim, device='cpu', n_envs=1): @@ -43,15 +45,26 @@ class BaseBuffer(object): self.pos = 0 self.full = False - def sample(self, batch_size): + def sample(self, batch_size, env=None): upper_bound = self.buffer_size if self.full else self.pos batch_inds = th.LongTensor( np.random.randint(0, upper_bound, size=batch_size)) - return self._get_samples(batch_inds) + return self._get_samples(batch_inds, env=env) - def _get_samples(self, batch_inds): + def _get_samples(self, batch_inds, env=None): raise NotImplementedError() + def _normalize_obs(self, obs, env=None): + if env is not None: + # TODO: get rid of pytorch - numpy conversion + return th.FloatTensor(env.normalize_obs(obs.numpy())) + return obs + + def _normalize_reward(self, reward, env=None): + if env is not None: + return th.FloatTensor(env.normalize_reward(reward.numpy())) + return reward + class ReplayBuffer(BaseBuffer): """ @@ -81,12 +94,12 @@ class ReplayBuffer(BaseBuffer): self.full = True self.pos = 0 - def _get_samples(self, batch_inds): - return (self.observations[batch_inds, 0, :].to(self.device), + def _get_samples(self, batch_inds, env=None): + return (self._normalize_obs(self.observations[batch_inds, 0, :], env).to(self.device), self.actions[batch_inds, 0, :].to(self.device), - self.next_observations[batch_inds, 0, :].to(self.device), + self._normalize_obs(self.next_observations[batch_inds, 0, :], env).to(self.device), self.dones[batch_inds].to(self.device), - self.rewards[batch_inds].to(self.device)) + self._normalize_reward(self.rewards[batch_inds], env).to(self.device)) class RolloutBuffer(BaseBuffer): @@ -184,7 +197,7 @@ class RolloutBuffer(BaseBuffer): yield self._get_samples(indices[start_idx:start_idx + batch_size]) start_idx += batch_size - def _get_samples(self, batch_inds): + def _get_samples(self, batch_inds, env=None): return (self.observations[batch_inds].to(self.device), self.actions[batch_inds].to(self.device), self.values[batch_inds].flatten().to(self.device), diff --git a/torchy_baselines/common/vec_env/__init__.py b/torchy_baselines/common/vec_env/__init__.py index 97f6022..2c542e5 100644 --- a/torchy_baselines/common/vec_env/__init__.py +++ b/torchy_baselines/common/vec_env/__init__.py @@ -1,7 +1,38 @@ # flake8: noqa F401 +from copy import deepcopy + from torchy_baselines.common.vec_env.base_vec_env import AlreadySteppingError, NotSteppingError,\ VecEnv, VecEnvWrapper, CloudpickleWrapper from torchy_baselines.common.vec_env.dummy_vec_env import DummyVecEnv from torchy_baselines.common.vec_env.subproc_vec_env import SubprocVecEnv from torchy_baselines.common.vec_env.vec_frame_stack import VecFrameStack from torchy_baselines.common.vec_env.vec_normalize import VecNormalize + + +def unwrap_vec_normalize(env): + """ + :param env: (gym.Env) + :return: (VecNormalize) + """ + env_tmp = env + while isinstance(env_tmp, VecEnvWrapper): + if isinstance(env_tmp, VecNormalize): + return env_tmp + env_tmp = env_tmp.venv + return None + + +# Define here to avoid circular import +def sync_envs_normalization(env, eval_env): + """ + Sync eval env and train env when using VecNormalize + + :param env: (gym.Env) + :param eval_env: (gym.Env) + """ + env_tmp, eval_env_tmp = env, eval_env + while isinstance(env_tmp, VecEnvWrapper): + if isinstance(env_tmp, VecNormalize): + eval_env_tmp.obs_rms = deepcopy(env_tmp.obs_rms) + env_tmp = env_tmp.venv + eval_env_tmp.venv diff --git a/torchy_baselines/common/vec_env/vec_normalize.py b/torchy_baselines/common/vec_env/vec_normalize.py index 6f2b2b3..0b9797f 100644 --- a/torchy_baselines/common/vec_env/vec_normalize.py +++ b/torchy_baselines/common/vec_env/vec_normalize.py @@ -36,6 +36,7 @@ class VecNormalize(VecEnvWrapper): self.norm_obs = norm_obs self.norm_reward = norm_reward self.old_obs = np.array([]) + self.old_reward = np.array([]) def step_wait(self): """ @@ -46,12 +47,13 @@ class VecNormalize(VecEnvWrapper): """ obs, rews, news, infos = self.venv.step_wait() self.ret = self.ret * self.gamma + rews - self.old_obs = obs + 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 = np.clip(rews / np.sqrt(self.ret_rms.var + self.epsilon), -self.clip_reward, self.clip_reward) + rews = self.normalize_reward(rews) self.ret[news] = 0 return obs, rews, news, infos @@ -62,12 +64,31 @@ class VecNormalize(VecEnvWrapper): if self.norm_obs: if self.training: self.obs_rms.update(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 + return self.normalize_obs(obs) else: return obs + def normalize_obs(self, obs): + 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) + return obs + + def normalize_reward(self, reward): + if self.norm_reward: + return np.clip(reward / np.sqrt(self.ret_rms.var + self.epsilon), -self.clip_reward, self.clip_reward) + return reward + + def unnormalize_obs(self, obs): + if self.norm_obs: + return (obs * np.sqrt(self.obs_rms.var + self.epsilon)) + self.obs_rms.mean + return obs + + def unnormalize_reward(self, reward): + if self.norm_reward: + return reward * np.sqrt(self.ret_rms.var + self.epsilon) + return reward + def get_original_obs(self): """ returns the unnormalized observation @@ -76,6 +97,14 @@ class VecNormalize(VecEnvWrapper): """ return self.old_obs + def get_original_reward(self): + """ + returns the unnormalized observation + + :return: (numpy float) + """ + return self.old_reward + def reset(self): """ Reset all environments diff --git a/torchy_baselines/ppo/ppo.py b/torchy_baselines/ppo/ppo.py index 2b2a5e6..f3a738a 100644 --- a/torchy_baselines/ppo/ppo.py +++ b/torchy_baselines/ppo/ppo.py @@ -1,6 +1,5 @@ import os import time -from copy import deepcopy import gym from gym import spaces @@ -17,7 +16,7 @@ from torchy_baselines.common.base_class import BaseRLModel from torchy_baselines.common.evaluation import evaluate_policy from torchy_baselines.common.buffers import RolloutBuffer from torchy_baselines.common.utils import explained_variance, get_schedule_fn -from torchy_baselines.common.vec_env import VecNormalize, VecEnvWrapper +from torchy_baselines.common.vec_env import sync_envs_normalization from torchy_baselines.common import logger from torchy_baselines.ppo.policies import PPOPolicy @@ -295,14 +294,8 @@ class PPO(BaseRLModel): # Evaluate agent if 0 < eval_freq <= timesteps_since_eval and eval_env is not None: timesteps_since_eval %= eval_freq - # TODO: move that to the base class - # Sync eval env and train env when using VecNormalize - env_tmp, eval_env_tmp = self.env, eval_env - while isinstance(env_tmp, VecEnvWrapper): - if isinstance(env_tmp, VecNormalize): - eval_env_tmp.obs_rms = deepcopy(env_tmp.obs_rms) - env_tmp = env_tmp.venv - eval_env_tmp.venv + sync_envs_normalization(self.env, eval_env) + mean_reward, _ = evaluate_policy(self, eval_env, n_eval_episodes) if self.tb_writer is not None: self.tb_writer.add_scalar('Eval/reward', mean_reward, self.num_timesteps) diff --git a/torchy_baselines/sac/sac.py b/torchy_baselines/sac/sac.py index bad470a..f4c5c34 100644 --- a/torchy_baselines/sac/sac.py +++ b/torchy_baselines/sac/sac.py @@ -8,6 +8,7 @@ from torchy_baselines.common.base_class import BaseRLModel from torchy_baselines.common.buffers import ReplayBuffer from torchy_baselines.common.evaluation import evaluate_policy from torchy_baselines.sac.policies import SACPolicy +from torchy_baselines.common.vec_env import sync_envs_normalization class SAC(BaseRLModel): @@ -162,7 +163,7 @@ class SAC(BaseRLModel): for gradient_step in range(gradient_steps): # Sample replay buffer - replay_data = self.replay_buffer.sample(batch_size) + replay_data = self.replay_buffer.sample(batch_size, env=self._vec_normalize_env) obs, action_batch, next_obs, done, reward = replay_data @@ -266,6 +267,7 @@ class SAC(BaseRLModel): # Evaluate episode if 0 < eval_freq <= timesteps_since_eval and eval_env is not None: timesteps_since_eval %= eval_freq + sync_envs_normalization(self.env, eval_env) mean_reward, _ = evaluate_policy(self, eval_env, n_eval_episodes) evaluations.append(mean_reward) if self.verbose > 0: diff --git a/torchy_baselines/td3/td3.py b/torchy_baselines/td3/td3.py index ce6b9a4..6c1dbcf 100644 --- a/torchy_baselines/td3/td3.py +++ b/torchy_baselines/td3/td3.py @@ -8,6 +8,7 @@ from torchy_baselines.common.base_class import BaseRLModel from torchy_baselines.common.buffers import ReplayBuffer from torchy_baselines.common.evaluation import evaluate_policy from torchy_baselines.td3.policies import TD3Policy +from torchy_baselines.common.vec_env import sync_envs_normalization class TD3(BaseRLModel): @@ -123,7 +124,7 @@ class TD3(BaseRLModel): for gradient_step in range(gradient_steps): # Sample replay buffer if replay_data is None: - obs, action, next_obs, done, reward = self.replay_buffer.sample(batch_size) + obs, action, next_obs, done, reward = self.replay_buffer.sample(batch_size, env=self._vec_normalize_env) else: obs, action, next_obs, done, reward = replay_data @@ -162,7 +163,7 @@ class TD3(BaseRLModel): for gradient_step in range(gradient_steps): # Sample replay buffer if replay_data is None: - obs, _, next_obs, done, reward = self.replay_buffer.sample(batch_size) + obs, _, next_obs, done, reward = self.replay_buffer.sample(batch_size, env=self._vec_normalize_env) else: obs, _, next_obs, done, reward = replay_data @@ -187,7 +188,7 @@ class TD3(BaseRLModel): for gradient_step in range(gradient_steps): # Sample replay buffer - replay_data = self.replay_buffer.sample(batch_size) + replay_data = self.replay_buffer.sample(batch_size, env=self._vec_normalize_env) self.train_critic(replay_data=replay_data) # Delayed policy updates @@ -278,6 +279,7 @@ class TD3(BaseRLModel): # Evaluate episode if 0 < eval_freq <= timesteps_since_eval and eval_env is not None: timesteps_since_eval %= eval_freq + sync_envs_normalization(self.env, eval_env) mean_reward, _ = evaluate_policy(self, eval_env, n_eval_episodes) evaluations.append(mean_reward) if self.verbose > 0: