Allow variable horizon trajectories

This commit is contained in:
Johannes Fischer
2021-09-10 11:34:36 +02:00
parent b2b2abafa2
commit 544ea4d15a
3 changed files with 3 additions and 2 deletions

View File

@@ -47,6 +47,7 @@ gail_trainer = adversarial.GAIL(
n_disc_updates_per_round=32,
discrim_kwargs={'discrim_net': MlpDiscriminator()},
gen_algo=sb3.PPO("MlpPolicy", venv, verbose=1, n_steps=4530),
allow_variable_horizon=True,
)
gail_trainer.train(total_timesteps=400000)
gail_trainer.gen_algo.save(model_name)

View File

@@ -47,6 +47,7 @@ gail_trainer = adversarial.GAIL(
#n_disc_updates_per_round=2048,
discrim_kwargs={'discrim_net': CnnDiscriminator(venv)},
gen_algo=sb3.PPO("CnnPolicy", venv, verbose=1, n_steps=1024),
allow_variable_horizon=True,
)
gail_trainer.train(total_timesteps=100000)
gail_trainer.gen_algo.save(model_name)

View File

@@ -47,8 +47,7 @@ def target_velocity_plan(current_v: float, target_v: float, t: int)
# for now, constant acceleration
a = (target_v - current_v) / t
return a*np.ones((t,))
Hello
def generate_plan(env, i):
"""Generate input profile for high-level action `i`."""
assert i < len(ALL_OPTIONS), "Invalid option index {i}"