Add SHAIL-PPO

This commit is contained in:
ebuehrle
2022-02-18 06:54:52 +01:00
parent 1624e1a349
commit 84351e77f2
4 changed files with 8 additions and 1 deletions

Binary file not shown.

View File

@@ -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}'

View File

@@ -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

View File

@@ -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)