From 2280597db6d9b85fc144062601fcd8448a80b5ad Mon Sep 17 00:00:00 2001 From: Johannes Fischer Date: Mon, 13 Sep 2021 19:42:41 +0200 Subject: [PATCH] Add GAIL with random agent data --- .../etienne/intersimple/gail_image_random.py | 70 +++++++++++++++++++ 1 file changed, 70 insertions(+) create mode 100644 scratch/etienne/intersimple/gail_image_random.py diff --git a/scratch/etienne/intersimple/gail_image_random.py b/scratch/etienne/intersimple/gail_image_random.py new file mode 100644 index 0000000..bff33c0 --- /dev/null +++ b/scratch/etienne/intersimple/gail_image_random.py @@ -0,0 +1,70 @@ +# %% +import pathlib +import pickle +import tempfile + +import stable_baselines3 as sb3 +from stable_baselines3.common.env_util import make_vec_env + +from imitation.algorithms import adversarial, bc +from imitation.data import rollout +from imitation.util import logger + +from intersim.envs.intersimple import NRasterizedRandomAgent + +from gail.discriminator import CnnDiscriminator + +model_name = 'gail_image_random' + +# %% +# Load pickled test demonstrations. +with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedRandomAgentw36h36mppx2.pkl", "rb") as f: + # This is a list of `imitation.data.types.Trajectory`, where + # every instance contains observations and actions for a single expert + # demonstration. + trajectories = pickle.load(f) + +# %% +# Convert List[types.Trajectory] to an instance of `imitation.data.types.Transitions`. +# This is a more general dataclass containing unordered +# (observation, actions, next_observation) transitions. +transitions = rollout.flatten_trajectories(trajectories) + +venv = make_vec_env(NRasterizedRandomAgent, n_envs=2, env_kwargs={'width': 36, 'height': 36, 'm_per_px': 2}) + +tempdir = tempfile.TemporaryDirectory(prefix="quickstart") +tempdir_path = pathlib.Path(tempdir.name) +print(f"All Tensorboards and logging are being written inside {tempdir_path}/.") + +# Train GAIL on expert data. +# GAIL, and AIRL also accept as `expert_data` any Pytorch-style DataLoader that +# iterates over dictionaries containing observations, actions, and next_observations. +logger.configure(tempdir_path / "GAIL/") +gail_trainer = adversarial.GAIL( + venv, + expert_data=transitions, + expert_batch_size=32, + #n_disc_updates_per_round=2048, + discrim_kwargs={'discrim_net': CnnDiscriminator(venv)}, + gen_algo=sb3.PPO("CnnPolicy", venv, verbose=1, n_steps=1024), + allow_variable_horizon=True, +) +gail_trainer.train(total_timesteps=100000) +gail_trainer.gen_algo.save(model_name) + +#del gail_trainer + +# %% +model = sb3.PPO.load(model_name) + +env = NRasterizedRandomAgent(width=36, height=36, m_per_px=2) + +obs = env.reset() +while True: + 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)