Add SHAIL-PPO
This commit is contained in:
BIN
checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt
Normal file
BIN
checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt
Normal file
Binary file not shown.
@@ -27,3 +27,6 @@ python -m src.eval_main --method=ogail-ppo --policy_file='checkpoints/gail-ppo-o
|
|||||||
|
|
||||||
# SHAIL
|
# 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}'
|
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}'
|
||||||
|
|||||||
@@ -72,6 +72,10 @@ def load_policy(method:str,
|
|||||||
policy = ReparamPolicy(policy)
|
policy = ReparamPolicy(policy)
|
||||||
policy.load_state_dict(torch.load(policy_file))
|
policy.load_state_dict(torch.load(policy_file))
|
||||||
policy.eval()
|
policy.eval()
|
||||||
|
elif method == 'sgail-ppo':
|
||||||
|
policy = SetMaskedDiscretePolicy(env.action_space.n)
|
||||||
|
policy.load_state_dict(torch.load(policy_file))
|
||||||
|
policy.eval()
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
return policy
|
return policy
|
||||||
|
|||||||
@@ -257,4 +257,4 @@ def NormalizedSafeOptionsEvalEnv(max_episode_steps=float('inf'), safe_actions_co
|
|||||||
n_rays=5,
|
n_rays=5,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
|
), 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)
|
||||||
|
|||||||
Reference in New Issue
Block a user