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.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