wrapping all environments in timelimit to stop runs longer than 100s, since some others were erroring

This commit is contained in:
Arec
2022-02-21 17:54:29 -08:00
parent e7f8385628
commit e7f4f6a871
7 changed files with 49 additions and 35 deletions

View File

@@ -14,26 +14,26 @@ python -m src.eval_main
python -m src.eval_main --method=idm python -m src.eval_main --method=idm
# behavior cloning # behavior cloning
python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=0 python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=0
#python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=1 #python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=1
#python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=2 #python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=2
#python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=3 #python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=3
#python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=4 #python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=4
#python -m src.evaluation.utils load_and_average out/bc #python -m src.evaluation.utils load_and_average out/bc
# GAIL # GAIL
python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=0 python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=0
#python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=1 #python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=1
#python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=2 #python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=2
#python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=3 #python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=3
#python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=4 #python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=4
#python -m src.evaluation.utils load_and_average out/gail #python -m src.evaluation.utils load_and_average out/gail
# options GAIL # 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}' 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,max_episode_steps:1000}'
# options GAIL-PPO # 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}' 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,max_episode_steps:1000}'
# SHAIL # SHAIL
python -m src.eval_main --method=sgail --policy_file='checkpoints/sgail-options-setobs2-Feb21_13-30-45.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-Feb21_13-30-45.pt' --env='NormalizedSafeOptionsEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}'
@@ -51,26 +51,26 @@ python -m src.eval_main --locations='[(0,4)]'
python -m src.eval_main --method=idm --locations='[(0,4)]' python -m src.eval_main --method=idm --locations='[(0,4)]'
# behavior cloning # behavior cloning
python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=0 --locations='[(0,4)]' python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=0 --locations='[(0,4)]'
#python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=1 --locations='[(0,4)]' #python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=1 --locations='[(0,4)]'
#python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=2 --locations='[(0,4)]' #python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=2 --locations='[(0,4)]'
#python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=3 --locations='[(0,4)]' #python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=3 --locations='[(0,4)]'
#python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=4 --locations='[(0,4)]' #python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=4 --locations='[(0,4)]'
#python -m src.evaluation.utils load_and_average out/bc #python -m src.evaluation.utils load_and_average out/bc
# GAIL # GAIL
python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=0 --locations='[(0,4)]' python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=0 --locations='[(0,4)]'
#python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=1 --locations='[(0,4)]' #python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=1 --locations='[(0,4)]'
#python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=2 --locations='[(0,4)]' #python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=2 --locations='[(0,4)]'
#python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=3 --locations='[(0,4)]' #python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=3 --locations='[(0,4)]'
#python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' --seed=4 --locations='[(0,4)]' #python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --seed=4 --locations='[(0,4)]'
#python -m src.evaluation.utils load_and_average out/gail #python -m src.evaluation.utils load_and_average out/gail
# options GAIL # 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}' --locations='[(0,4)]' 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,max_episode_steps:1000}' --locations='[(0,4)]'
# options GAIL-PPO # 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}' --locations='[(0,4)]' 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,max_episode_steps:1000}' --locations='[(0,4)]'
# SHAIL # SHAIL
python -m src.eval_main --method=sgail --policy_file='checkpoints/sgail-options-setobs2-Feb21_13-30-45.pt' --env='NormalizedSafeOptionsEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --locations='[(0,4)]' python -m src.eval_main --method=sgail --policy_file='checkpoints/sgail-options-setobs2-Feb21_13-30-45.pt' --env='NormalizedSafeOptionsEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' --locations='[(0,4)]'

View File

@@ -14,6 +14,7 @@ from src.core.reparam_module import ReparamPolicy, ReparamSafePolicy
from src.options import envs as options_envs2 from src.options import envs as options_envs2
from src.safe_options.policy import SetMaskedDiscretePolicy from src.safe_options.policy import SetMaskedDiscretePolicy
from src.safe_options import options as options_envs3 from src.safe_options import options as options_envs3
from src.util.wrappers import IntersimpleTimeLimit
from typing import Optional, List, Dict, Tuple from typing import Optional, List, Dict, Tuple
import torch import torch
@@ -214,6 +215,13 @@ def evaluate_policy(locations:List[Tuple[int,int]],
# initialize environment # initialize environment
Env = envs_dict[env_class] Env = envs_dict[env_class]
# wrap in TimeLimit
if 'max_episode_steps' in it_env_kwargs.keys():
steps = it_env_kwargs.pop('max_episode_steps')
eval_env = IntersimpleTimeLimit(Env(**it_env_kwargs),
max_episode_steps=steps)
else:
eval_env = Env(**it_env_kwargs) eval_env = Env(**it_env_kwargs)
evaluator = IntersimpleEvaluation(eval_env) evaluator = IntersimpleEvaluation(eval_env)
@@ -364,7 +372,8 @@ def eval_main(
np.random.seed(seed) np.random.seed(seed)
torch.manual_seed(seed) torch.manual_seed(seed)
pfilename = policy_file.split('/')[-1].split('.')[0] pfilename = policy_file.split('/')[-1].split('.')[0]
outbase = f'out/{method}/{pfilename}_seed{seed}' locstr = 'loc_'+'_'.join([f'r{ro}t{tr}' for (ro,tr) in locations])
outbase = f'out/{method}/{locstr}/{pfilename}_seed{seed}'
# load expert metrics # load expert metrics
expert_metrics = generate_expert_metrics(locations) expert_metrics = generate_expert_metrics(locations)

View File

@@ -6,8 +6,9 @@ from typing import Callable, Dict, Optional
import os import os
import pickle import pickle
from tqdm import tqdm from tqdm import tqdm
from src.util.wrappers import IntersimpleTimeLimit
from src.options.envs import OptionsEnv from src.options.envs import OptionsEnv
from src.util.wrappers import OptionsTimeLimit from src.safe_options.options import SafeOptionsEnv
class IntersimpleEvaluation: class IntersimpleEvaluation:
""" """
@@ -36,7 +37,10 @@ class IntersimpleEvaluation:
self.env = eval_env self.env = eval_env
self.n_episodes = eval_env.nv self.n_episodes = eval_env.nv
self.use_pbar = use_pbar self.use_pbar = use_pbar
self.is_options_env = isinstance(self.env, (OptionsEnv, OptionsTimeLimit)) if isinstance(self.env, IntersimpleTimeLimit):
self.is_options_env = isinstance(self.env.env, (OptionsEnv, SafeOptionsEnv))
else:
self.is_options_env = isinstance(self.env, (OptionsEnv, SafeOptionsEnv))
# metrics present on every step of every episode # metrics present on every step of every episode
self.metric_keys_all = ['x_all', 'y_all', 'v_all', 'a_all', 'col_all'] self.metric_keys_all = ['x_all', 'y_all', 'v_all', 'a_all', 'col_all']

View File

@@ -24,7 +24,7 @@ def rwse(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> Dict
assert len(expert) == len(policy) assert len(expert) == len(policy)
# calculate rwse # calculate rwse
times = [1,2,5,10,15,20] times = [1,2,5,10,15,20,25,30]
time_indices = [int(t/dt) for t in times] time_indices = [int(t/dt) for t in times]
rwse_dict_keys = [f'rwse_{t}s' for t in times]+['rwse_end'] rwse_dict_keys = [f'rwse_{t}s' for t in times]+['rwse_end']
se_dict = {key:[] for key in rwse_dict_keys} se_dict = {key:[] for key in rwse_dict_keys}

View File

@@ -2,6 +2,7 @@ import pickle
import os import os
import numpy as np import numpy as np
from typing import List,Dict from typing import List,Dict
def save_metrics(metrics:dict, filestr:str): def save_metrics(metrics:dict, filestr:str):
""" """
Save metric dict to filestr Save metric dict to filestr

View File

@@ -14,7 +14,7 @@ from src.options.envs import OptionsEnv
from src.safe_options.collisions import feasible from src.safe_options.collisions import feasible
from intersim.envs import IntersimpleLidarFlatIncrementingAgent from intersim.envs import IntersimpleLidarFlatIncrementingAgent
from src.util.wrappers import OptionsTimeLimit, Setobs, TransformObservation from src.util.wrappers import Setobs, TransformObservation
@dataclass @dataclass
class Buffer: class Buffer:
@@ -251,10 +251,10 @@ obs_max = np.array([
[50, np.pi, 20, 20, np.pi, 1e-1], [50, np.pi, 20, 20, np.pi, 1e-1],
]).reshape(-1) ]).reshape(-1)
def NormalizedSafeOptionsEvalEnv(max_episode_steps=float('inf'), safe_actions_collision_method=None, abort_unsafe_collision_method=None, **kwargs): def NormalizedSafeOptionsEvalEnv(safe_actions_collision_method=None, abort_unsafe_collision_method=None, **kwargs):
return OptionsTimeLimit(SafeOptionsEnv(Setobs( return SafeOptionsEnv(Setobs(
TransformObservation(IntersimpleLidarFlatIncrementingAgent( TransformObservation(IntersimpleLidarFlatIncrementingAgent(
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), (10, 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)

View File

@@ -9,7 +9,7 @@ class TransformObservation(gym.wrappers.TransformObservation):
def __getattr__(self, name): def __getattr__(self, name):
return getattr(self.env, name) return getattr(self.env, name)
class OptionsTimeLimit(gym.wrappers.TimeLimit): class IntersimpleTimeLimit(gym.wrappers.TimeLimit):
def __getattr__(self, name): def __getattr__(self, name):
return getattr(self.env, name) return getattr(self.env, name)