From 2aaaad36f0b24ab1a5964a21165a33885aa18e96 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Tue, 2 Nov 2021 16:15:34 +0100 Subject: [PATCH] Fix rendering - support predict() - move discount to OptionsEnv - fix RenderOptions --- scratch/etienne/intersimple/gail/options2.py | 11 +++++----- .../gail_options_image_random_location.py | 21 ++++++++++++------- src/policies/options.py | 17 +++++++++------ 3 files changed, 30 insertions(+), 19 deletions(-) diff --git a/scratch/etienne/intersimple/gail/options2.py b/scratch/etienne/intersimple/gail/options2.py index d4ae929..6d3728e 100644 --- a/scratch/etienne/intersimple/gail/options2.py +++ b/scratch/etienne/intersimple/gail/options2.py @@ -15,7 +15,7 @@ def imitation_discriminator(discriminator): class OptionsEnv(gym.Wrapper): - def __init__(self, env, options, discriminator, ll_buffer_capacity, *args, **kwargs): + def __init__(self, env, options, discriminator, discount, ll_buffer_capacity, *args, **kwargs): super().__init__(env, *args, **kwargs) self.options = options @@ -27,6 +27,7 @@ class OptionsEnv(gym.Wrapper): }) self.discriminator = discriminator + self.discount = discount self.ll_buffer_capacity = ll_buffer_capacity self.ll_buffer = deque(maxlen=ll_buffer_capacity) @@ -85,11 +86,11 @@ class OptionsEnv(gym.Wrapper): class RenderOptions(OptionsEnv): - def __init__(self, options, *args, **kwargs): - super().__init__(options, discriminator=lambda s, a, n, d: 0, ll_buffer_capacity=0, *args, **kwargs) + def __init__(self, env, options, *args, **kwargs): + super().__init__(env, options, discriminator=lambda s, a, n, d: 0, discount=1, ll_buffer_capacity=0, *args, **kwargs) - def _ll_step(self): - out = super()._ll_step() + def _ll_step(self, action): + out = super()._ll_step(action) self.env.render() return out diff --git a/scratch/etienne/intersimple/gail_options_image_random_location.py b/scratch/etienne/intersimple/gail_options_image_random_location.py index f1d550d..501ee54 100644 --- a/scratch/etienne/intersimple/gail_options_image_random_location.py +++ b/scratch/etienne/intersimple/gail_options_image_random_location.py @@ -34,7 +34,6 @@ def train( epochs=200, ): env = NRasterizedRouteSpeedRandomAgentLocation(**env_settings) - env.discount = discount tempdir = tempfile.TemporaryDirectory(prefix="quickstart") tempdir_path = pathlib.Path(tempdir.name) @@ -51,10 +50,12 @@ def train( gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused ) - options_env = TimeLimit(OptionsEnv(env, - discriminator=imitation_discriminator(discriminator), + options_env = TimeLimit(OptionsEnv( + env, options=ALL_OPTIONS, - ll_buffer_capacity=expert_batch_size + discriminator=imitation_discriminator(discriminator), + discount=discount, + ll_buffer_capacity=expert_batch_size, ), max_episode_steps=10) generator = stable_baselines3.PPO( OptionsCnnPolicy, @@ -79,11 +80,15 @@ def train( return generator def video(model_name, env): - model = stable_baselines3.PPO.load(model_name) env = RenderOptions(env, options=ALL_OPTIONS) - for s in env.sample_ll(model): - if s['dones']: - break + model = stable_baselines3.PPO.load(model_name) + + done = False + obs = env.reset() + while not done: + action, _ = model.predict(obs) + obs, _, done, _ = env.step(action) + env.close(filestr='render/'+model_name) def evaluate(): diff --git a/src/policies/options.py b/src/policies/options.py index 851edd4..44e9d77 100644 --- a/src/policies/options.py +++ b/src/policies/options.py @@ -1,12 +1,13 @@ -import stable_baselines3 +from stable_baselines3.common.policies import ActorCriticPolicy, ActorCriticCnnPolicy from torch.distributions import Categorical -class OptionsCnnPolicy(stable_baselines3.common.policies.ActorCriticCnnPolicy): +class OptionsCnnPolicy(ActorCriticPolicy): """ Class for high-level options policy (generator) """ def __init__(self, observation_space, *args, **kwargs): - super().__init__(observation_space['obs'], *args, **kwargs) + super().__init__(observation_space, *args, **kwargs) + self.cnn_policy = ActorCriticCnnPolicy(observation_space['obs'], *args, **kwargs) def _prior_distribution(self, s): """ @@ -17,9 +18,9 @@ class OptionsCnnPolicy(stable_baselines3.common.policies.ActorCriticCnnPolicy): values (torch.tensor): values from critic dist (torch.distributions): prior distribution over actions """ - latent_pi, latent_vf, latent_sde = self._get_latent(s) - distribution = self._get_action_dist_from_latent(latent_pi, latent_sde) - values = self.value_net(latent_vf) + latent_pi, latent_vf, latent_sde = self.cnn_policy._get_latent(s) + distribution = self.cnn_policy._get_action_dist_from_latent(latent_pi, latent_sde) + values = self.cnn_policy.value_net(latent_vf) return values, distribution.distribution def forward(self, obs, eps=1e-6): @@ -39,6 +40,10 @@ class OptionsCnnPolicy(stable_baselines3.common.policies.ActorCriticCnnPolicy): posterior = Categorical((prior.probs + eps) * m) ch = posterior.sample() return ch, values, posterior.log_prob(ch) + + def _predict(self, obs, deterministic=False): + action, _, _ = self.forward(obs) + return action def evaluate_actions(self, obs, ch, eps=1e-6): """