From 6d2ab54b6e6ed5395fd997872cbe1ceeec0f9542 Mon Sep 17 00:00:00 2001 From: Johannes Fischer Date: Fri, 8 Oct 2021 18:18:42 +0200 Subject: [PATCH] Add callback to gail image random to report metrics Currently nothing is reported yet --- .../etienne/intersimple/gail_image_random.py | 17 +++++++++++++---- 1 file changed, 13 insertions(+), 4 deletions(-) diff --git a/scratch/etienne/intersimple/gail_image_random.py b/scratch/etienne/intersimple/gail_image_random.py index bff33c0..79794e9 100644 --- a/scratch/etienne/intersimple/gail_image_random.py +++ b/scratch/etienne/intersimple/gail_image_random.py @@ -10,7 +10,9 @@ from imitation.algorithms import adversarial, bc from imitation.data import rollout from imitation.util import logger -from intersim.envs.intersimple import NRasterizedRandomAgent +from intersim.envs.intersimple import NRasterizedRandomAgent, IntersimpleReward, speed_reward +import functools +from stable_baselines3.common.evaluation import evaluate_policy from gail.discriminator import CnnDiscriminator @@ -30,7 +32,8 @@ with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedRandomAgentw36h36mp # (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}) +env_kwargs = {'width': 36, 'height': 36, 'm_per_px': 2} +venv = make_vec_env(NRasterizedRandomAgent, n_envs=2, env_kwargs=env_kwargs) tempdir = tempfile.TemporaryDirectory(prefix="quickstart") tempdir_path = pathlib.Path(tempdir.name) @@ -40,16 +43,22 @@ print(f"All Tensorboards and logging are being written inside {tempdir_path}/.") # 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/") +generator = sb3.PPO("CnnPolicy", venv, verbose=1, n_steps=1024) 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), + gen_algo=generator, allow_variable_horizon=True, ) -gail_trainer.train(total_timesteps=100000) +def callback(round): + eval_env = NRasterizedRandomAgent(reward=functools.partial(speed_reward, collision_penalty=0.), **env_kwargs) + #sync_envs_normalization(self.training_env, self.eval_env) + episode_rewards, episode_lengths = evaluate_policy(generator, eval_env, return_episode_rewards=True) + +gail_trainer.train(total_timesteps=100000, callback=callback) gail_trainer.gen_algo.save(model_name) #del gail_trainer