Vanilla GAIL on rasterized observation

This commit is contained in:
ebuehrle
2021-09-08 20:28:18 +02:00
parent 802d4a4301
commit 50916aec05
5 changed files with 14 additions and 13 deletions

View File

@@ -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}' --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.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'

View File

@@ -43,12 +43,12 @@ logger.configure(tempdir_path / "GAIL/")
gail_trainer = adversarial.GAIL( gail_trainer = adversarial.GAIL(
venv, venv,
expert_data=transitions, expert_data=transitions,
expert_batch_size=220, expert_batch_size=150,
#n_disc_updates_per_round=32, n_disc_updates_per_round=32,
discrim_kwargs={'discrim_net': MlpDiscriminator()}, 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) gail_trainer.gen_algo.save(model_name)
#del gail_trainer #del gail_trainer

View File

@@ -10,7 +10,7 @@ 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 NRasterized from intersimple.intersimple import NRasterized
from gail.discriminator import CnnDiscriminator from gail.discriminator import CnnDiscriminator
@@ -18,7 +18,7 @@ model_name = 'gail_image'
# %% # %%
# Load pickled test demonstrations. # 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 # This is a list of `imitation.data.types.Trajectory`, where
# every instance contains observations and actions for a single expert # every instance contains observations and actions for a single expert
# demonstration. # demonstration.
@@ -30,7 +30,7 @@ with open("data/NormalizedIntersimpleExpert_NRasterizedAgent51.pkl", "rb") as f:
# (observation, actions, next_observation) transitions. # (observation, actions, next_observation) transitions.
transitions = rollout.flatten_trajectories(trajectories) 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 = tempfile.TemporaryDirectory(prefix="quickstart")
tempdir_path = pathlib.Path(tempdir.name) tempdir_path = pathlib.Path(tempdir.name)
@@ -43,10 +43,10 @@ logger.configure(tempdir_path / "GAIL/")
gail_trainer = adversarial.GAIL( gail_trainer = adversarial.GAIL(
venv, venv,
expert_data=transitions, expert_data=transitions,
expert_batch_size=200, 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=128), gen_algo=sb3.PPO("CnnPolicy", venv, verbose=1, n_steps=1024),
) )
gail_trainer.train(total_timesteps=100000) gail_trainer.train(total_timesteps=100000)
gail_trainer.gen_algo.save(model_name) gail_trainer.gen_algo.save(model_name)
@@ -56,7 +56,7 @@ gail_trainer.gen_algo.save(model_name)
# %% # %%
model = sb3.PPO.load(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() obs = env.reset()
while True: while True:

Binary file not shown.