Add callback to gail image random to report metrics

Currently nothing is reported yet
This commit is contained in:
Johannes Fischer
2021-10-08 18:18:42 +02:00
parent 4b9a81080b
commit 6d2ab54b6e

View File

@@ -10,7 +10,9 @@ from imitation.algorithms import adversarial, bc
from imitation.data import rollout from imitation.data import rollout
from imitation.util import logger 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 from gail.discriminator import CnnDiscriminator
@@ -30,7 +32,8 @@ with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedRandomAgentw36h36mp
# (observation, actions, next_observation) transitions. # (observation, actions, next_observation) transitions.
transitions = rollout.flatten_trajectories(trajectories) 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 = tempfile.TemporaryDirectory(prefix="quickstart")
tempdir_path = pathlib.Path(tempdir.name) 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 # GAIL, and AIRL also accept as `expert_data` any Pytorch-style DataLoader that
# iterates over dictionaries containing observations, actions, and next_observations. # iterates over dictionaries containing observations, actions, and next_observations.
logger.configure(tempdir_path / "GAIL/") logger.configure(tempdir_path / "GAIL/")
generator = sb3.PPO("CnnPolicy", venv, verbose=1, n_steps=1024)
gail_trainer = adversarial.GAIL( gail_trainer = adversarial.GAIL(
venv, venv,
expert_data=transitions, expert_data=transitions,
expert_batch_size=32, expert_batch_size=32,
#n_disc_updates_per_round=2048, #n_disc_updates_per_round=2048,
discrim_kwargs={'discrim_net': CnnDiscriminator(venv)}, discrim_kwargs={'discrim_net': CnnDiscriminator(venv)},
gen_algo=sb3.PPO("CnnPolicy", venv, verbose=1, n_steps=1024), gen_algo=generator,
allow_variable_horizon=True, 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) gail_trainer.gen_algo.save(model_name)
#del gail_trainer #del gail_trainer