Add callback to gail image random to report metrics
Currently nothing is reported yet
This commit is contained in:
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user