mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-24 19:43:50 +00:00
Testing off policy normalization
This commit is contained in:
parent
8aac10f3fa
commit
5278a6f3f8
10 changed files with 125 additions and 29 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Reference in a new issue