diff --git a/scratch/etienne/intersimple/gail_options_image_random.py b/scratch/etienne/intersimple/gail_options_image_random.py new file mode 100644 index 0000000..ac67a4e --- /dev/null +++ b/scratch/etienne/intersimple/gail_options_image_random.py @@ -0,0 +1,87 @@ +# %% +from gail.discriminator import CnnDiscriminatorFlatAction +from imitation.algorithms import adversarial +import stable_baselines3 +import torch.utils.data +import numpy as np +from intersim.envs.intersimple import NRasterizedRandomAgent +import itertools +from torch.distributions import Categorical +import gym +import torch +import pickle +import imitation.data.rollout as rollout +import tempfile +import pathlib +from imitation.util import logger +from stable_baselines3.common.env_util import make_vec_env +from tqdm import tqdm +from gail.policy import OptionsCnnPolicy +from gail.options import OptionsEnv, LLOptions, HLOptions, RenderOptions +from gail.train import train_discriminator, train_generator + +model_name = 'gail_options_image_random' +env_settings = {'width': 36, 'height': 36, 'm_per_px': 2} + +ALL_OPTIONS = [(v,t) for v in [0,2,4,6,8] for t in [5, 10, 20]] # option 0 is safe fallback + +def train(expert_data, epochs=20, expert_batch_size=32, generator_steps=1024, discount=0.99): + env = NRasterizedRandomAgent(**env_settings) + env.discount = discount + + tempdir = tempfile.TemporaryDirectory(prefix="quickstart") + tempdir_path = pathlib.Path(tempdir.name) + logger.configure(tempdir_path / "GAIL/") + print(f"All Tensorboards and logging are being written inside {tempdir_path}/.") + + venv = make_vec_env(NRasterizedRandomAgent, n_envs=1, env_kwargs=env_settings) + discriminator = adversarial.GAIL( + expert_data=expert_data, + expert_batch_size=expert_batch_size, + discrim_kwargs={'discrim_net': CnnDiscriminatorFlatAction(venv)}, + #discrim_kwargs={'discrim_net': CnnDiscriminator(venv)}, + venv=venv, # unused + gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused + ) + + generator = stable_baselines3.PPO( + OptionsCnnPolicy, + OptionsEnv(env, options=ALL_OPTIONS), + verbose=1, + n_steps=generator_steps, + ) + + # PPO.train requires logger as set up in + # PPO._setup_learn (called by PPO.learn) + generator._logger = stable_baselines3.common.utils.configure_logger( + generator.verbose, + generator.tensorboard_log, + ) + + for _ in tqdm(range(epochs)): + train_discriminator(LLOptions(env, options=ALL_OPTIONS), generator, discriminator, num_samples=expert_batch_size) + train_generator(HLOptions(env, options=ALL_OPTIONS), generator, discriminator, num_samples=generator_steps) + + return generator + +# %% +if __name__ == '__main__': + # %% + + with open("data/NormalizedIntersimpleExpertMu.001N10000_NRasterizedRandomAgentw36h36mppx2.pkl", "rb") as f: + trajectories = pickle.load(f) + transitions = rollout.flatten_trajectories(trajectories) + generator = train(transitions, epochs=100) + + generator.save(model_name) + + # %% + model = stable_baselines3.PPO.load(model_name) + + env = RenderOptions(NRasterizedRandomAgent(**env_settings)) + + for s in env.sample_ll(model): + if s['dones']: + break + + env.close(filestr='render/'+model_name)