diff --git a/scratch/etienne/intersimple/data/generate.sh b/scratch/etienne/intersimple/data/generate.sh index 96c4fff..ef44050 100644 --- a/scratch/etienne/intersimple/data/generate.sh +++ b/scratch/etienne/intersimple/data/generate.sh @@ -1,3 +1,4 @@ #python -m intersimple.expert --env=IntersimpleReward --min_timesteps=200 --env_args='{agent:51}' --path='NormalizedIntersimpleExpert_IntersimpleRewardAgent51.pkl' #python -m intersimple.expert --env=IntersimpleReward --min_timesteps=200 --env_args='{agent:51}' --policy_args='{mu:0.005}' --path='NormalizedIntersimpleExpert_IntersimpleRewardAgent51Mu.005.pkl' -python -m intersimple.expert --env=IntersimpleReward --min_timesteps=200 --env_args='{agent:51}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpert_IntersimpleRewardAgent51Mu.001.pkl' --video +#python -m intersimple.expert --env=IntersimpleReward --min_timesteps=200 --env_args='{agent:51}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpert_IntersimpleRewardAgent51Mu.001.pkl' +python -m intersimple.expert --env=NRasterized --min_timesteps=200 --env_args='{agent:51,width:36,height:36,m_per_px:2}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl' diff --git a/scratch/etienne/intersimple/gail_flat.py b/scratch/etienne/intersimple/gail_flat.py index 86af50e..3996649 100644 --- a/scratch/etienne/intersimple/gail_flat.py +++ b/scratch/etienne/intersimple/gail_flat.py @@ -43,12 +43,12 @@ logger.configure(tempdir_path / "GAIL/") gail_trainer = adversarial.GAIL( venv, expert_data=transitions, - expert_batch_size=220, - #n_disc_updates_per_round=32, + expert_batch_size=150, + n_disc_updates_per_round=32, discrim_kwargs={'discrim_net': MlpDiscriminator()}, - gen_algo=sb3.PPO("MlpPolicy", venv, verbose=1, n_steps=4096), + gen_algo=sb3.PPO("MlpPolicy", venv, verbose=1, n_steps=4530), ) -gail_trainer.train(total_timesteps=80000) +gail_trainer.train(total_timesteps=400000) gail_trainer.gen_algo.save(model_name) #del gail_trainer @@ -66,4 +66,4 @@ while True: if done: break -env.close(filestr='render/'+model_name) +env.close(filestr='render/'+model_name) \ No newline at end of file diff --git a/scratch/etienne/intersimple/gail_image.py b/scratch/etienne/intersimple/gail_image.py index 1f0ba6a..2054ae6 100644 --- a/scratch/etienne/intersimple/gail_image.py +++ b/scratch/etienne/intersimple/gail_image.py @@ -10,7 +10,7 @@ from imitation.algorithms import adversarial, bc from imitation.data import rollout from imitation.util import logger -from intersim.envs.intersimple import NRasterized +from intersimple.intersimple import NRasterized from gail.discriminator import CnnDiscriminator @@ -18,7 +18,7 @@ model_name = 'gail_image' # %% # Load pickled test demonstrations. -with open("data/NormalizedIntersimpleExpert_NRasterizedAgent51.pkl", "rb") as f: +with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f: # This is a list of `imitation.data.types.Trajectory`, where # every instance contains observations and actions for a single expert # demonstration. @@ -30,7 +30,7 @@ with open("data/NormalizedIntersimpleExpert_NRasterizedAgent51.pkl", "rb") as f: # (observation, actions, next_observation) transitions. transitions = rollout.flatten_trajectories(trajectories) -venv = make_vec_env(NRasterized, n_envs=2, env_kwargs={'agent': 51}) +venv = make_vec_env(NRasterized, n_envs=2, env_kwargs={'agent': 51, 'width': 36, 'height': 36, 'm_per_px': 2}) tempdir = tempfile.TemporaryDirectory(prefix="quickstart") tempdir_path = pathlib.Path(tempdir.name) @@ -43,10 +43,10 @@ logger.configure(tempdir_path / "GAIL/") gail_trainer = adversarial.GAIL( venv, expert_data=transitions, - expert_batch_size=200, - n_disc_updates_per_round=2048, + 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=128), + gen_algo=sb3.PPO("CnnPolicy", venv, verbose=1, n_steps=1024), ) gail_trainer.train(total_timesteps=100000) gail_trainer.gen_algo.save(model_name) @@ -56,7 +56,7 @@ gail_trainer.gen_algo.save(model_name) # %% model = sb3.PPO.load(model_name) -env = NRasterized(agent=51) +env = NRasterized(agent=51, width=36, height=36, m_per_px=2) obs = env.reset() while True: diff --git a/scratch/etienne/intersimple/render/gail_image_ani.mp4 b/scratch/etienne/intersimple/render/gail_image_ani.mp4 new file mode 100644 index 0000000..1816bbb Binary files /dev/null and b/scratch/etienne/intersimple/render/gail_image_ani.mp4 differ diff --git a/scratch/etienne/intersimple/render/gail_image_observation.mp4 b/scratch/etienne/intersimple/render/gail_image_observation.mp4 new file mode 100644 index 0000000..52719da Binary files /dev/null and b/scratch/etienne/intersimple/render/gail_image_observation.mp4 differ