From 70e55327dcb60359e9d23bb536e9f5f392f905c9 Mon Sep 17 00:00:00 2001 From: Arec Date: Mon, 11 Oct 2021 08:47:18 -0700 Subject: [PATCH] fixing flataction discriminator to convert to float beforehand, adding necessary forward calls in expert, adding Fire to video creator from model, and trying full run of options gail with new discrimination model --- scratch/etienne/intersimple/data/expert.py | 6 +++ .../etienne/intersimple/gail/discriminator.py | 2 +- .../etienne/intersimple/gail_options_image.py | 8 ++-- .../intersimple/render_env_from_model.py | 38 ++++++++++++------- 4 files changed, 37 insertions(+), 17 deletions(-) diff --git a/scratch/etienne/intersimple/data/expert.py b/scratch/etienne/intersimple/data/expert.py index a8edbaf..56105af 100644 --- a/scratch/etienne/intersimple/data/expert.py +++ b/scratch/etienne/intersimple/data/expert.py @@ -79,6 +79,12 @@ class DummyVecEnvPolicy(BasePolicy): actions = [p[0] for p in predictions] states = [p[1] for p in predictions] return actions, states + + def forward(self, *args, **kwargs): + raise NotImplementedError() + + def _predict(self, *args, **kwargs): + raise NotImplementedError() def save_video(env, expert): env.reset() diff --git a/scratch/etienne/intersimple/gail/discriminator.py b/scratch/etienne/intersimple/gail/discriminator.py index 02cb83a..1c9664c 100644 --- a/scratch/etienne/intersimple/gail/discriminator.py +++ b/scratch/etienne/intersimple/gail/discriminator.py @@ -75,7 +75,7 @@ class CnnDiscriminatorFlatAction(torch.nn.Module): return sa def forward(self, state, action): - s = self.cnn(state) + s = self.cnn(state.float()) sa = self._concatenate(s, action) assert sa.ndim == 2 return self.decoder(sa).squeeze(1) diff --git a/scratch/etienne/intersimple/gail_options_image.py b/scratch/etienne/intersimple/gail_options_image.py index ae70490..b35eb3f 100644 --- a/scratch/etienne/intersimple/gail_options_image.py +++ b/scratch/etienne/intersimple/gail_options_image.py @@ -1,5 +1,5 @@ # %% -from gail.discriminator import CnnDiscriminator +from gail.discriminator import CnnDiscriminator, CnnDiscriminatorFlatAction from imitation.algorithms import adversarial import stable_baselines3 import torch.utils.data @@ -271,7 +271,8 @@ def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048, di discriminator = adversarial.GAIL( expert_data=expert_data, expert_batch_size=expert_batch_size, - discrim_kwargs={'discrim_net': CnnDiscriminator(venv)}, + discrim_kwargs={'discrim_net': CnnDiscriminatorFlatAction(venv)}, + #discrim_kwargs={'discrim_net': CnnDiscriminator(venv)}, venv=venv, # unused gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused ) @@ -299,10 +300,11 @@ def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048, di # %% if __name__ == '__main__': # %% + with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f: trajectories = pickle.load(f) transitions = rollout.flatten_trajectories(trajectories) - generator = train(transitions, epochs=2, expert_batch_size=2, generator_steps=2) + generator = train(transitions) generator.save(model_name) diff --git a/scratch/etienne/intersimple/render_env_from_model.py b/scratch/etienne/intersimple/render_env_from_model.py index f44a670..0214667 100644 --- a/scratch/etienne/intersimple/render_env_from_model.py +++ b/scratch/etienne/intersimple/render_env_from_model.py @@ -3,19 +3,31 @@ import stable_baselines3 as sb3 from intersim.envs.intersimple import NRasterized -model_name = 'gail_image_singleagent_nocollision' -model = sb3.PPO.load(model_name) +def render_env(model_name='gail_image_multiagent_nocollision', agent=51, environment=NRasterized): + """ + Render a video from an model, agent, and environment + Args: + model_name (str): name of the model + agent (int): agent to start the video from + environment (gym.Env): gym environment class to render environment on + """ -env = NRasterized(stop_on_collision=False, width=36, height=36, m_per_px=2, agent=51) + model = sb3.PPO.load(model_name) -obs = env.reset() -i=0 -while True and i < 600: - i+=1 - action, _states = model.predict(obs) - obs, rewards, done, info = env.step(action) - env.render(mode='post') - if done: - break + env = environment(stop_on_collision=False, width=36, height=36, m_per_px=2, agent=agent) -env.close(filestr='render/'+model_name) \ No newline at end of file + obs = env.reset() + i=0 + while True and i < 600: + i+=1 + action, _states = model.predict(obs) + obs, rewards, done, info = env.step(action) + env.render(mode='post') + if done: + break + + env.close(filestr='render/'+model_name+'_agent%i'%(agent)) + +if __name__ == '__main__': + import fire + fire.Fire(render_env) \ No newline at end of file