diff --git a/checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt b/checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt new file mode 100644 index 0000000..33df960 Binary files /dev/null and b/checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt differ diff --git a/evaluate_models.sh b/evaluate_models.sh index c4e5f72..edb84c4 100755 --- a/evaluate_models.sh +++ b/evaluate_models.sh @@ -27,3 +27,6 @@ python -m src.eval_main --method=ogail-ppo --policy_file='checkpoints/gail-ppo-o # SHAIL python -m src.eval_main --method=sgail --policy_file='checkpoints/sgail-options-setobs2.pt' --env='NormalizedSafeOptionsEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' + +# SHAIL-PPO +python -m src.eval_main --method=sgail-ppo --policy_file='checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt' --env='NormalizedSafeOptionsEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' diff --git a/src/eval_main.py b/src/eval_main.py index 381b061..4a68bcd 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -72,6 +72,10 @@ def load_policy(method:str, policy = ReparamPolicy(policy) policy.load_state_dict(torch.load(policy_file)) policy.eval() + elif method == 'sgail-ppo': + policy = SetMaskedDiscretePolicy(env.action_space.n) + policy.load_state_dict(torch.load(policy_file)) + policy.eval() else: raise NotImplementedError return policy diff --git a/src/safe_options/options.py b/src/safe_options/options.py index eeb5bae..87ab047 100644 --- a/src/safe_options/options.py +++ b/src/safe_options/options.py @@ -257,4 +257,4 @@ def NormalizedSafeOptionsEvalEnv(max_episode_steps=float('inf'), safe_actions_co n_rays=5, **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)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method), max_episode_steps=max_episode_steps) + ), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method), max_episode_steps=max_episode_steps)