Allow variable horizon trajectories
This commit is contained in:
@@ -47,6 +47,7 @@ gail_trainer = adversarial.GAIL(
|
|||||||
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=4530),
|
gen_algo=sb3.PPO("MlpPolicy", venv, verbose=1, n_steps=4530),
|
||||||
|
allow_variable_horizon=True,
|
||||||
)
|
)
|
||||||
gail_trainer.train(total_timesteps=400000)
|
gail_trainer.train(total_timesteps=400000)
|
||||||
gail_trainer.gen_algo.save(model_name)
|
gail_trainer.gen_algo.save(model_name)
|
||||||
|
|||||||
@@ -47,6 +47,7 @@ gail_trainer = adversarial.GAIL(
|
|||||||
#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=1024),
|
gen_algo=sb3.PPO("CnnPolicy", venv, verbose=1, n_steps=1024),
|
||||||
|
allow_variable_horizon=True,
|
||||||
)
|
)
|
||||||
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)
|
||||||
|
|||||||
@@ -47,7 +47,6 @@ def target_velocity_plan(current_v: float, target_v: float, t: int)
|
|||||||
# for now, constant acceleration
|
# for now, constant acceleration
|
||||||
a = (target_v - current_v) / t
|
a = (target_v - current_v) / t
|
||||||
return a*np.ones((t,))
|
return a*np.ones((t,))
|
||||||
Hello
|
|
||||||
|
|
||||||
def generate_plan(env, i):
|
def generate_plan(env, i):
|
||||||
"""Generate input profile for high-level action `i`."""
|
"""Generate input profile for high-level action `i`."""
|
||||||
|
|||||||
Reference in New Issue
Block a user