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