diff --git a/checkpoints/bc-intersimple-setobs2.pt b/checkpoints/bc-intersimple-setobs2.pt new file mode 100644 index 0000000..944cd70 Binary files /dev/null and b/checkpoints/bc-intersimple-setobs2.pt differ diff --git a/checkpoints/gail-options-setobs2-Feb15_18-49-05.pt b/checkpoints/gail-options-setobs2-Feb15_18-49-05.pt new file mode 100644 index 0000000..76f70af Binary files /dev/null and b/checkpoints/gail-options-setobs2-Feb15_18-49-05.pt differ diff --git a/checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt b/checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt new file mode 100644 index 0000000..a2bd8b6 Binary files /dev/null and b/checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt differ diff --git a/checkpoints/wgail-options-setobs2-Feb16_01-06-27.pt b/checkpoints/wgail-options-setobs2-Feb16_01-06-27.pt new file mode 100644 index 0000000..ac6db8c Binary files /dev/null and b/checkpoints/wgail-options-setobs2-Feb16_01-06-27.pt differ diff --git a/checkpoints/wgail-ppo-options-setobs2-Feb16_04-02-56.pt b/checkpoints/wgail-ppo-options-setobs2-Feb16_04-02-56.pt new file mode 100644 index 0000000..255d031 Binary files /dev/null and b/checkpoints/wgail-ppo-options-setobs2-Feb16_04-02-56.pt differ diff --git a/evaluate_models.sh b/evaluate_models.sh index 830752f..5f64408 100755 --- a/evaluate_models.sh +++ b/evaluate_models.sh @@ -13,4 +13,11 @@ python -m src.eval_main # idm python -m src.eval_main --method=idm -python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-setobs2-15-02-2022.pt' --env='NormalizedOptionsEvalEnv' +# options GAIL +python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-setobs2-Feb15_18-49-05.pt' --env='NormalizedOptionsEvalEnv' --env_kwargs='{stop_on_collision:True}' + +# options GAIL-PPO +python -m src.eval_main --method=ogail-ppo --policy_file='checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt' --env='NormalizedOptionsEvalEnv' --env_kwargs='{stop_on_collision:True}' + +# behavior cloning +python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' diff --git a/src/evaluation/evaluation.py b/src/evaluation/evaluation.py index 3682759..20a6db9 100644 --- a/src/evaluation/evaluation.py +++ b/src/evaluation/evaluation.py @@ -122,7 +122,7 @@ class IntersimpleEvaluation: done = local_vars['done'] _agent = info['agent'] env = local_vars['env'].envs[venv_i] - assert isinstance(env, Intersimple) + # assert isinstance(env, Intersimple) self.eval_policy_step(info, done, _agent) diff --git a/src/gail2/envs.py b/src/gail2/envs.py index b9cef2b..6f6bdd4 100644 --- a/src/gail2/envs.py +++ b/src/gail2/envs.py @@ -25,11 +25,18 @@ def NormalizedOptionsEvalEnv(**kwargs): return OptionsEnv(Setobs( TransformObservation(IntersimpleLidarFlatIncrementingAgent( n_rays=5, - stop_on_collision=False, **kwargs, ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) ), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5)]) +def NormalizedContinuousEvalEnv(**kwargs): + return Setobs( + TransformObservation(IntersimpleLidarFlatIncrementingAgent( + n_rays=5, + **kwargs, + ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) + ) + class OptionsEnv(Wrapper): def __init__(self, env, options): @@ -65,7 +72,7 @@ class OptionsEnv(Wrapper): o, r, d, i = super().step(u) actions[k] = u rewards[k] = r - env_done[k+1] = d + env_done[k] = d infos.append(i) observations[k+1] = o @@ -84,7 +91,7 @@ class OptionsEnv(Wrapper): ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) hl_obs = ll_obs[ll_steps] hl_reward = (ll_rewards * ~ll_plan_done).sum().item() - hl_done = ll_env_done[ll_steps].item() + hl_done = ll_env_done[ll_steps-1].item() hl_infos = { 'll': { 'observations': ll_obs,