From e7f838562821caba4c9475d3f47ef8c553853cb0 Mon Sep 17 00:00:00 2001 From: Arec Date: Mon, 21 Feb 2022 15:55:16 -0800 Subject: [PATCH] updating rwse to work at different times, updating correct testing environment from roundabout, removing the assertion that a collision implies done in the evaluator, using nanmean and nanstd in averaging --- evaluate_models.sh | 57 +++++++++++++++++++++++++++++------- src/eval_main.py | 8 +++-- src/evaluation/evaluation.py | 4 +-- src/evaluation/metrics.py | 54 +++++++++++++++++++++++++++++++--- src/evaluation/utils.py | 4 +-- 5 files changed, 106 insertions(+), 21 deletions(-) diff --git a/evaluate_models.sh b/evaluate_models.sh index 6e9ec23..4fcfd9b 100755 --- a/evaluate_models.sh +++ b/evaluate_models.sh @@ -15,19 +15,19 @@ python -m src.eval_main --method=idm # 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}' --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}' --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.evaluation.utils load_and_average out/bc +#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}' --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}' --seed=4 +#python -m src.evaluation.utils load_and_average out/bc # 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}' --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}' --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.evaluation.utils load_and_average out/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=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}' --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.evaluation.utils load_and_average out/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}' @@ -40,3 +40,40 @@ python -m src.eval_main --method=sgail --policy_file='checkpoints/sgail-options- # 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}' + + + +### Same checkpoint files applied out of distribution (e.g. to track 5) +# expert +python -m src.eval_main --locations='[(0,4)]' + +# idm +python -m src.eval_main --method=idm --locations='[(0,4)]' + +# 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}' --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}' --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.evaluation.utils load_and_average out/bc + +# 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}' --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}' --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.evaluation.utils load_and_average out/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)]' + +# 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)]' + +# 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)]' + +# 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}' --locations='[(0,4)]' diff --git a/src/eval_main.py b/src/eval_main.py index 2e38932..c5bb771 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -214,7 +214,7 @@ def evaluate_policy(locations:List[Tuple[int,int]], # initialize environment Env = envs_dict[env_class] - eval_env = Env(**env_kwargs) + eval_env = Env(**it_env_kwargs) evaluator = IntersimpleEvaluation(eval_env) # load policy @@ -303,7 +303,7 @@ def comparison_metrics(policy_metrics:List[Dict[str,list]], expert_traj.append(np.vstack((expert_metrics[iR]['x_all'][iTraj], expert_metrics[iR]['y_all'][iTraj]))) policy_traj.append(np.vstack((policy_metrics[iR]['x_all'][iTraj], policy_metrics[iR]['y_all'][iTraj]))) assert len(expert_traj)==len(policy_traj) - comparison_metrics['rwse'] = rwse(expert_traj, policy_traj) + comparison_metrics.update(rwse(expert_traj, policy_traj)) # average velocity shortfall expert_vavg = np.array(sum([d['v_avg'] for d in expert_metrics],[])) @@ -356,7 +356,9 @@ def eval_main( env (str): environment class method (str): method (expert, bc, gail, rail, hgail, hrail) """ - print(f'Evaluating {method} on {env}') + print(f'#############################################################################') + print(f'Evaluating {method} from file {policy_file} on {env} at locations {locations}') + print(f'#############################################################################') # set seed np.random.seed(seed) diff --git a/src/evaluation/evaluation.py b/src/evaluation/evaluation.py index b63d8cd..d0065fc 100644 --- a/src/evaluation/evaluation.py +++ b/src/evaluation/evaluation.py @@ -135,8 +135,8 @@ class IntersimpleEvaluation: self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item()) col = info['collision'] - if col: - assert done + # if col: # commenting out if we dont want to end on collision + # assert done self._metrics['col_all'][_agent].append(col) if done and self.use_pbar: diff --git a/src/evaluation/metrics.py b/src/evaluation/metrics.py index 4098229..46eda94 100644 --- a/src/evaluation/metrics.py +++ b/src/evaluation/metrics.py @@ -4,10 +4,58 @@ import numpy as np import matplotlib.pyplot as plt from torch.utils.data import DataLoader from intersim import collisions -from typing import List +from typing import List, Dict # import tikzplotlib -def rwse(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> float: +def rwse(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> Dict[str,float]: + """ + Calculate average mean squared displacement error + + Args: + expert (List[np.ndarray]): all position trajectories for all expert rollouts + policy (List[np.ndarray]): all position trajectories for all policy rollouts + + each trajectory in the list should have shape (2, T). however expert[i] might have a + different T than policy[i] + + Returns + rwse_dict (Dict[str,float]): dict of different RWSEs + """ + assert len(expert) == len(policy) + + # calculate rwse + times = [1,2,5,10,15,20] + time_indices = [int(t/dt) for t in times] + rwse_dict_keys = [f'rwse_{t}s' for t in times]+['rwse_end'] + se_dict = {key:[] for key in rwse_dict_keys} + for expert_trajectory, policy_trajectory in zip(expert, policy): + _, T1 = expert_trajectory.shape + _, T2 = policy_trajectory.shape + minT = min(T1, T2) + + crop_expert_trajectory = expert_trajectory[:, :minT] + crop_policy_trajectory = policy_trajectory[:, :minT] + + # square error along every time + se = ((crop_policy_trajectory - crop_expert_trajectory)**2).sum(0) + + # add to dict with appropriate indexing + for time, idx in zip(times, time_indices): + if minT >= idx: + se_dict[f'rwse_{time}s'].append(se[idx-1]) + se_dict['rwse_end'].append(se[-1]) + + assert len(se_dict['rwse_end']) == len(expert) + + # print how many trajectories of each time: + for key in rwse_dict_keys: + print('%s has %i elements'%(key, len(se_dict[key]))) + + rwse_dict = {key:np.mean(np.array(se_dict[key]))**0.5 for key in rwse_dict_keys} + + return rwse_dict + +def rwse_basic(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> float: """ Calculate average mean squared displacement error @@ -42,8 +90,6 @@ def rwse(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> floa return avg_rwse - - def visualize_distribution(expert, policy, filestr): """ Visualize two distributions diff --git a/src/evaluation/utils.py b/src/evaluation/utils.py index 7c78ca5..aa72b53 100644 --- a/src/evaluation/utils.py +++ b/src/evaluation/utils.py @@ -47,8 +47,8 @@ def average_metrics(metric_list:List[Dict[str,float]]): for i in range(N): master_dict[key].append(metric_list[i][key]) master_dict[key] = np.array(master_dict[key]) - mu = np.mean(master_dict[key]) - std2 = np.std(master_dict[key])*2 + mu = np.nanmean(master_dict[key]) + std2 = np.nanstd(master_dict[key])*2 print(f'{key}: {mu} \pm {std2}')