Make options env compatible with PPO

This commit is contained in:
ebuehrle
2021-10-29 16:36:12 +02:00
parent fc2cd936a8
commit 66c10f5280
2 changed files with 96 additions and 125 deletions

View File

@@ -20,8 +20,9 @@ from imitation.util import logger
from stable_baselines3.common.env_util import make_vec_env
from tqdm import tqdm
from src.policies.options import OptionsCnnPolicy
from src.gail.options import OptionsEnv, LLOptions, HLOptions, RenderOptions
from src.gail.train import train_discriminator, train_generator
from src.gail.train import flatten_transitions
from gail.options2 import OptionsEnv, RenderOptions, imitation_discriminator
model_name = 'gail_options_image_random_location'
env_settings = {'width': 70, 'height': 70, 'm_per_px': 1, 'map_color': 128, 'mu': 0.001}
@@ -30,12 +31,13 @@ ALL_OPTIONS = [(v,t) for v in [0,2,4,6,8] for t in [5, 10, 20]] # option 0 is sa
def train(
expert_data,
epochs=200,
expert_batch_size=256,
generator_steps=1024,
discount=0.99,
n_disc_updates_per_round=2,
generator_steps=512,
generator_total_steps=2048,
n_gen_updates_per_round=10,
discount=0.99,
epochs=100,
):
env = NRasterizedRouteRandomAgentLocation(**env_settings)
env.discount = discount
@@ -55,24 +57,25 @@ def train(
gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused
)
options_env = OptionsEnv(env, discriminator=imitation_discriminator(discriminator), options=ALL_OPTIONS, ll_buffer_capacity=expert_batch_size)
generator = stable_baselines3.PPO(
OptionsCnnPolicy,
OptionsEnv(env, options=ALL_OPTIONS),
options_env,
verbose=1,
n_steps=generator_steps,
n_epochs=n_gen_updates_per_round,
)
# PPO.train requires logger as set up in
# PPO._setup_learn (called by PPO.learn)
generator._logger = stable_baselines3.common.utils.configure_logger(
generator.verbose,
generator.tensorboard_log,
)
for _ in tqdm(range(epochs)):
train_discriminator(LLOptions(env, options=ALL_OPTIONS), generator, discriminator, num_samples=expert_batch_size, n_updates=n_disc_updates_per_round)
train_generator(HLOptions(env, options=ALL_OPTIONS), generator, discriminator, num_samples=generator_steps)
# train generator
generator.learn(total_timesteps=generator_total_steps)
# train discriminator
generator_samples = options_env.sample_ll(expert_batch_size)
generator_samples = flatten_transitions(generator_samples)
for _ in range(n_disc_updates_per_round):
discriminator.train_disc(gen_samples=generator_samples)
generator.save(model_name)
return generator