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

This commit is contained in:
Arec
2022-02-21 15:55:16 -08:00
parent daa4825f17
commit e7f8385628
5 changed files with 106 additions and 21 deletions

View File

@@ -15,19 +15,19 @@ 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}' --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=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=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=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}' --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}' --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=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=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=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}' --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}'
@@ -40,3 +40,40 @@ python -m src.eval_main --method=sgail --policy_file='checkpoints/sgail-options-
# SHAIL-PPO # 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}' 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)]'

View File

@@ -214,7 +214,7 @@ def evaluate_policy(locations:List[Tuple[int,int]],
# initialize environment # initialize environment
Env = envs_dict[env_class] Env = envs_dict[env_class]
eval_env = Env(**env_kwargs) eval_env = Env(**it_env_kwargs)
evaluator = IntersimpleEvaluation(eval_env) evaluator = IntersimpleEvaluation(eval_env)
# load policy # 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]))) 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]))) policy_traj.append(np.vstack((policy_metrics[iR]['x_all'][iTraj], policy_metrics[iR]['y_all'][iTraj])))
assert len(expert_traj)==len(policy_traj) 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 # average velocity shortfall
expert_vavg = np.array(sum([d['v_avg'] for d in expert_metrics],[])) expert_vavg = np.array(sum([d['v_avg'] for d in expert_metrics],[]))
@@ -356,7 +356,9 @@ def eval_main(
env (str): environment class env (str): environment class
method (str): method (expert, bc, gail, rail, hgail, hrail) 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 # set seed
np.random.seed(seed) np.random.seed(seed)

View File

@@ -135,8 +135,8 @@ class IntersimpleEvaluation:
self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item()) self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item())
col = info['collision'] col = info['collision']
if col: # if col: # commenting out if we dont want to end on collision
assert done # assert done
self._metrics['col_all'][_agent].append(col) self._metrics['col_all'][_agent].append(col)
if done and self.use_pbar: if done and self.use_pbar:

View File

@@ -4,10 +4,58 @@ import numpy as np
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
from torch.utils.data import DataLoader from torch.utils.data import DataLoader
from intersim import collisions from intersim import collisions
from typing import List from typing import List, Dict
# import tikzplotlib # 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 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 return avg_rwse
def visualize_distribution(expert, policy, filestr): def visualize_distribution(expert, policy, filestr):
""" """
Visualize two distributions Visualize two distributions

View File

@@ -47,8 +47,8 @@ def average_metrics(metric_list:List[Dict[str,float]]):
for i in range(N): for i in range(N):
master_dict[key].append(metric_list[i][key]) master_dict[key].append(metric_list[i][key])
master_dict[key] = np.array(master_dict[key]) master_dict[key] = np.array(master_dict[key])
mu = np.mean(master_dict[key]) mu = np.nanmean(master_dict[key])
std2 = np.std(master_dict[key])*2 std2 = np.nanstd(master_dict[key])*2
print(f'{key}: {mu} \pm {std2}') print(f'{key}: {mu} \pm {std2}')