From 8c4ff03208a4bd12d3b4a935c87d665a9143fb3a Mon Sep 17 00:00:00 2001 From: Arec Date: Mon, 21 Feb 2022 00:06:39 -0800 Subject: [PATCH] adding average absolute delta v, and tracking positions and setting up architecture to implement rwse --- src/eval_main.py | 24 +++++++++++++++++++----- src/evaluation/evaluation.py | 8 ++++++-- src/evaluation/metrics.py | 23 +++++++++++++++++++++++ 3 files changed, 48 insertions(+), 7 deletions(-) diff --git a/src/eval_main.py b/src/eval_main.py index 4874f83..c303f29 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -7,7 +7,7 @@ from stable_baselines3.common.base_class import BaseAlgorithm from src.baselines import IDMRulePolicy from src.evaluation import IntersimpleEvaluation import src.gail.options as options_envs -from src.evaluation.metrics import divergence, visualize_distribution +from src.evaluation.metrics import divergence, visualize_distribution, rwse from src.evaluation.utils import save_metrics from src.core.policy import SetPolicy, SetDiscretePolicy from src.core.reparam_module import ReparamPolicy @@ -103,7 +103,7 @@ def form_expert_metrics(states:torch.Tensor, actions:torch.Tensor) ->Dict[str, l hard_brake = -3. timestep = 0.1 - keys = ['col_all','v_all', 'a_all','j_all', 'v_avg', 'a_avg', 'col', 'brake', 't'] + keys = ['col_all','x_all','y_all','v_all', 'a_all','j_all', 'v_avg', 'a_avg', 'col', 'brake', 't'] metrics = {key:[None]*nv for key in keys} for i in range(nv): @@ -111,6 +111,8 @@ def form_expert_metrics(states:torch.Tensor, actions:torch.Tensor) ->Dict[str, l nni = ~torch.isnan(states[:,i,0]) metrics['col_all'][i] = [False] * sum(nni) + metrics['x_all'][i] = states[nni,i,0].numpy() + metrics['y_all'][i] = states[nni,i,1].numpy() metrics['v_all'][i] = states[nni,i,2].numpy() metrics['a_all'][i] = actions[nni,i,0].numpy() @@ -243,7 +245,7 @@ def summary_metrics(metrics:List[Dict[str,list]]) -> Dict[str,float]: # average acceleration all_aalls = np.concatenate([np.concatenate(d['a_all']) for d in metrics]) - summary_metrics['mean acceleartion'] = np.mean(all_aalls) + summary_metrics['mean acceleration'] = np.mean(all_aalls) # average +acceleration pos_accels = all_aalls[all_aalls>0] @@ -261,11 +263,11 @@ def summary_metrics(metrics:List[Dict[str,list]]) -> Dict[str,float]: summary_metrics['mean |jerk|'] = np.mean(np.abs(all_jerks)) # collision rate - all_collisions = sum([d['col'] for d in metrics],[]) # aggregate to single list + all_collisions = sum([d['col'] for d in metrics],[]) # aggregate to single list summary_metrics['collision rate'] = sum(all_collisions)/len(all_collisions) # hard brake rate - all_hard_brakes = sum([d['brake'] for d in metrics],[]) # aggregate to single list + all_hard_brakes = sum([d['brake'] for d in metrics],[]) # aggregate to single list summary_metrics['hard brake rate'] = sum(all_hard_brakes)/len(all_hard_brakes) # average number of timesteps @@ -294,11 +296,23 @@ def comparison_metrics(policy_metrics:List[Dict[str,list]], """ comparison_metrics = {} + # rwse + expert_traj, policy_traj = [], [] + for iR in range(len(policy_metrics)): + for iTraj in range(len(policy_metrics[iR]['x_all'])): + 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) + # average velocity shortfall expert_vavg = np.array(sum([d['v_avg'] for d in expert_metrics],[])) policy_vavg = np.array(sum([d['v_avg'] for d in policy_metrics],[])) assert len(expert_vavg)==len(policy_vavg) comparison_metrics['mean shortfall velocity'] = np.mean(expert_vavg - policy_vavg) + + # Average |Delta V average| + comparison_metrics['average absolute average velocity'] = np.mean(np.abs(expert_vavg - policy_vavg)) # velocity JSD expert_vs = torch.tensor(np.concatenate([np.concatenate(d['v_all']) for d in expert_metrics])) diff --git a/src/evaluation/evaluation.py b/src/evaluation/evaluation.py index 96b8437..b63d8cd 100644 --- a/src/evaluation/evaluation.py +++ b/src/evaluation/evaluation.py @@ -39,7 +39,7 @@ class IntersimpleEvaluation: self.is_options_env = isinstance(self.env, (OptionsEnv, OptionsTimeLimit)) # metrics present on every step of every episode - self.metric_keys_all = ['v_all', 'a_all', 'col_all'] + self.metric_keys_all = ['x_all', 'y_all', 'v_all', 'a_all', 'col_all'] # metrics calculated after the fact, with one per episode self.metric_keys_single = ['j_all', 'v_avg','a_avg', 'col','brake', 't'] @@ -129,6 +129,8 @@ class IntersimpleEvaluation: def eval_policy_step(self, info, done, _agent): # Increase collision counter if episode terminated with a collision + self._metrics['x_all'][_agent].append(info['prev_state'][_agent,0].item()) + self._metrics['y_all'][_agent].append(info['prev_state'][_agent,1].item()) self._metrics['v_all'][_agent].append(info['prev_state'][_agent,2].item()) self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item()) col = info['collision'] @@ -144,10 +146,12 @@ class IntersimpleEvaluation: """ Postprocess and metrics after simulation episodes """ - # self.metric_keys_all = ['v_all', 'a_all', 'col_all'] + # self.metric_keys_all = ['x_all','y_all','v_all', 'a_all', 'col_all'] # self.metric_keys_single = ['j_all', 'v_avg','a_avg', 'col','brake', 't'] for i in range(self.n_episodes): + self._metrics['x_all'][i] = np.array(self._metrics['x_all'][i]) + self._metrics['y_all'][i] = np.array(self._metrics['y_all'][i]) self._metrics['v_all'][i] = np.array(self._metrics['v_all'][i]) self._metrics['a_all'][i] = np.array(self._metrics['a_all'][i]) diff --git a/src/evaluation/metrics.py b/src/evaluation/metrics.py index 5ef5d30..477a22b 100644 --- a/src/evaluation/metrics.py +++ b/src/evaluation/metrics.py @@ -4,8 +4,31 @@ import numpy as np import matplotlib.pyplot as plt from torch.utils.data import DataLoader from intersim import collisions +from typing import List # import tikzplotlib +def rwse(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> float: + """ + Calculated time-weighted + + 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 (float): rwse of positions + """ + assert len(expert) == len(policy) + + # calculate rwse + + return 0.0 + + + def visualize_distribution(expert, policy, filestr): """ Visualize two distributions