Vanilla GAIL on rasterized observation
This commit is contained in:
@@ -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'
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -66,4 +66,4 @@ while True:
|
|||||||
if done:
|
if done:
|
||||||
break
|
break
|
||||||
|
|
||||||
env.close(filestr='render/'+model_name)
|
env.close(filestr='render/'+model_name)
|
||||||
@@ -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:
|
||||||
|
|||||||
BIN
scratch/etienne/intersimple/render/gail_image_ani.mp4
Normal file
BIN
scratch/etienne/intersimple/render/gail_image_ani.mp4
Normal file
Binary file not shown.
BIN
scratch/etienne/intersimple/render/gail_image_observation.mp4
Normal file
BIN
scratch/etienne/intersimple/render/gail_image_observation.mp4
Normal file
Binary file not shown.
Reference in New Issue
Block a user