Testing off policy normalization

This commit is contained in:
Antonin Raffin 2019-11-14 14:35:00 +01:00
parent 8aac10f3fa
commit 5278a6f3f8
10 changed files with 125 additions and 29 deletions

View file

@ -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

View file

@ -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)

View file

@ -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)

View file

@ -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:

View file

@ -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),

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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:

View file

@ -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: