From 1916a8fe691c1099ac992eed5c89cb0efcb76cb4 Mon Sep 17 00:00:00 2001 From: Johannes Fischer Date: Tue, 3 Aug 2021 15:24:24 +0200 Subject: [PATCH] Implement metrics and write to tensorboard summary at test time --- src/main.py | 6 +++++- src/metrics.py | 29 ++++++++++++++++++++++++----- 2 files changed, 29 insertions(+), 6 deletions(-) diff --git a/src/main.py b/src/main.py index e69d52c..943f8a7 100644 --- a/src/main.py +++ b/src/main.py @@ -4,6 +4,7 @@ import gym import intersim import numpy as np from tqdm import tqdm +from torch.utils.tensorboard import SummaryWriter from src import InteractionDatasetSingleAgent, metrics from intersim.utils import get_map_path, get_svt @@ -63,7 +64,10 @@ def main(config, method='bc', train=False, test=False, loc=0, datadir='./expert_ # run test metrics test_dataset = InteractionDatasetSingleAgent(output_dir=datadir, loc=loc, tracks=[track]) - metrics(filestr, test_dataset, policy) + writer = SummaryWriter(filestr) + info = metrics(filestr, test_dataset, policy) + for k, m in info.items(): + writer.add_scalar('test/{}'.format(k), m, 0) def simulate_policy(policy, loc=0, track=0, filestr='', nframes=float('inf'), graph=None): diff --git a/src/metrics.py b/src/metrics.py index 5e4997f..56a157e 100644 --- a/src/metrics.py +++ b/src/metrics.py @@ -12,7 +12,10 @@ def metrics(filestr: str, test_dataset, policy): filestr (str): base string to outputs of a simulation test_dataset: a dataset held for testing policy: policy + Returns: + info (dict): metrics in a dictionary """ + info = {} # compute metrics using either # a) simulation files that were saved under the trained policy with prefix 'policy' @@ -27,26 +30,41 @@ def metrics(filestr: str, test_dataset, policy): # count collisions (from function in intersim.collisions) n_collisions = collisions.count_collisions_trajectory(states, lengths, widths) + info['n_collisions'] = n_collisions # calculate average velocity avg_v = average_velocity(states) + info['average_velocity'] = avg_v - # calculate divergence between velocity distributions - - # calcuate divergence between acceleration distributions + # convert policy dtype between float32 and float64 policy.policy = policy.policy.type(test_dataset[0]['state'].dtype) # generate actions in test dataset true_actions, pred_actions = [], [] + true_velocities = [] test_loader = DataLoader(test_dataset, batch_size=1024) with torch.no_grad(): for (batch_idx, batch) in enumerate(test_loader): pred_actions.append(policy(batch)) true_actions.append(batch['action']) + true_velocities.append(batch['state'][:,2]) - true_actions, pred_actions = torch.cat(true_actions,dim=0), torch.cat(pred_actions,dim=0) + true_actions, pred_actions = torch.cat(true_actions,dim=0), torch.cat(pred_actions, dim=0) visualize_distribution(true_actions[:,0], pred_actions[:,0], filestr+'_action_viz') + # calculate divergence between acceleration distributions + acceleration_kl = divergence(pred_actions, true_actions, type='kl', n_components=-1) + info['acceleration_kl'] = acceleration_kl + + # calculate divergence between velocity distributions + sim_velocities = states[:,:,2] + sim_velocities = sim_velocities[~torch.isnan(sim_velocities)].flatten() + true_velocities = torch.cat(true_velocities, dim=0) + velocity_kl = divergence(sim_velocities, true_velocities, type='kl', n_components=-1) + info['velocity_kl'] = velocity_kl + + return info + def visualize_distribution(true, pred, filestr): """ @@ -68,11 +86,12 @@ def average_velocity(states): """ Compute average of average velocity over all vehicles. Args: - states (torch.tensor): (T,nv,5) vehicle states + states (torch.tensor): (T,nv,5) vehicle states where T is the number of time steps and nv the number of vehicles Returns avg_v (float): average velocity """ velocities = states[:,:,2] + # average velocity per vehicle vehicle_avg_v = nanmean(velocities, dim=0) arg_v = nanmean(vehicle_avg_v) return arg_v